Initial commit

This commit is contained in:
2026-06-30 15:14:37 +08:00
commit 15dab96872
311 changed files with 95639 additions and 0 deletions
+93
View File
@@ -0,0 +1,93 @@
package connect
import (
"context"
"fmt"
"net"
"time"
"go.uber.org/zap"
)
// DirectFactory Direct-UDP 直连工厂(Layer 1
type DirectFactory struct {
stunServers []string
logger *zap.Logger
}
// NewDirectFactory 创建 Direct-UDP 工厂
func NewDirectFactory(stunServers []string, logger *zap.Logger) *DirectFactory {
return &DirectFactory{
stunServers: stunServers,
logger: logger,
}
}
// Layer 返回传输层类型
func (f *DirectFactory) Layer() Layer {
return LayerDirectUDP
}
// Name 返回传输方式名称
func (f *DirectFactory) Name() string {
return "Direct-UDP"
}
// Dial 建立 Direct-UDP 直连
func (f *DirectFactory) Dial(ctx context.Context, config *DialConfig) (net.Conn, error) {
f.logger.Info("开始建立 Direct-UDP 直连",
zap.String("peer_id", config.PeerID))
servers := f.stunServers
if len(servers) == 0 {
servers = config.STUNServers
}
if len(servers) == 0 {
f.logger.Warn("未配置 STUN 服务器列表,仅尝试内部 P2P 打洞")
}
var candidates []string
if len(servers) > 0 {
// 1. 创建 STUN 客户端收集候选地址
stun := NewSTUNClient(servers, f.logger)
candidates = stun.CollectCandidates()
}
if len(candidates) == 0 {
f.logger.Warn("未能收集到任何 STUN 候选地址,回退至 PeerID")
// As a fallback, maybe PeerID contains IP:PORT
candidates = append(candidates, config.PeerID)
}
f.logger.Info("STUN 候选地址收集完成",
zap.Strings("candidates", candidates))
// 2. 实际 P2P 连接尝试
f.logger.Warn("当前尝试所有候选地址...")
dialer := &net.Dialer{
Timeout: 5 * time.Second,
}
// 建立 UDP 连接并返回最近可用的
var lastErr error
for _, candidate := range candidates {
if candidate == "" { continue }
conn, err := dialer.DialContext(ctx, "udp", candidate)
if err == nil {
f.logger.Info("Direct-UDP 直连建立成功",
zap.String("peer_id", config.PeerID),
zap.String("remote_addr", conn.RemoteAddr().String()))
return conn, nil
}
f.logger.Warn("候选地址连接失败",
zap.String("candidate", candidate),
zap.Error(err))
lastErr = err
}
return nil, fmt.Errorf("所有候选地址 UDP 连接均失败,最后错误:%w", lastErr)
}
+219
View File
@@ -0,0 +1,219 @@
package connect
import (
"context"
"encoding/binary"
"fmt"
"io"
"net"
"sync"
"time"
"go.uber.org/zap"
)
// FakeTCPConn FakeTCP 连接(UDP 包封装为 TCP 流)
type FakeTCPConn struct {
conn net.Conn
mu sync.Mutex
closed bool
readBuffer []byte
}
// NewFakeTCPConn 创建 FakeTCP 连接
func NewFakeTCPConn(conn net.Conn) *FakeTCPConn {
return &FakeTCPConn{
conn: conn,
readBuffer: make([]byte, 0),
}
}
// Read 读取数据(带长度前缀解析)
func (c *FakeTCPConn) Read(b []byte) (int, error) {
c.mu.Lock()
defer c.mu.Unlock()
// 如果缓冲区有数据,直接返回
if len(c.readBuffer) > 0 {
n := copy(b, c.readBuffer)
c.readBuffer = c.readBuffer[n:]
return n, nil
}
// 读取长度前缀(4 字节)
var length uint32
if err := binary.Read(c.conn, binary.BigEndian, &length); err != nil {
return 0, err
}
// 限制最大长度(防止恶意攻击)
if length > 65535 {
return 0, fmt.Errorf("packet too large: %d bytes", length)
}
// 读取实际数据
data := make([]byte, length)
if _, err := io.ReadFull(c.conn, data); err != nil {
return 0, err
}
// 返回请求的数据
n := copy(b, data)
if n < len(data) {
// 剩余数据存入缓冲区
c.readBuffer = data[n:]
}
return n, nil
}
// Write 写入数据(添加 4 字节长度前缀)
func (c *FakeTCPConn) Write(b []byte) (int, error) {
c.mu.Lock()
defer c.mu.Unlock()
if c.closed {
return 0, fmt.Errorf("connection closed")
}
// 写入长度前缀
length := uint32(len(b))
if err := binary.Write(c.conn, binary.BigEndian, length); err != nil {
return 0, err
}
// 写入实际数据
n, err := c.conn.Write(b)
return n, err
}
// Close 关闭连接
func (c *FakeTCPConn) Close() error {
c.mu.Lock()
defer c.mu.Unlock()
c.closed = true
return c.conn.Close()
}
// LocalAddr 本地地址
func (c *FakeTCPConn) LocalAddr() net.Addr {
return c.conn.LocalAddr()
}
// RemoteAddr 远程地址
func (c *FakeTCPConn) RemoteAddr() net.Addr {
return c.conn.RemoteAddr()
}
// SetDeadline 设置截止时间
func (c *FakeTCPConn) SetDeadline(t time.Time) error {
return c.conn.SetDeadline(t)
}
// SetReadDeadline 设置读截止时间
func (c *FakeTCPConn) SetReadDeadline(t time.Time) error {
return c.conn.SetReadDeadline(t)
}
// SetWriteDeadline 设置写截止时间
func (c *FakeTCPConn) SetWriteDeadline(t time.Time) error {
return c.conn.SetWriteDeadline(t)
}
// DialFakeTCP 拨号 FakeTCP 连接
func DialFakeTCP(ctx context.Context, network, addr string, logger *zap.Logger) (net.Conn, error) {
logger.Debug("dialing FakeTCP", zap.String("addr", addr))
// 建立 TCP 连接
conn, err := (&net.Dialer{}).DialContext(ctx, network, addr)
if err != nil {
return nil, fmt.Errorf("failed to dial TCP: %w", err)
}
// 包装为 FakeTCP 连接
return NewFakeTCPConn(conn), nil
}
// ListenFakeTCP 监听 FakeTCP 端口
func ListenFakeTCP(network, addr string, logger *zap.Logger) (net.Listener, error) {
logger.Info("listening FakeTCP", zap.String("addr", addr))
// 监听 TCP 端口
listener, err := net.Listen(network, addr)
if err != nil {
return nil, fmt.Errorf("failed to listen TCP: %w", err)
}
return &fakeTCPListener{
Listener: listener,
logger: logger,
}, nil
}
// fakeTCPListener FakeTCP 监听器
type fakeTCPListener struct {
net.Listener
logger *zap.Logger
}
// Accept 接受连接并包装为 FakeTCPConn
func (l *fakeTCPListener) Accept() (net.Conn, error) {
conn, err := l.Listener.Accept()
if err != nil {
return nil, err
}
l.logger.Debug("accepted FakeTCP connection", zap.String("addr", conn.RemoteAddr().String()))
return NewFakeTCPConn(conn), nil
}
// FakeTCPFactory FakeTCP 传输工厂
type FakeTCPFactory struct {
logger *zap.Logger
}
// NewFakeTCPFactory 创建 FakeTCP 工厂
func NewFakeTCPFactory(logger *zap.Logger) *FakeTCPFactory {
return &FakeTCPFactory{
logger: logger,
}
}
// Layer 返回传输层类型
func (f *FakeTCPFactory) Layer() Layer {
return LayerFakeTCP
}
// Name 返回名称
func (f *FakeTCPFactory) Name() string {
return "FakeTCP"
}
// Dial 建立 FakeTCP 连接
func (f *FakeTCPFactory) Dial(ctx context.Context, config *DialConfig) (net.Conn, error) {
f.logger.Info("开始建立 FakeTCP 连接",
zap.String("peer_id", config.PeerID))
// 1. 解析对端地址(PeerID 格式应为 "ip:port"
if config.PeerID == "" {
return nil, fmt.Errorf("PeerID 为空")
}
// 2. 建立 TCP 连接
dialer := &net.Dialer{Timeout: config.Timeout}
conn, err := dialer.DialContext(ctx, "tcp", config.PeerID)
if err != nil {
return nil, fmt.Errorf("TCP 连接失败:%w", err)
}
// 3. 包装为 FakeTCP 连接(UDP 包封装为 TCP 流)
fakeConn := NewFakeTCPConn(conn)
f.logger.Info("FakeTCP 连接建立成功",
zap.String("peer_id", config.PeerID),
zap.String("local_addr", conn.LocalAddr().String()),
zap.String("remote_addr", conn.RemoteAddr().String()))
return fakeConn, nil
}
+585
View File
@@ -0,0 +1,585 @@
package connect
import (
"context"
"encoding/json"
"fmt"
"io"
"net"
"sync"
"time"
"github.com/pion/webrtc/v3"
"go.uber.org/zap"
)
// ICEConfig ICE 配置
type ICEConfig struct {
STUNServers []string
TURNServers []TURNServerConfig
}
// TURNServerConfig TURN 服务器配置
type TURNServerConfig struct {
URLs []string
Username string
Credential string
}
// ICEServer ICE 服务器配置(别名,保持兼容)
type ICEICEServer = TURNServerConfig
// ICEClient ICE 客户端(ICE协商 + WebRTC DataChannel
type ICEClient struct {
config *ICEConfig
logger *zap.Logger
api *webrtc.API
peerConns map[string]*webrtc.PeerConnection // peerID -> PeerConnection
dataChannels map[string]*webrtc.DataChannel // peerID -> DataChannel
signalingCh map[string]chan SignalMessage // peerID -> 信令通道
mu sync.RWMutex
onSignal func(peerID string, signal SignalMessage) // 信令回调
}
// SignalMessage 信令消息
type SignalMessage struct {
Type string `json:"type"` // "offer" | "answer" | "candidate"
SDP string `json:"sdp,omitempty"`
Candidate string `json:"candidate,omitempty"`
Target string `json:"target"` // 目标 PeerID
Source string `json:"source"` // 来源 PeerID
}
// NewICEClient 创建 ICE 客户端
func NewICEClient(config *ICEConfig, logger *zap.Logger) *ICEClient {
// 创建 WebRTC API(使用默认配置)
api := webrtc.NewAPI()
return &ICEClient{
config: config,
logger: logger,
api: api,
peerConns: make(map[string]*webrtc.PeerConnection),
dataChannels: make(map[string]*webrtc.DataChannel),
signalingCh: make(map[string]chan SignalMessage),
}
}
// createPeerConnection 创建 PeerConnection
func (c *ICEClient) createPeerConnection(peerID string) (*webrtc.PeerConnection, error) {
// 构建 ICE 服务器配置
var iceServers []webrtc.ICEServer
// 添加 STUN 服务器
for _, stun := range c.config.STUNServers {
iceServers = append(iceServers, webrtc.ICEServer{
URLs: []string{stun},
})
}
// 添加 TURN 服务器
for _, turn := range c.config.TURNServers {
iceServers = append(iceServers, webrtc.ICEServer{
URLs: turn.URLs,
Username: turn.Username,
Credential: turn.Credential,
})
}
// 创建 PeerConnection 配置
config := webrtc.Configuration{
ICEServers: iceServers,
}
// 创建 PeerConnection
pc, err := c.api.NewPeerConnection(config)
if err != nil {
return nil, fmt.Errorf("创建 PeerConnection 失败:%w", err)
}
// 存储 PeerConnection
c.mu.Lock()
c.peerConns[peerID] = pc
c.mu.Unlock()
c.logger.Info("创建 PeerConnection",
zap.String("peer_id", peerID),
zap.Int("ice_servers", len(iceServers)))
return pc, nil
}
// CreateOffer 创建 Offer(主动发起方)
func (c *ICEClient) CreateOffer(ctx context.Context, peerID string) (*SignalMessage, error) {
pc, err := c.createPeerConnection(peerID)
if err != nil {
return nil, err
}
// 创建 DataChannel
dc, err := pc.CreateDataChannel("meshray", nil)
if err != nil {
pc.Close()
return nil, fmt.Errorf("创建 DataChannel 失败:%w", err)
}
// 设置 DataChannel 处理器
dc.OnOpen(func() {
c.logger.Info("DataChannel 已打开", zap.String("peer_id", peerID))
})
dc.OnClose(func() {
c.logger.Info("DataChannel 已关闭", zap.String("peer_id", peerID))
})
// 存储 DataChannel
c.mu.Lock()
c.dataChannels[peerID] = dc
c.mu.Unlock()
// 创建 Offer
offer, err := pc.CreateOffer(nil)
if err != nil {
pc.Close()
return nil, fmt.Errorf("创建 Offer 失败:%w", err)
}
// 设置本地描述
if err := pc.SetLocalDescription(offer); err != nil {
pc.Close()
return nil, fmt.Errorf("设置本地描述失败:%w", err)
}
// 设置 ICE 候选回调
c.setupICECandidateHandler(pc, peerID)
return &SignalMessage{
Type: "offer",
SDP: offer.SDP,
}, nil
}
// HandleAnswer 处理 Answer(主动发起方收到应答)
func (c *ICEClient) HandleAnswer(peerID string, answer SignalMessage) error {
c.mu.RLock()
pc, exists := c.peerConns[peerID]
c.mu.RUnlock()
if !exists {
return fmt.Errorf("未找到 PeerConnection%s", peerID)
}
// 设置远程描述
if err := pc.SetRemoteDescription(webrtc.SessionDescription{
Type: webrtc.SDPTypeAnswer,
SDP: answer.SDP,
}); err != nil {
return fmt.Errorf("设置远程描述失败:%w", err)
}
c.logger.Info("已设置 Answer", zap.String("peer_id", peerID))
return nil
}
// HandleOffer 处理 Offer(被动接收方)
func (c *ICEClient) HandleOffer(ctx context.Context, peerID string, offer SignalMessage) (*SignalMessage, error) {
pc, err := c.createPeerConnection(peerID)
if err != nil {
return nil, err
}
// 设置远程描述
if err := pc.SetRemoteDescription(webrtc.SessionDescription{
Type: webrtc.SDPTypeOffer,
SDP: offer.SDP,
}); err != nil {
pc.Close()
return nil, fmt.Errorf("设置远程描述失败:%w", err)
}
// 监听 DataChannel
pc.OnDataChannel(func(dc *webrtc.DataChannel) {
c.logger.Info("收到 DataChannel", zap.String("peer_id", peerID), zap.String("label", dc.Label()))
// 存储 DataChannel
c.mu.Lock()
c.dataChannels[peerID] = dc
c.mu.Unlock()
dc.OnOpen(func() {
c.logger.Info("DataChannel 已打开", zap.String("peer_id", peerID))
})
})
// 创建 Answer
answer, err := pc.CreateAnswer(nil)
if err != nil {
pc.Close()
return nil, fmt.Errorf("创建 Answer 失败:%w", err)
}
// 设置本地描述
if err := pc.SetLocalDescription(answer); err != nil {
pc.Close()
return nil, fmt.Errorf("设置本地描述失败:%w", err)
}
// 设置 ICE 候选回调
c.setupICECandidateHandler(pc, peerID)
return &SignalMessage{
Type: "answer",
SDP: answer.SDP,
}, nil
}
// HandleICECandidate 处理 ICE 候选
func (c *ICEClient) HandleICECandidate(peerID string, candidate SignalMessage) error {
c.mu.RLock()
pc, exists := c.peerConns[peerID]
c.mu.RUnlock()
if !exists {
return fmt.Errorf("未找到 PeerConnection%s", peerID)
}
// 添加 ICE 候选
if err := pc.AddICECandidate(webrtc.ICECandidateInit{
Candidate: candidate.Candidate,
}); err != nil {
return fmt.Errorf("添加 ICE 候选失败:%w", err)
}
c.logger.Debug("已添加 ICE 候选", zap.String("peer_id", peerID))
return nil
}
// setupICECandidateHandler 设置 ICE 候选处理器
func (c *ICEClient) setupICECandidateHandler(pc *webrtc.PeerConnection, peerID string) {
pc.OnICECandidate(func(candidate *webrtc.ICECandidate) {
if candidate == nil {
return
}
c.logger.Debug("发现 ICE 候选",
zap.String("peer_id", peerID),
zap.String("candidate", candidate.String()))
// 触发信令回调
if c.onSignal != nil {
c.onSignal(peerID, SignalMessage{
Type: "candidate",
Candidate: candidate.ToJSON().Candidate,
})
}
})
}
// WaitForConnection 等待连接建立
func (c *ICEClient) WaitForConnection(ctx context.Context, peerID string, timeout time.Duration) error {
c.mu.RLock()
pc, exists := c.peerConns[peerID]
c.mu.RUnlock()
if !exists {
return fmt.Errorf("未找到 PeerConnection%s", peerID)
}
// 创建超时上下文
ctx, cancel := context.WithTimeout(ctx, timeout)
defer cancel()
// 创建连接状态通道
stateCh := make(chan webrtc.PeerConnectionState, 1)
// 监听连接状态
pc.OnConnectionStateChange(func(state webrtc.PeerConnectionState) {
c.logger.Info("连接状态变化",
zap.String("peer_id", peerID),
zap.String("state", state.String()))
select {
case stateCh <- state:
default:
}
})
// 检查当前状态
if pc.ConnectionState() == webrtc.PeerConnectionStateConnected {
return nil
}
// 等待连接建立
for {
select {
case state := <-stateCh:
switch state {
case webrtc.PeerConnectionStateConnected:
return nil
case webrtc.PeerConnectionStateFailed, webrtc.PeerConnectionStateDisconnected:
return fmt.Errorf("连接失败:%s", state.String())
}
case <-ctx.Done():
return fmt.Errorf("等待连接超时")
}
}
}
// GetDataChannel 获取 DataChannel
func (c *ICEClient) GetDataChannel(peerID string) (*webrtc.DataChannel, bool) {
c.mu.RLock()
defer c.mu.RUnlock()
dc, ok := c.dataChannels[peerID]
return dc, ok
}
// ClosePeer 关闭指定 Peer 的连接
func (c *ICEClient) ClosePeer(peerID string) error {
c.mu.Lock()
defer c.mu.Unlock()
if pc, ok := c.peerConns[peerID]; ok {
delete(c.peerConns, peerID)
if dc, ok := c.dataChannels[peerID]; ok {
dc.Close()
delete(c.dataChannels, peerID)
}
return pc.Close()
}
return nil
}
// SetOnSignal 设置信令回调
func (c *ICEClient) SetOnSignal(callback func(peerID string, signal SignalMessage)) {
c.onSignal = callback
}
// Close 关闭所有连接
func (c *ICEClient) Close() error {
c.mu.Lock()
defer c.mu.Unlock()
var errs []error
for peerID, pc := range c.peerConns {
if dc, ok := c.dataChannels[peerID]; ok {
dc.Close()
}
if err := pc.Close(); err != nil {
errs = append(errs, fmt.Errorf("关闭 %s 失败:%w", peerID, err))
}
}
c.peerConns = make(map[string]*webrtc.PeerConnection)
c.dataChannels = make(map[string]*webrtc.DataChannel)
if len(errs) > 0 {
return fmt.Errorf("关闭连接时发生错误:%v", errs)
}
return nil
}
// WebRTCFactory WebRTC 工厂
type WebRTCFactory struct {
client *ICEClient
logger *zap.Logger
config *ICEConfig
}
// NewWebRTCFactory 创建 WebRTC 工厂
func NewWebRTCFactory(config *ICEConfig, logger *zap.Logger) *WebRTCFactory {
return &WebRTCFactory{
client: NewICEClient(config, logger),
logger: logger,
config: config,
}
}
// Layer 返回传输层类型
func (f *WebRTCFactory) Layer() Layer {
return LayerWebRTC
}
// Name 返回名称
func (f *WebRTCFactory) Name() string {
return "WebRTC"
}
// Dial 建立 WebRTC 连接
// 注意:WebRTC 需要信令服务器交换 SDP,这里提供简化的直连模式
// 实际使用时需要通过信令服务器交换 Offer/Answer
func (f *WebRTCFactory) Dial(ctx context.Context, config *DialConfig) (net.Conn, error) {
f.logger.Info("开始建立 WebRTC 连接",
zap.String("peer_id", config.PeerID))
// WebRTC 需要信令服务器支持
// 这里返回错误,提示需要使用信令服务器
return nil, fmt.Errorf("WebRTC 需要信令服务器交换 SDP,请使用 ICEClient 配合信令服务")
}
// GetClient 获取 ICE 客户端
func (f *WebRTCFactory) GetClient() *ICEClient {
return f.client
}
// DataChannelConn DataChannel net.Conn 包装器
type DataChannelConn struct {
dc *webrtc.DataChannel
localAddr net.Addr
remoteAddr net.Addr
readCh chan []byte
readBuf []byte
mu sync.Mutex
closed bool
onClose func()
}
// NewDataChannelConn 创建 DataChannel 连接
func NewDataChannelConn(dc *webrtc.DataChannel, onClose func()) *DataChannelConn {
conn := &DataChannelConn{
dc: dc,
readCh: make(chan []byte, 100),
onClose: onClose,
}
// 设置消息处理
dc.OnMessage(func(msg webrtc.DataChannelMessage) {
conn.mu.Lock()
if conn.closed {
conn.mu.Unlock()
return
}
select {
case conn.readCh <- msg.Data:
default:
// 缓冲区满,丢弃消息
}
conn.mu.Unlock()
})
// 设置关闭处理
dc.OnClose(func() {
conn.Close()
})
return conn
}
// Read 从 DataChannel 读取数据
func (c *DataChannelConn) Read(b []byte) (n int, err error) {
c.mu.Lock()
if c.closed {
c.mu.Unlock()
return 0, io.EOF
}
// 如果有缓冲数据,先返回
if len(c.readBuf) > 0 {
n = copy(b, c.readBuf)
c.readBuf = c.readBuf[n:]
c.mu.Unlock()
return n, nil
}
c.mu.Unlock()
// 等待新数据
select {
case data := <-c.readCh:
c.mu.Lock()
if c.closed {
c.mu.Unlock()
return 0, io.EOF
}
n = copy(b, data)
if n < len(data) {
// 缓冲剩余数据
c.readBuf = data[n:]
}
c.mu.Unlock()
return n, nil
case <-time.After(30 * time.Second):
return 0, fmt.Errorf("读取超时")
}
}
// Write 写入 DataChannel
func (c *DataChannelConn) Write(b []byte) (n int, err error) {
c.mu.Lock()
defer c.mu.Unlock()
if c.closed {
return 0, io.EOF
}
if err := c.dc.Send(b); err != nil {
return 0, fmt.Errorf("发送失败:%w", err)
}
return len(b), nil
}
// Close 关闭连接
func (c *DataChannelConn) Close() error {
c.mu.Lock()
defer c.mu.Unlock()
if c.closed {
return nil
}
c.closed = true
if c.onClose != nil {
c.onClose()
}
return c.dc.Close()
}
// LocalAddr 返回本地地址
func (c *DataChannelConn) LocalAddr() net.Addr {
if c.localAddr == nil {
return &net.TCPAddr{IP: net.IPv4zero, Port: 0}
}
return c.localAddr
}
// RemoteAddr 返回远程地址
func (c *DataChannelConn) RemoteAddr() net.Addr {
if c.remoteAddr == nil {
return &net.TCPAddr{IP: net.IPv4zero, Port: 0}
}
return c.remoteAddr
}
// SetDeadline 设置截止时间
func (c *DataChannelConn) SetDeadline(t time.Time) error {
return nil
}
// SetReadDeadline 设置读取截止时间
func (c *DataChannelConn) SetReadDeadline(t time.Time) error {
return nil
}
// SetWriteDeadline 设置写入截止时间
func (c *DataChannelConn) SetWriteDeadline(t time.Time) error {
return nil
}
// MarshalJSON 序列化信令消息
func (m SignalMessage) MarshalJSON() ([]byte, error) {
type Alias SignalMessage
return json.Marshal((*Alias)(&m))
}
// UnmarshalJSON 反序列化信令消息
func (m *SignalMessage) UnmarshalJSON(data []byte) error {
type Alias SignalMessage
var tmp Alias
if err := json.Unmarshal(data, &tmp); err != nil {
return err
}
*m = SignalMessage(tmp)
return nil
}
+181
View File
@@ -0,0 +1,181 @@
package connect
import (
"context"
"fmt"
"net"
"sync"
"time"
"go.uber.org/zap"
)
// RealTCPConn 真正的 TCP 连接(用于传输 WireGuard 密文)
// 与 FakeTCP 不同,RealTCP 不封装 UDP 包,直接传输原始数据
type RealTCPConn struct {
conn net.Conn
closed bool
mu sync.Mutex
}
// NewRealTCPConn 创建 RealTCP 连接
func NewRealTCPConn(conn net.Conn) *RealTCPConn {
return &RealTCPConn{
conn: conn,
}
}
// Read 读取数据
func (c *RealTCPConn) Read(b []byte) (int, error) {
c.mu.Lock()
defer c.mu.Unlock()
if c.closed {
return 0, fmt.Errorf("connection closed")
}
return c.conn.Read(b)
}
// Write 写入数据
func (c *RealTCPConn) Write(b []byte) (int, error) {
c.mu.Lock()
defer c.mu.Unlock()
if c.closed {
return 0, fmt.Errorf("connection closed")
}
return c.conn.Write(b)
}
// Close 关闭连接
func (c *RealTCPConn) Close() error {
c.mu.Lock()
defer c.mu.Unlock()
c.closed = true
return c.conn.Close()
}
// RealTCPFactory RealTCP 传输工厂
type RealTCPFactory struct {
logger *zap.Logger
}
// NewRealTCPFactory 创建 RealTCP 工厂
func NewRealTCPFactory(logger *zap.Logger) *RealTCPFactory {
return &RealTCPFactory{
logger: logger,
}
}
// Layer 返回传输层类型
func (f *RealTCPFactory) Layer() Layer {
return LayerRealTCP
}
// Name 返回名称
func (f *RealTCPFactory) Name() string {
return "RealTCP"
}
// Dial 建立 RealTCP 连接
func (f *RealTCPFactory) Dial(ctx context.Context, config *DialConfig) (net.Conn, error) {
f.logger.Info("开始建立 RealTCP 连接",
zap.String("peer_id", config.PeerID))
// 1. 解析对端地址(PeerID 格式应为 "ip:port"
if config.PeerID == "" {
return nil, fmt.Errorf("PeerID 为空")
}
// 2. 建立 TCP 连接
dialer := &net.Dialer{Timeout: config.Timeout}
conn, err := dialer.DialContext(ctx, "tcp", config.PeerID)
if err != nil {
return nil, fmt.Errorf("TCP 连接失败:%w", err)
}
// 3. 包装为 RealTCP 连接(直接传输原始数据)
realConn := NewRealTCPConn(conn)
f.logger.Info("RealTCP 连接建立成功",
zap.String("peer_id", config.PeerID),
zap.String("local_addr", conn.LocalAddr().String()),
zap.String("remote_addr", conn.RemoteAddr().String()))
return realConn, nil
}
// LocalAddr 本地地址
func (c *RealTCPConn) LocalAddr() net.Addr {
return c.conn.LocalAddr()
}
// RemoteAddr 远程地址
func (c *RealTCPConn) RemoteAddr() net.Addr {
return c.conn.RemoteAddr()
}
// SetDeadline 设置截止时间
func (c *RealTCPConn) SetDeadline(t time.Time) error {
return c.conn.SetDeadline(t)
}
// SetReadDeadline 设置读截止时间
func (c *RealTCPConn) SetReadDeadline(t time.Time) error {
return c.conn.SetReadDeadline(t)
}
// SetWriteDeadline 设置写截止时间
func (c *RealTCPConn) SetWriteDeadline(t time.Time) error {
return c.conn.SetWriteDeadline(t)
}
// DialRealTCP 拨号 RealTCP 连接
func DialRealTCP(ctx context.Context, network, addr string, logger *zap.Logger) (net.Conn, error) {
logger.Debug("dialing RealTCP", zap.String("addr", addr))
// 建立 TCP 连接
conn, err := (&net.Dialer{}).DialContext(ctx, network, addr)
if err != nil {
return nil, fmt.Errorf("failed to dial TCP: %w", err)
}
// 包装为 RealTCP 连接
return NewRealTCPConn(conn), nil
}
// ListenRealTCP 监听 RealTCP 端口
func ListenRealTCP(network, addr string, logger *zap.Logger) (net.Listener, error) {
logger.Info("listening RealTCP", zap.String("addr", addr))
// 监听 TCP 端口
listener, err := net.Listen(network, addr)
if err != nil {
return nil, fmt.Errorf("failed to listen TCP: %w", err)
}
return &realTCPListener{
Listener: listener,
logger: logger,
}, nil
}
// realTCPListener RealTCP 监听器
type realTCPListener struct {
net.Listener
logger *zap.Logger
}
// Accept 接受连接并包装为 RealTCPConn
func (l *realTCPListener) Accept() (net.Conn, error) {
conn, err := l.Listener.Accept()
if err != nil {
return nil, err
}
l.logger.Debug("accepted RealTCP connection", zap.String("addr", conn.RemoteAddr().String()))
return NewRealTCPConn(conn), nil
}
+754
View File
@@ -0,0 +1,754 @@
package connect
import (
"context"
"fmt"
"net"
"sync"
"time"
"go.uber.org/zap"
)
// Layer 传输层类型(9 层策略)
type Layer int
const (
// LayerDirectUDP Direct-UDP 直连(WireGuard over UDP- 最高效
LayerDirectUDP Layer = iota
// LayerFakeTCP Direct-FakeTCPUDP 封装 TCP 头部,欺骗防火墙)
LayerFakeTCP
// LayerRealTCP Direct-RealTCPP2P TCP 直连)
LayerRealTCP
// LayerTURNUDP TURN-UDP 中继(标准 RFC 5766
LayerTURNUDP
// LayerTURNQUIC TURN-QUIC 中继(私有扩展,RFC 9000
LayerTURNQUIC
// LayerTURNTCP TURN-TCP 中继(TCP 中继)
LayerTURNTCP
// LayerTURNTLS TURN-TLS 中继(TLS 加密,RFC 8656
LayerTURNTLS
// LayerWebRTC WebRTC DataChannelDTLS 加密)
LayerWebRTC
// LayerWS WS/WSS 兜底(仅 80/443 端口,终极兜底)
LayerWS
// LayerCount 传输层总数
LayerCount
)
// String 实现 Stringer 接口
func (l Layer) String() string {
switch l {
case LayerDirectUDP:
return "Direct-UDP"
case LayerFakeTCP:
return "Direct-FakeTCP"
case LayerRealTCP:
return "Direct-RealTCP"
case LayerTURNUDP:
return "TURN-UDP"
case LayerTURNQUIC:
return "TURN-QUIC"
case LayerTURNTCP:
return "TURN-TCP"
case LayerTURNTLS:
return "TURN-TLS"
case LayerWebRTC:
return "WebRTC"
case LayerWS:
return "WS/WSS"
default:
return "Unknown"
}
}
// DefaultLayerOrder 默认优先级顺序(从最优到兜底)
// 根据 MeshRay_项目文档 v2.0.1 第 132-153 行定义
var DefaultLayerOrder = []Layer{
LayerDirectUDP, // 1. Direct-UDP - 公网/锥型 NAT,首选链路
LayerFakeTCP, // 2. Direct-FakeTCP - 校园网、酒店 Wi-Fi、UDP 被 QoS 限速
LayerRealTCP, // 3. Direct-RealTCP - 完全禁用 UDP,仅允许 TCP 出站
LayerTURNUDP, // 4. TURN-UDP 中继 - 无 P2P 直连,但 UDP 可通
LayerTURNQUIC, // 5. TURN-QUIC 中继 - UDP 可通但弱网(4G/5G、高丢包)【私有扩展】
LayerTURNTCP, // 6. TURN-TCP 中继 - UDP 封禁,仅放行 TCP
LayerTURNTLS, // 7. TURN-TLS 中继 - 企业防火墙 DPI,仅放行 HTTPS
LayerWebRTC, // 8. WebRTC 终极兜底 - 最严格隔离内网、代理环境
LayerWS, // 9. WS/WSS 兜底 - 仅放行 80/443 端口,且封锁 TURN
}
// TransportFactory 传输工厂接口 - 每种传输方式必须实现
type TransportFactory interface {
// Layer 返回传输层类型
Layer() Layer
// Dial 建立连接到对端
// 返回标准的 net.Conn 接口
Dial(ctx context.Context, config *DialConfig) (net.Conn, error)
// Name 返回传输方式名称(用于日志)
Name() string
}
// DialConfig 拨号配置
type DialConfig struct {
// PeerID 对端标识
PeerID string
// PeerPublicKey 对端公钥
PeerPublicKey string
// STUNServers STUN 服务器列表(用于 P2P)
STUNServers []string
// TURNServers TURN 服务器列表
TURNServers []string
// WSServers WebSocket 服务器列表
WSServers []string
// SignalingServers WebRTC 第三方信令服务器列表
SignalingServers []string
// ICESServers ICE 服务器列表(STUN+TURN 的组合)
ICESServers []string
// Timeout 连接超时
Timeout time.Duration
// Logger 日志记录器
Logger *zap.Logger
}
// StrategyScheduler 9 层策略调度器(主动调度层)
// 职责:
// 1. 按优先级选择链路(P2P → Mesh中继 → TURN-UDP → ... → WS/WSS
// 2. 根据网络环境自动切换(500ms 超时 / 10s 丢包率 > 10%
// 3. 切换后探测恢复并自动切回高性能链路(30s)
type StrategyScheduler struct {
layerFactories map[Layer]TransportFactory // 各层的工厂
layerOrder []Layer // 优先级顺序
logger *zap.Logger
// 每个 Peer 的降级控制器
fallbackControllers map[string]*FallbackController // peerID -> controller
fallbackMu sync.RWMutex
// 当前活跃连接
activeConnections map[string]activeConn // peerID -> 连接信息
connMu sync.RWMutex
// 统计
stats *SchedulerStats
// 连接变更回调(通知上层 ConnManager
OnConnectionUpdate func(peerID string, conn net.Conn, err error)
}
// activeConn 活跃连接信息
type activeConn struct {
conn net.Conn
layer Layer
peerID string
established time.Time
config *DialConfig
}
// SchedulerStats 调度器统计
type SchedulerStats struct {
mu sync.RWMutex
totalDials int64
successDials int64
fallbackCount int64
recoveryCount int64
layerDialCount map[Layer]int64
layerFailCount map[Layer]int64
}
// NewStrategyScheduler 创建策略调度器
func NewStrategyScheduler(logger *zap.Logger) *StrategyScheduler {
return &StrategyScheduler{
layerFactories: make(map[Layer]TransportFactory),
layerOrder: DefaultLayerOrder,
logger: logger,
fallbackControllers: make(map[string]*FallbackController),
activeConnections: make(map[string]activeConn),
stats: &SchedulerStats{
layerDialCount: make(map[Layer]int64),
layerFailCount: make(map[Layer]int64),
},
}
}
// RegisterFactory 注册传输工厂
func (s *StrategyScheduler) RegisterFactory(factory TransportFactory) {
layer := factory.Layer()
s.layerFactories[layer] = factory
s.logger.Debug("注册传输工厂",
zap.String("layer", layer.String()),
zap.String("name", factory.Name()))
}
// SetLayerOrder 设置优先级顺序
func (s *StrategyScheduler) SetLayerOrder(order []Layer) {
if len(order) == 0 {
s.logger.Warn("空的层级顺序,使用默认顺序")
return
}
s.layerOrder = order
s.logger.Info("更新传输层优先级顺序", zap.Any("order", order))
}
// Dial 按优先级顺序尝试建立连接
// 这是核心方法,实现了 9 层策略调度
func (s *StrategyScheduler) Dial(config *DialConfig) (net.Conn, error) {
ctx := context.Background()
if config.Timeout > 0 {
var cancel context.CancelFunc
ctx, cancel = context.WithTimeout(ctx, config.Timeout)
defer cancel()
}
s.logger.Info("开始 8 层策略调度连接",
zap.String("peer_id", config.PeerID),
zap.Int("total_layers", len(s.layerOrder)))
// 统计
s.stats.mu.Lock()
s.stats.totalDials++
s.stats.mu.Unlock()
var lastErr error
for i, layer := range s.layerOrder {
factory, ok := s.layerFactories[layer]
if !ok {
s.logger.Debug("该传输层未注册,跳过",
zap.String("layer", layer.String()))
continue
}
s.logger.Debug("尝试第 N 层传输",
zap.Int("index", i),
zap.String("layer", layer.String()),
zap.String("name", factory.Name()))
// 统计该层拨号次数
s.stats.mu.Lock()
s.stats.layerDialCount[layer]++
s.stats.mu.Unlock()
startTime := time.Now()
conn, err := factory.Dial(ctx, config)
duration := time.Since(startTime)
if err == nil {
// 成功!
s.stats.mu.Lock()
s.stats.successDials++
s.stats.mu.Unlock()
// 记录活跃连接
s.connMu.Lock()
s.activeConnections[config.PeerID] = activeConn{
conn: conn,
layer: layer,
peerID: config.PeerID,
established: time.Now(),
config: config,
}
s.connMu.Unlock()
// 创建或更新降级控制器
s.ensureFallbackController(config.PeerID, layer)
s.logger.Info("连接建立成功",
zap.String("layer", layer.String()),
zap.String("name", factory.Name()),
zap.String("peer_id", config.PeerID),
zap.String("remote_addr", conn.RemoteAddr().String()),
zap.Duration("duration", duration))
// 包装连接,用于监控
return newMonitoredConn(conn, config.PeerID, layer, s), nil
}
// 失败,统计
s.stats.mu.Lock()
s.stats.layerFailCount[layer]++
s.stats.mu.Unlock()
// 记录失败并继续尝试下一层
lastErr = err
s.logger.Warn("该传输层连接失败,尝试下一层",
zap.String("layer", layer.String()),
zap.Duration("duration", duration),
zap.Error(err))
}
// 所有层都失败
return nil, fmt.Errorf("所有传输层均失败,最后错误:%w", lastErr)
}
// reconnectToLayer 触发重连到指定层级
func (s *StrategyScheduler) reconnectToLayer(peerID string, toLayer Layer) {
s.connMu.RLock()
ac, exists := s.activeConnections[peerID]
s.connMu.RUnlock()
if !exists || ac.config == nil {
s.logger.Warn("重连失败:找不到活跃连接配置", zap.String("peer_id", peerID))
return
}
factory, ok := s.layerFactories[toLayer]
if !ok {
s.logger.Error("重连失败:找不到目标层级工厂", zap.String("layer", toLayer.String()))
return
}
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
defer cancel()
conn, err := factory.Dial(ctx, ac.config)
if err != nil {
s.logger.Error("降级重连失败", zap.Error(err))
if s.OnConnectionUpdate != nil {
s.OnConnectionUpdate(peerID, nil, err)
}
return
}
wrappedConn := newMonitoredConn(conn, peerID, toLayer, s)
s.connMu.Lock()
if oldAc, exists := s.activeConnections[peerID]; exists {
oldAc.conn.Close()
}
s.activeConnections[peerID] = activeConn{
conn: wrappedConn,
layer: toLayer,
peerID: peerID,
established: time.Now(),
config: ac.config,
}
s.connMu.Unlock()
if s.OnConnectionUpdate != nil {
s.OnConnectionUpdate(peerID, wrappedConn, nil)
}
}
// ensureFallbackController 确保对端有降级控制器
func (s *StrategyScheduler) ensureFallbackController(peerID string, initialLayer Layer) {
s.fallbackMu.Lock()
defer s.fallbackMu.Unlock()
if _, exists := s.fallbackControllers[peerID]; !exists {
controller := NewFallbackController(
initialLayer,
func(from, to Layer) {
// 降级回调
s.stats.mu.Lock()
s.stats.fallbackCount++
s.stats.mu.Unlock()
s.logger.Warn("链路降级",
zap.String("peer_id", peerID),
zap.String("from_layer", from.String()),
zap.String("to_layer", to.String()))
// 触发重连到新层级
go s.reconnectToLayer(peerID, to)
},
func(to Layer) {
// 恢复回调
s.stats.mu.Lock()
s.stats.recoveryCount++
s.stats.mu.Unlock()
s.logger.Info("链路恢复",
zap.String("peer_id", peerID),
zap.String("to_layer", to.String()))
// 回调处理已经在 probeHighLayers 中完成并传递了新连接
},
s.logger,
s, // pass scheduler to access activeConnections
)
s.fallbackControllers[peerID] = controller
}
}
// RecordLatency 记录延迟(供 MonitoredConn 调用)
func (s *StrategyScheduler) RecordLatency(peerID string, success bool, duration time.Duration) {
s.fallbackMu.RLock()
controller, exists := s.fallbackControllers[peerID]
s.fallbackMu.RUnlock()
if exists {
controller.CheckAndFallback(success, duration)
}
}
// GetActiveLayer 获取当前活跃的传输层(用于监控)
func (s *StrategyScheduler) GetActiveLayer() Layer {
// 返回第一个活跃连接的层级
s.connMu.RLock()
defer s.connMu.RUnlock()
for _, ac := range s.activeConnections {
return ac.layer
}
return LayerDirectUDP // 默认值
}
// GetAllActiveLayers 获取所有 Peer 的活跃层级(用于全局监控)
func (s *StrategyScheduler) GetAllActiveLayers() map[string]Layer {
s.connMu.RLock()
defer s.connMu.RUnlock()
result := make(map[string]Layer)
for peerID, ac := range s.activeConnections {
result[peerID] = ac.layer
}
return result
}
// GetPeerLayer 获取指定 Peer 的当前层级
func (s *StrategyScheduler) GetPeerLayer(peerID string) Layer {
s.connMu.RLock()
defer s.connMu.RUnlock()
if ac, exists := s.activeConnections[peerID]; exists {
return ac.layer
}
return LayerDirectUDP // 默认值
}
// GetStats 获取统计信息
func (s *StrategyScheduler) GetStats() map[string]interface{} {
s.stats.mu.RLock()
defer s.stats.mu.RUnlock()
layerStats := make(map[string]int64)
for layer, count := range s.stats.layerDialCount {
layerStats[layer.String()+"_dial"] = count
}
for layer, count := range s.stats.layerFailCount {
layerStats[layer.String()+"_fail"] = count
}
return map[string]interface{}{
"total_dials": s.stats.totalDials,
"success_dials": s.stats.successDials,
"fallback_count": s.stats.fallbackCount,
"recovery_count": s.stats.recoveryCount,
"layer_stats": layerStats,
}
}
// ClosePeer 关闭指定 Peer 的连接和控制器
func (s *StrategyScheduler) ClosePeer(peerID string) {
// 关闭连接
s.connMu.Lock()
if ac, exists := s.activeConnections[peerID]; exists {
ac.conn.Close()
delete(s.activeConnections, peerID)
}
s.connMu.Unlock()
// 移除降级控制器
s.fallbackMu.Lock()
if controller, exists := s.fallbackControllers[peerID]; exists {
// 停止恢复探测器
if controller.recoveryTimer != nil {
controller.recoveryTimer.Stop()
}
delete(s.fallbackControllers, peerID)
}
s.fallbackMu.Unlock()
s.logger.Debug("已关闭 Peer 连接和控制器",
zap.String("peer_id", peerID))
}
// MonitoredConn 带监控的连接包装器
type MonitoredConn struct {
net.Conn
peerID string
layer Layer
scheduler *StrategyScheduler
}
// newMonitoredConn 创建带监控的连接
func newMonitoredConn(conn net.Conn, peerID string, layer Layer, scheduler *StrategyScheduler) *MonitoredConn {
return &MonitoredConn{
Conn: conn,
peerID: peerID,
layer: layer,
scheduler: scheduler,
}
}
// Read 重写 Read 方法,记录延迟
func (c *MonitoredConn) Read(b []byte) (n int, err error) {
start := time.Now()
n, err = c.Conn.Read(b)
duration := time.Since(start)
// 记录成功/失败
c.scheduler.RecordLatency(c.peerID, err == nil, duration)
return n, err
}
// Write 重写 Write 方法,记录延迟
func (c *MonitoredConn) Write(b []byte) (n int, err error) {
start := time.Now()
n, err = c.Conn.Write(b)
duration := time.Since(start)
// 记录成功/失败
c.scheduler.RecordLatency(c.peerID, err == nil, duration)
return n, err
}
// FallbackController 降级控制器
type FallbackController struct {
currentLayer Layer // 当前使用的层
windowStart time.Time // 滑动窗口起始时间
packetCount int // 总包数
lostPacketCount int // 丢包数
mu chan struct{} // 互斥锁(用 channel 实现)
triggerFallback func(Layer, Layer) // 降级触发回调
triggerRecovery func(Layer) // 恢复触发回调
logger *zap.Logger
recoveryTimer *time.Timer // 恢复探测定时器
scheduler *StrategyScheduler
}
const (
// TimeoutThreshold 单次超时阈值
TimeoutThreshold = 500 * time.Millisecond
// PacketLossThreshold 丢包率阈值
PacketLossThreshold = 0.10 // 10%
// RecoveryInterval 恢复探测间隔
RecoveryInterval = 30 * time.Second
// SlidingWindowDuration 滑动窗口时长
SlidingWindowDuration = 10 * time.Second
)
// NewFallbackController 创建降级控制器
func NewFallbackController(
initialLayer Layer,
onFallback func(Layer, Layer),
onRecovery func(Layer),
logger *zap.Logger,
scheduler *StrategyScheduler,
) *FallbackController {
fc := &FallbackController{
currentLayer: initialLayer,
mu: make(chan struct{}, 1),
triggerFallback: onFallback,
triggerRecovery: onRecovery,
logger: logger,
scheduler: scheduler,
}
// 启动恢复探测
fc.startRecoveryProbe()
return fc
}
// CheckAndFallback 检查是否需要降级
// 在每次连接操作后调用
func (fc *FallbackController) CheckAndFallback(success bool, duration time.Duration) {
select {
case fc.mu <- struct{}{}:
defer func() { <-fc.mu }()
default:
// 锁被占用,说明正在处理,直接返回
return
}
// 重置滑动窗口
if time.Since(fc.windowStart) > SlidingWindowDuration {
fc.windowStart = time.Now()
fc.packetCount = 0
fc.lostPacketCount = 0
}
// 统计
fc.packetCount++
if !success || duration > TimeoutThreshold {
fc.lostPacketCount++
}
// 检查是否达到阈值
if fc.packetCount >= 10 {
lossRate := float64(fc.lostPacketCount) / float64(fc.packetCount)
if lossRate > PacketLossThreshold {
fc.triggerFallbackLocked()
}
}
}
// triggerFallbackLocked 执行降级(已持有锁)
func (fc *FallbackController) triggerFallbackLocked() {
currentIndex := int(fc.currentLayer)
if currentIndex >= int(LayerCount)-1 {
// 已经是最低优先级,无法降级
fc.logger.Warn("已是最底层级,无法降级",
zap.String("current_layer", fc.currentLayer.String()))
return
}
nextLayer := Layer(currentIndex + 1)
// 在更新 currentLayer 之前保存旧值用于回调
oldLayer := fc.currentLayer
fc.logger.Warn("触发降级",
zap.String("from_layer", oldLayer.String()),
zap.String("to_layer", nextLayer.String()))
fc.currentLayer = nextLayer
fc.resetWindow()
if fc.triggerFallback != nil {
fc.triggerFallback(oldLayer, nextLayer)
}
// 重置恢复定时器
fc.startRecoveryProbe()
}
// startRecoveryProbe 启动恢复探测
func (fc *FallbackController) startRecoveryProbe() {
if fc.recoveryTimer != nil {
fc.recoveryTimer.Stop()
}
fc.recoveryTimer = time.AfterFunc(RecoveryInterval, func() {
fc.probeHigherLayers()
})
}
// probeHigherLayers 探测更高层级
func (fc *FallbackController) probeHigherLayers() {
select {
case fc.mu <- struct{}{}:
defer func() { <-fc.mu }()
default:
return
}
currentIndex := int(fc.currentLayer)
if currentIndex == 0 {
// 已经是最高优先级,无需探测
return
}
// 尝试上一层
higherLayer := Layer(currentIndex - 1)
fc.logger.Info("探测更高层级",
zap.String("current_layer", fc.currentLayer.String()),
zap.String("probe_layer", higherLayer.String()))
// 获取 PeerID 及 Config
fc.scheduler.connMu.RLock()
var peerID string
var config *DialConfig
for pid, ac := range fc.scheduler.activeConnections {
if ac.layer == fc.currentLayer {
peerID = pid
config = ac.config
break
}
}
fc.scheduler.connMu.RUnlock()
if config == nil {
fc.logger.Warn("探测更高层级失败:找不到有效 DialConfig")
return
}
factory, ok := fc.scheduler.layerFactories[higherLayer]
if !ok {
fc.logger.Debug("更高层级未注册工厂,跳过探测")
return
}
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
defer cancel()
conn, err := factory.Dial(ctx, config)
if err == nil {
fc.logger.Info("更高层级探测成功,准备切换")
wrappedConn := newMonitoredConn(conn, peerID, higherLayer, fc.scheduler)
fc.scheduler.connMu.Lock()
if ac, exists := fc.scheduler.activeConnections[peerID]; exists {
ac.conn.Close() // Close old
ac.conn = wrappedConn
ac.layer = higherLayer
fc.scheduler.activeConnections[peerID] = ac
}
fc.scheduler.connMu.Unlock()
if fc.scheduler.OnConnectionUpdate != nil {
fc.scheduler.OnConnectionUpdate(peerID, wrappedConn, nil)
}
// 触发恢复回调
fc.triggerRecoveryLocked(higherLayer)
} else {
fc.logger.Debug("更高层级探测失败", zap.Error(err))
}
}
// triggerRecoveryLocked 执行恢复(已持有锁)
func (fc *FallbackController) triggerRecoveryLocked(higherLayer Layer) {
fc.logger.Info("触发恢复",
zap.String("from_layer", fc.currentLayer.String()),
zap.String("to_layer", higherLayer.String()))
fc.currentLayer = higherLayer
fc.resetWindow()
if fc.triggerRecovery != nil {
fc.triggerRecovery(higherLayer)
}
}
// resetWindow 重置滑动窗口
func (fc *FallbackController) resetWindow() {
fc.windowStart = time.Now()
fc.packetCount = 0
fc.lostPacketCount = 0
}
// GetCurrentLayer 获取当前层级
func (fc *FallbackController) GetCurrentLayer() Layer {
select {
case fc.mu <- struct{}{}:
defer func() { <-fc.mu }()
default:
return fc.currentLayer
}
return fc.currentLayer
}
+119
View File
@@ -0,0 +1,119 @@
package connect
import (
"fmt"
"net"
"time"
"github.com/pion/stun"
"go.uber.org/zap"
)
// STUNClient STUN 客户端 - 用于 NAT 探测和候选地址采集
type STUNClient struct {
servers []string
logger *zap.Logger
timeout time.Duration
}
// NewSTUNClient 创建 STUN 客户端
func NewSTUNClient(servers []string, logger *zap.Logger) *STUNClient {
return &STUNClient{
servers: servers,
logger: logger,
timeout: 5 * time.Second,
}
}
// DiscoverAddress 发现外部地址(通过单个 STUN 服务器)
func (c *STUNClient) DiscoverAddress(server string) (*net.UDPAddr, error) {
host, port, err := net.SplitHostPort(server)
if err != nil {
return nil, fmt.Errorf("STUN 服务器地址格式错误:%w", err)
}
udpAddr, err := net.ResolveUDPAddr("udp4", net.JoinHostPort(host, port))
if err != nil {
return nil, fmt.Errorf("解析 UDP 地址失败:%w", err)
}
conn, err := net.DialUDP("udp4", nil, udpAddr)
if err != nil {
return nil, fmt.Errorf("连接 STUN 服务器失败:%w", err)
}
defer conn.Close()
conn.SetDeadline(time.Now().Add(c.timeout))
// 构建 STUN Binding Request
msg, err := stun.Build(stun.BindingRequest, stun.TransactionID)
if err != nil {
return nil, fmt.Errorf("构建 STUN 请求失败:%w", err)
}
// 发送请求
if _, err := conn.Write(msg.Raw); err != nil {
return nil, fmt.Errorf("发送 STUN 请求失败:%w", err)
}
// 读取响应
buf := make([]byte, 1024)
n, err := conn.Read(buf)
if err != nil {
return nil, fmt.Errorf("读取 STUN 响应失败:%w", err)
}
// 解析响应
res := &stun.Message{Raw: buf[:n]}
if err := res.Decode(); err != nil {
return nil, fmt.Errorf("解码 STUN 响应失败:%w", err)
}
// 提取 XOR-MAPPED-ADDRESS
var xorAddr stun.XORMappedAddress
if err := xorAddr.GetFrom(res); err != nil {
return nil, fmt.Errorf("提取外部地址失败:%w", err)
}
c.logger.Debug("STUN 查询成功",
zap.String("server", server),
zap.String("external_addr", xorAddr.String()))
return &net.UDPAddr{
IP: xorAddr.IP,
Port: xorAddr.Port,
}, nil
}
// CollectCandidates 收集候选地址(通过多个 STUN 服务器)
func (c *STUNClient) CollectCandidates() []string {
var candidates []string
for _, server := range c.servers {
addr, err := c.DiscoverAddress(server)
if err != nil {
c.logger.Debug("STUN 服务器查询失败",
zap.String("server", server),
zap.Error(err))
continue
}
candidates = append(candidates, addr.String())
c.logger.Debug("收集到候选地址",
zap.String("server", server),
zap.String("candidate", addr.String()))
}
return candidates
}
// GetExternalIP 获取外部 IP(兼容旧 API)
func (c *STUNClient) GetExternalIP() (string, error) {
for _, server := range c.servers {
addr, err := c.DiscoverAddress(server)
if err == nil {
return addr.IP.String(), nil
}
}
return "", fmt.Errorf("所有 STUN 服务器均查询失败")
}
+358
View File
@@ -0,0 +1,358 @@
package connect
import (
"context"
"fmt"
"net"
"sync"
"time"
"github.com/pion/turn/v2"
"go.uber.org/zap"
)
// TURNProtocol TURN 协议类型
type TURNProtocol string
const (
TURNProtocolUDP TURNProtocol = "udp"
TURNProtocolTCP TURNProtocol = "tcp"
TURNProtocolTLS TURNProtocol = "tls"
)
// TURNFactory TURN 工厂(Layer 4-6: TURN-UDP/TCP/TLS
// 自包含实现:TURN 协议协商 + 建连
type TURNFactory struct {
protocol TURNProtocol
servers []string
username string
password string
logger *zap.Logger
}
// NewTURNFactory 创建 TURN 工厂
func NewTURNFactory(protocol TURNProtocol, servers []string, username, password string, logger *zap.Logger) *TURNFactory {
return &TURNFactory{
protocol: protocol,
servers: servers,
username: username,
password: password,
logger: logger,
}
}
// Layer 返回传输层类型
func (f *TURNFactory) Layer() Layer {
switch f.protocol {
case TURNProtocolUDP:
return LayerTURNUDP
case TURNProtocolTCP:
return LayerTURNTCP
case TURNProtocolTLS:
return LayerTURNTLS
default:
return LayerTURNUDP
}
}
// Name 返回名称
func (f *TURNFactory) Name() string {
switch f.protocol {
case TURNProtocolUDP:
return "TURN-UDP"
case TURNProtocolTCP:
return "TURN-TCP"
case TURNProtocolTLS:
return "TURN-TLS"
default:
return "TURN-UDP"
}
}
// Dial 建立 TURN 中继连接
func (f *TURNFactory) Dial(ctx context.Context, config *DialConfig) (net.Conn, error) {
f.logger.Info("开始建立 TURN 中继连接",
zap.String("peer_id", config.PeerID),
zap.String("protocol", string(f.protocol)))
servers := f.servers
if len(servers) == 0 {
servers = config.TURNServers
}
if len(servers) == 0 {
return nil, fmt.Errorf("未配置 TURN 服务器")
}
// 解析第一个 TURN 服务器
server := servers[0]
host, port := parseServerAddr(server)
// 根据协议类型建立连接
var relayConn net.PacketConn
var err error
switch f.protocol {
case TURNProtocolUDP:
relayConn, err = f.allocateUDP(ctx, host, port, config)
case TURNProtocolTCP:
relayConn, err = f.allocateTCP(ctx, host, port, config)
case TURNProtocolTLS:
return nil, fmt.Errorf("TURN-TLS 尚未实现")
default:
return nil, fmt.Errorf("不支持的 TURN 协议:%s", f.protocol)
}
if err != nil {
return nil, fmt.Errorf("TURN 分配失败:%w", err)
}
f.logger.Info("TURN 中继连接建立成功",
zap.String("peer_id", config.PeerID),
zap.String("relay_addr", relayConn.LocalAddr().String()))
// 包装成 net.Conn 返回
return newTURNConn(relayConn, f.logger), nil
}
// allocateUDP UDP TURN 分配
func (f *TURNFactory) allocateUDP(ctx context.Context, host, port string, config *DialConfig) (net.PacketConn, error) {
udpAddr, err := net.ResolveUDPAddr("udp", net.JoinHostPort(host, port))
if err != nil {
return nil, fmt.Errorf("解析 UDP 地址失败:%w", err)
}
conn, err := net.DialUDP("udp", nil, udpAddr)
if err != nil {
return nil, fmt.Errorf("创建 UDP 连接失败:%w", err)
}
clientConfig := &turn.ClientConfig{
STUNServerAddr: net.JoinHostPort(host, port),
TURNServerAddr: net.JoinHostPort(host, port),
Username: f.username,
Password: f.password,
Conn: conn,
}
client, err := turn.NewClient(clientConfig)
if err != nil {
conn.Close()
return nil, fmt.Errorf("创建 TURN 客户端失败:%w", err)
}
if err := client.Listen(); err != nil {
client.Close()
conn.Close()
return nil, fmt.Errorf("TURN 客户端监听失败:%w", err)
}
relayConn, err := client.Allocate()
if err != nil {
client.Close()
conn.Close()
return nil, fmt.Errorf("分配 TURN 中继失败:%w", err)
}
// 创建 Permission(允许特定对端地址使用中继)
// 这是 TURN 协议的关键步骤,否则无法收发数据
// 注意:PeerID 在这里应该是对端的公网地址(由信使服务器转发)
if config != nil && config.PeerID != "" {
peerAddr, err := net.ResolveUDPAddr("udp", config.PeerID)
if err == nil {
if permErr := client.CreatePermission(peerAddr); permErr != nil {
f.logger.Warn("CreatePermission 失败",
zap.String("peer_addr", peerAddr.String()),
zap.Error(permErr))
// 注意:CreatePermission 失败不影响连接建立,只是警告
} else {
f.logger.Debug("CreatePermission 成功",
zap.String("peer_addr", peerAddr.String()))
}
}
}
f.logger.Debug("TURN-UDP 分配成功",
zap.String("relay_addr", relayConn.LocalAddr().String()))
return relayConn, nil
}
// allocateTCP TCP TURN 分配
func (f *TURNFactory) allocateTCP(ctx context.Context, host, port string, config *DialConfig) (net.PacketConn, error) {
dialer := &net.Dialer{Timeout: 10 * time.Second}
conn, err := dialer.DialContext(ctx, "tcp", net.JoinHostPort(host, port))
if err != nil {
return nil, fmt.Errorf("TCP 连接失败:%w", err)
}
packetConn := newTCPPacketConn(conn, f.logger)
clientConfig := &turn.ClientConfig{
STUNServerAddr: net.JoinHostPort(host, port),
TURNServerAddr: net.JoinHostPort(host, port),
Username: f.username,
Password: f.password,
Conn: packetConn,
}
client, err := turn.NewClient(clientConfig)
if err != nil {
conn.Close()
return nil, fmt.Errorf("创建 TURN 客户端失败:%w", err)
}
if err := client.Listen(); err != nil {
client.Close()
conn.Close()
return nil, fmt.Errorf("TURN 客户端监听失败:%w", err)
}
relayConn, err := client.Allocate()
if err != nil {
client.Close()
conn.Close()
return nil, fmt.Errorf("分配 TURN 中继失败:%w", err)
}
// 创建 Permission(允许特定对端地址使用中继)
if config != nil && config.PeerID != "" {
peerAddr, err := net.ResolveTCPAddr("tcp", config.PeerID)
if err == nil {
if permErr := client.CreatePermission(peerAddr); permErr != nil {
f.logger.Warn("CreatePermission 失败",
zap.String("peer_addr", peerAddr.String()),
zap.Error(permErr))
} else {
f.logger.Debug("CreatePermission 成功",
zap.String("peer_addr", peerAddr.String()))
}
}
}
f.logger.Debug("TURN-TCP 分配成功",
zap.String("relay_addr", relayConn.LocalAddr().String()))
return relayConn, nil
}
// parseServerAddr 解析服务器地址
func parseServerAddr(server string) (host, port string) {
h, p, _ := net.SplitHostPort(server)
if h == "" {
h = server
p = "3478" // 默认 TURN 端口
}
return h, p
}
// turnConn TURN 连接包装器
type turnConn struct {
relay net.PacketConn
remoteAddr net.Addr // 对端地址
buffer []byte
logger *zap.Logger
mu sync.Mutex
}
// newTURNConn 创建 TURN 连接
func newTURNConn(relay net.PacketConn, logger *zap.Logger) *turnConn {
return &turnConn{
relay: relay,
buffer: make([]byte, 65535),
logger: logger,
}
}
// SetRemoteAddr 设置对端地址(必须在 Write 之前调用)
func (c *turnConn) SetRemoteAddr(addr net.Addr) {
c.mu.Lock()
defer c.mu.Unlock()
c.remoteAddr = addr
}
func (c *turnConn) Read(b []byte) (n int, err error) {
c.mu.Lock()
defer c.mu.Unlock()
n, _, err = c.relay.ReadFrom(b)
return n, err
}
func (c *turnConn) Write(b []byte) (n int, err error) {
c.mu.Lock()
defer c.mu.Unlock()
if c.remoteAddr == nil {
return 0, fmt.Errorf("未设置对端地址,请先调用 SetRemoteAddr()")
}
n, err = c.relay.WriteTo(b, c.remoteAddr)
return n, err
}
func (c *turnConn) Close() error {
return c.relay.Close()
}
func (c *turnConn) LocalAddr() net.Addr {
return c.relay.LocalAddr()
}
func (c *turnConn) RemoteAddr() net.Addr {
return nil // TURN 中继没有固定的 RemoteAddr
}
func (c *turnConn) SetDeadline(t time.Time) error {
return c.relay.SetDeadline(t)
}
func (c *turnConn) SetReadDeadline(t time.Time) error {
return c.relay.SetReadDeadline(t)
}
func (c *turnConn) SetWriteDeadline(t time.Time) error {
return c.relay.SetWriteDeadline(t)
}
// tcpPacketConn TCP PacketConn 包装器
type tcpPacketConn struct {
conn net.Conn
logger *zap.Logger
}
func newTCPPacketConn(conn net.Conn, logger *zap.Logger) *tcpPacketConn {
return &tcpPacketConn{
conn: conn,
logger: logger,
}
}
func (p *tcpPacketConn) ReadFrom(b []byte) (n int, addr net.Addr, err error) {
n, err = p.conn.Read(b)
return n, p.conn.RemoteAddr(), err
}
func (p *tcpPacketConn) WriteTo(b []byte, addr net.Addr) (n int, err error) {
return p.conn.Write(b)
}
func (p *tcpPacketConn) Close() error {
return p.conn.Close()
}
func (p *tcpPacketConn) LocalAddr() net.Addr {
return p.conn.LocalAddr()
}
func (p *tcpPacketConn) SetDeadline(t time.Time) error {
return p.conn.SetDeadline(t)
}
func (p *tcpPacketConn) SetReadDeadline(t time.Time) error {
return p.conn.SetReadDeadline(t)
}
func (p *tcpPacketConn) SetWriteDeadline(t time.Time) error {
return p.conn.SetWriteDeadline(t)
}
+253
View File
@@ -0,0 +1,253 @@
package connect
import (
"context"
"crypto/rand"
"crypto/rsa"
"crypto/tls"
"crypto/x509"
"encoding/pem"
"fmt"
"math/big"
"net"
"time"
"github.com/quic-go/quic-go"
"go.uber.org/zap"
)
// QUICListener QUIC 监听器
type QUICListener struct {
listener *quic.Listener
}
// NewQUICListener 创建 QUIC 监听器
func NewQUICListener(addr string, logger *zap.Logger) (*QUICListener, error) {
// 生成自签名证书(用于测试)
cert, err := generateSelfSignedCert()
if err != nil {
return nil, fmt.Errorf("生成证书失败:%w", err)
}
tlsConf := &tls.Config{
Certificates: []tls.Certificate{cert},
NextProtos: []string{"meshray-quic"},
}
udpAddr, err := net.ResolveUDPAddr("udp", addr)
if err != nil {
return nil, err
}
udpConn, err := net.ListenUDP("udp", udpAddr)
if err != nil {
return nil, err
}
listener, err := quic.Listen(udpConn, tlsConf, nil)
if err != nil {
return nil, fmt.Errorf("创建 QUIC 监听器失败:%w", err)
}
logger.Info("QUIC 监听器已启动", zap.String("addr", addr))
return &QUICListener{
listener: listener,
}, nil
}
// Accept 接受 QUIC 连接
func (l *QUICListener) Accept(ctx context.Context) (*quic.Conn, error) {
return l.listener.Accept(ctx)
}
// Close 关闭监听器
func (l *QUICListener) Close() error {
return l.listener.Close()
}
// QUICClient QUIC 客户端
type QUICClient struct {
servers []string
logger *zap.Logger
}
// NewQUICClient 创建 QUIC 客户端
func NewQUICClient(servers []string, logger *zap.Logger) *QUICClient {
return &QUICClient{
servers: servers,
logger: logger,
}
}
// Connect 建立 QUIC 连接
func (c *QUICClient) Connect(ctx context.Context) (net.Conn, error) {
if len(c.servers) == 0 {
return nil, fmt.Errorf("未配置 QUIC 服务器")
}
// 使用不安全的 TLS 配置(跳过证书验证,用于测试)
tlsConf := &tls.Config{
InsecureSkipVerify: true,
NextProtos: []string{"meshray-quic"},
}
// 尝试连接第一个服务器
for _, server := range c.servers {
_, err := net.ResolveUDPAddr("udp", server)
if err != nil {
c.logger.Warn("解析 QUIC 服务器地址失败",
zap.String("server", server),
zap.Error(err))
continue
}
var conn *quic.Conn
conn, err = quic.DialAddr(ctx, server, tlsConf, nil)
if err == nil {
c.logger.Info("QUIC 连接已建立",
zap.String("server", server),
zap.String("local_addr", conn.LocalAddr().String()))
return newQUICConn(conn), nil
}
c.logger.Warn("QUIC 连接失败",
zap.String("server", server),
zap.Error(err))
}
return nil, fmt.Errorf("所有 QUIC 服务器连接失败")
}
// quicConn QUIC 连接包装器(实现 net.Conn
type quicConn struct {
conn *quic.Conn // quic-go v0.59.0 使用 *quic.Conn
stream *quic.Stream // 使用 *quic.Stream
}
// newQUICConn 创建 QUIC 连接包装器
func newQUICConn(conn *quic.Conn) *quicConn {
return &quicConn{
conn: conn,
}
}
// OpenStream 打开流
func (c *quicConn) OpenStream() error {
stream, err := c.conn.OpenStreamSync(context.Background())
if err != nil {
return err
}
c.stream = stream
return nil
}
// Read 实现 net.Conn
func (c *quicConn) Read(b []byte) (n int, err error) {
if c.stream == nil {
stream, err := c.conn.OpenStreamSync(context.Background())
if err != nil {
return 0, err
}
c.stream = stream
}
return c.stream.Read(b)
}
// Write 实现 net.Conn
func (c *quicConn) Write(b []byte) (n int, err error) {
if c.stream == nil {
stream, err := c.conn.OpenStreamSync(context.Background())
if err != nil {
return 0, err
}
c.stream = stream
}
return c.stream.Write(b)
}
// Close 实现 net.Conn
func (c *quicConn) Close() error {
if c.stream != nil {
c.stream.Close()
}
return c.conn.CloseWithError(0, "closed")
}
// LocalAddr 实现 net.Conn
func (c *quicConn) LocalAddr() net.Addr {
return c.conn.LocalAddr()
}
// RemoteAddr 实现 net.Conn
func (c *quicConn) RemoteAddr() net.Addr {
return c.conn.RemoteAddr()
}
// SetDeadline 实现 net.Conn
func (c *quicConn) SetDeadline(t time.Time) error {
if c.stream != nil {
return (*c.stream).SetDeadline(t)
}
return nil
}
// SetReadDeadline 实现 net.Conn
func (c *quicConn) SetReadDeadline(t time.Time) error {
if c.stream != nil {
return (*c.stream).SetReadDeadline(t)
}
return nil
}
// SetWriteDeadline 实现 net.Conn
func (c *quicConn) SetWriteDeadline(t time.Time) error {
if c.stream != nil {
return (*c.stream).SetWriteDeadline(t)
}
return nil
}
// generateSelfSignedCert 生成自签名证书(仅用于测试)
func generateSelfSignedCert() (tls.Certificate, error) {
// 生成私钥
priv, err := rsa.GenerateKey(rand.Reader, 2048)
if err != nil {
return tls.Certificate{}, err
}
// 生成证书模板
template := x509.Certificate{
SerialNumber: big.NewInt(1),
NotBefore: time.Now(),
NotAfter: time.Now().Add(365 * 24 * time.Hour),
DNSNames: []string{"localhost"},
}
// 自签名
certDER, err := x509.CreateCertificate(rand.Reader, &template, &template, &priv.PublicKey, priv)
if err != nil {
return tls.Certificate{}, err
}
// 编码证书和私钥
certPEM := pem.EncodeToMemory(&pem.Block{
Type: "CERTIFICATE",
Bytes: certDER,
})
keyPEM := pem.EncodeToMemory(&pem.Block{
Type: "RSA PRIVATE KEY",
Bytes: x509.MarshalPKCS1PrivateKey(priv),
})
// 加载证书
return tls.X509KeyPair(certPEM, keyPEM)
}
// NewTURNFactoryQUIC 创建 QUIC TURN 工厂(用于 9 层降级策略)
// 注意:当前版本暂不启用 QUIC 支持,返回 nil
func NewTURNFactoryQUIC(servers []string, username, password string, logger *zap.Logger) *TURNFactory {
logger.Warn("QUIC 传输模式暂不支持,已跳过")
return nil // 暂时返回 nil,未来实现 QUIC 支持时再完善
}
+216
View File
@@ -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
}