Initial commit
This commit is contained in:
+272
@@ -0,0 +1,272 @@
|
||||
# MeshRay-Core 架构规范
|
||||
|
||||
> Core 是通用的数据传输引擎,通过 ProtocolPlugin 接口适配不同协议。当前默认内置 WG 插件。
|
||||
|
||||
---
|
||||
|
||||
## 一、目录结构
|
||||
|
||||
```
|
||||
core/
|
||||
├── core.go # 进程入口
|
||||
├── engine.go # 引擎实例
|
||||
├── metrics.go # 监控指标
|
||||
│
|
||||
├── connect/ # 建连层
|
||||
│ ├── strategy.go # 策略调度
|
||||
│ ├── stun.go # STUN 协议
|
||||
│ ├── direct.go # Layer 1
|
||||
│ ├── fake_tcp.go # Layer 2
|
||||
│ ├── real_tcp.go # Layer 3
|
||||
│ ├── turn.go # Layer 4/6/7
|
||||
│ ├── turn_quic.go # Layer 5
|
||||
│ ├── ice.go # Layer 8
|
||||
│ └── ws.go # Layer 9
|
||||
│
|
||||
├── transport/ # 传输层
|
||||
│ ├── conn_manager.go # 连接索引
|
||||
│ ├── relay.go # 转发循环
|
||||
│ └── plugin.go # Plugin 接口
|
||||
│
|
||||
└── plugins/ # 协议插件
|
||||
└── wg/
|
||||
└── wgparse.go # WG 插件实现
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## 二、各文件职责
|
||||
|
||||
### 2.1 根目录(4 个文件)
|
||||
|
||||
| 文件 | 职责 | 持有什么 | 不做什么 |
|
||||
|------|------|---------|----------|
|
||||
| `core.go` | 进程入口,管理多个 Engine | `map[engineID]*Engine` | 不做建连、不转发数据 |
|
||||
| `engine.go` | 一个组网的引擎实例 | strategy、conn_manager、relay、plugin | 不直接调用 connect,由 relay 调用 |
|
||||
| `metrics.go` | 监控指标采集 | 原子计数器(连接数、字节数、切换次数) | 不做业务逻辑 |
|
||||
|
||||
**关键关系**:
|
||||
- `core.go` 持有多个 `engine.go`
|
||||
- `internal/ctr/ctr.go` 直接调用 `core.go` 和 `engine.go` (进程内函数调用)
|
||||
- `engine.go` 持有 connect/、transport/、plugins/ 的实例
|
||||
|
||||
### 2.2 connect/(9 个文件)— 建连层
|
||||
|
||||
**职责**:通过各种网络方式建立连接,最终返回 `net.Conn`。
|
||||
|
||||
**对外暴露的唯一入口**:`strategy.go` 的 `Connect()` 方法。其他 connect 文件只被 `strategy.go` 调用。
|
||||
|
||||
| 文件 | 对应层级 | 职责 | 返回什么 |
|
||||
|------|---------|------|---------|
|
||||
| `strategy.go` | 全部 | 按优先级尝试各层,不通自动切换,定期探测恢复 | `net.Conn` + 当前层级名 |
|
||||
| `stun.go` | 被 direct/ice 调用 | STUN 协议:发送 Binding Request,获取本机公网地址 | `*net.UDPAddr`(地址,不是连接) |
|
||||
| `direct.go` | Layer 1 | 调用 stun 获取候选地址,然后 UDP 打洞 | `net.Conn` |
|
||||
| `fake_tcp.go` | Layer 2 | UDP 包外层封装 TCP 头部,欺骗防火墙 | `net.Conn` |
|
||||
| `real_tcp.go` | Layer 3 | 真正的 TCP 直连打洞 | `net.Conn` |
|
||||
| `turn.go` | Layer 4/6/7 | TURN 协议协商(Allocate/Permission/ChannelBind),参数区分 UDP/TCP/TLS | `net.Conn` |
|
||||
| `turn_quic.go` | Layer 5 | TURN-QUIC 私有扩展(RFC 9000) | `net.Conn` |
|
||||
| `ice.go` | Layer 8 | ICE 协商 + WebRTC DataChannel,内部调用 stun 收集候选 | `net.Conn` |
|
||||
| `ws.go` | Layer 9 | WS/WSS 握手 + 帧收发 + 身份标识 | `net.Conn` |
|
||||
|
||||
**9 层完整编号**:
|
||||
|
||||
| 层级 | 链路名称 | 文件 | 传输方式 | 穿透力 |
|
||||
|------|---------|------|---------|--------|
|
||||
| 1 | Direct-UDP | direct.go | P2P 直连 | 弱(性能最好) |
|
||||
| 2 | Direct-FakeTCP | fake_tcp.go | P2P 直连 | 弱 |
|
||||
| 3 | Direct-RealTCP | real_tcp.go | P2P 直连 | 中 |
|
||||
| 4 | TURN-UDP | turn.go | 中继 | 中 |
|
||||
| 5 | TURN-QUIC | turn_quic.go | 中继 | 中 |
|
||||
| 6 | TURN-TCP | turn.go | 中继 | 强 |
|
||||
| 7 | TURN-TLS | turn.go | 中继 | 强 |
|
||||
| 8 | WebRTC | ice.go | ICE/TURN | 强 |
|
||||
| 9 | WS/WSS | ws.go | 隧道 | 最强(兜底) |
|
||||
|
||||
**strategy.go 的自动切换逻辑**:
|
||||
- 单包超时 500ms → 切到下一层
|
||||
- 10s 滑动窗口丢包率 > 10% → 切到下一层
|
||||
- 当前在第 N 层时,每 30s 探测 Layer 1 → 连续 2 次成功直接切回 Layer 1(不逐层回退)
|
||||
|
||||
**stun.go 的特殊地位**:唯一被多处调用的 connect 文件(direct.go 和 ice.go 都需要它),所以独立存在。
|
||||
|
||||
### 2.3 transport/(3 个文件)— 传输层
|
||||
|
||||
**职责**:用 `net.Conn` 转发数据。不感知具体协议,通过 ProtocolPlugin 接口适配。
|
||||
|
||||
| 文件 | 职责 | 不做什么 |
|
||||
|------|------|---------|
|
||||
| `conn_manager.go` | 连接索引:`peer_key → net.Conn` 的映射 | 不做建连、不转发数据 |
|
||||
| `relay.go` | Read/Write 循环:从本地端口收包 → 查路由 → 通过 conn 发送;从 conn 收包 → 发到本地端口 | 不做建连 |
|
||||
| `plugin.go` | 定义 ProtocolPlugin 接口 | 不实现任何协议 |
|
||||
|
||||
**relay.go 的工作流程**:
|
||||
|
||||
```
|
||||
本地端口收到 WG 密文包
|
||||
→ 调用 plugin.IsControlPacket() 判断包类型
|
||||
→ true:控制包,按已建链路透传
|
||||
→ 调用 plugin.IsDataPacket() 判断
|
||||
→ true:调用 plugin.ExtractRouteID() 提取路由标识
|
||||
→ 查路由标识映射表 → 发往对应本地端口
|
||||
→ 都不是:丢弃
|
||||
```
|
||||
|
||||
**关键**:relay.go 不知道 WireGuard,不知道 receiver index,只知道 route_id。
|
||||
|
||||
### 2.4 plugins/wg/(1 个文件)— 协议插件
|
||||
|
||||
**职责**:实现 ProtocolPlugin 接口,处理 WG 协议特有的包解析。
|
||||
|
||||
| 文件 | 职责 | 不做什么 |
|
||||
|------|------|---------|
|
||||
| `wgparse.go` | 实现 `IsDataPacket`/`ExtractRouteID`/`IsControlPacket` | 不做建连、不转发数据 |
|
||||
|
||||
**WG 插件的具体实现**:
|
||||
|
||||
| 方法 | 逻辑 |
|
||||
|------|------|
|
||||
| `IsControlPacket(packet)` | `packet[0]` ∈ {1, 2, 3} → true |
|
||||
| `IsDataPacket(packet)` | `packet[0]` == 4 → true |
|
||||
| `ExtractRouteID(packet)` | 读取 `packet[4:8]`,网络字节序解析为 uint32(即 WG receiver index) |
|
||||
|
||||
### 2.5 plugins/wg/(1 个文件)
|
||||
|
||||
| 文件 | 职责 |
|
||||
|------|------|
|
||||
| `wgparse.go` | WireGuard 数据包解析和封装 |
|
||||
|
||||
---
|
||||
|
||||
## 三、分层架构图
|
||||
|
||||
```
|
||||
┌─────────────────────────────────────────────────────────┐
|
||||
│ internal/ctr/ctr.go │
|
||||
│ 直接调用 Core (进程内函数调用) │
|
||||
└────────────────────────┬────────────────────────────────┘
|
||||
│ 函数调用
|
||||
┌────────────────────────▼────────────────────────────────┐
|
||||
│ core.go │
|
||||
│ 管理多个 Engine 实例 │
|
||||
└───┬─────────────────────────────────────────────────────┘
|
||||
│ 每个组网一个 Engine
|
||||
▼
|
||||
┌─────────────────────────────────────────────────────────┐
|
||||
│ engine.go │
|
||||
│ 持有:strategy + conn_manager + relay + plugin │
|
||||
│ 协调 connect/ 和 transport/ 工作 │
|
||||
└───┬─────────────────────────────────────────────────────┘
|
||||
│
|
||||
├──────────────────────────────────────────┐
|
||||
▼ ▼
|
||||
┌───────────────────────┐ ┌───────────────────────┐
|
||||
│ connect/ │ │ transport/ │
|
||||
│ 建连层 │ │ 传输层 │
|
||||
│ │ │ │
|
||||
│ strategy.go │ 返回 │ relay.go │
|
||||
│ ├─ direct.go (L1) │ net.Conn├─ plugin.go │
|
||||
│ ├─ fake_tcp.go (L2) │────────►│ │ ProtocolPlugin │
|
||||
│ ├─ real_tcp.go (L3) │ │ │ │
|
||||
│ ├─ turn.go (L4/6/7) │ │ 插件调用 │
|
||||
│ ├─ turn_quic.go(L5) │ │ ▼ │
|
||||
│ ├─ ice.go (L8) │ │ plugins/wg/ │
|
||||
│ └─ ws.go (L9) │ │ └─ wgparse.go │
|
||||
│ │ │ │
|
||||
│ stun.go (被 direct/ice 调用) │ conn_manager.go │
|
||||
└───────────────────────┘ └───────────────────────┘
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## 四、调用关系
|
||||
|
||||
### 4.1 Engine 创建时
|
||||
|
||||
```
|
||||
engine.go
|
||||
→ 创建 WGPlugin(plugins/wg/wgparse.go)
|
||||
→ 创建 Relay,传入 plugin(transport/relay.go)
|
||||
→ 创建 Strategy(connect/strategy.go)
|
||||
→ 创建 ConnManager(transport/conn_manager.go)
|
||||
```
|
||||
|
||||
### 4.2 Bind 流程
|
||||
|
||||
```
|
||||
ctr 直接调用:Bind()
|
||||
→ engine.go 接收函数调用
|
||||
→ engine.go 调用 strategy.Connect()
|
||||
→ strategy 按优先级尝试各层
|
||||
→ Layer 1: direct.go 调用 stun.go 获取候选,尝试 UDP 打洞
|
||||
→ 不通?→ Layer 4: turn.go 调用 TURN 协商
|
||||
→ 不通?→ Layer 9: ws.go 调用 WS 握手
|
||||
→ 返回 net.Conn + 当前层级名
|
||||
→ engine.go 把 net.Conn 注册到 conn_manager
|
||||
→ engine.go 启动 relay 的 Read/Write 循环
|
||||
```
|
||||
|
||||
### 4.3 数据转发流程
|
||||
|
||||
```
|
||||
WG 发出密文包 → 本地端口
|
||||
→ relay.go 收到
|
||||
→ 调用 plugin.IsControlPacket()
|
||||
→ true:按已建链路透传(conn_manager 查 conn)
|
||||
→ 调用 plugin.IsDataPacket()
|
||||
→ true:调用 plugin.ExtractRouteID() 获取 route_id
|
||||
→ 查路由标识映射表 → 找到本地端口 → 发送
|
||||
→ 都不是:丢弃
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## 五、通用层与插件层的边界
|
||||
|
||||
| 层 | 知道什么 | 不知道什么 |
|
||||
|---|---------|-----------|
|
||||
| **connect/** | 网络协议(STUN/TURN/WS/WebRTC) | WireGuard、route_id |
|
||||
| **transport/** | net.Conn、route_id、ProtocolPlugin 接口 | WireGuard、receiver index |
|
||||
| **plugins/wg/** | WG 包格式、receiver index | 网络连接、net.Conn |
|
||||
| **engine.go** | 协调 connect/ 和 transport/ | WG 包格式细节 |
|
||||
|
||||
**如果将来要支持其他协议**:
|
||||
- 新建 `plugins/xxx/xxxparse.go`
|
||||
- 实现 `ProtocolPlugin` 接口的三个方法
|
||||
- `engine.go` 里换成 `xxx.NewPlugin()`
|
||||
- connect/、transport/、core.go 的代码完全不用改
|
||||
|
||||
---
|
||||
|
||||
## 六、Core 接口(直接被 ctr 调用)
|
||||
|
||||
| 方法 | 调用方 | 说明 |
|
||||
|------|--------|------|
|
||||
| `CreateEngine` | ctr | 创建一个 Engine 实例(直接函数调用) |
|
||||
| `Bind` | ctr | 为每个 Peer 开启本地端口,开始建链 |
|
||||
| `Unbind` | ctr | 停止指定 Peer 的端口监听 |
|
||||
| `Start` | ctr | 启动转发主循环 |
|
||||
| `Stop` | ctr | 停止 Engine |
|
||||
| `GetStatus` | ctr | 查询 Engine 状态 |
|
||||
| `NotifyPeerInfo` | ctr | 下发对端候选地址和 route_id |
|
||||
|
||||
**实现位置**:
|
||||
- 所有方法都在 `engine.go` 中实现
|
||||
- `core.go` 提供 Engine 实例管理
|
||||
- ctr通过`coreInst.CreateEngine(...)`直接调用
|
||||
|
||||
---
|
||||
|
||||
## 七、文件清单汇总
|
||||
|
||||
| 目录 | 文件数 | 文件 |
|
||||
|------|--------|------|
|
||||
| 根目录 | 3 | core.go, engine.go, metrics.go |
|
||||
| connect/ | 9 | strategy.go, stun.go, direct.go, fake_tcp.go, real_tcp.go, turn.go, turn_quic.go, ice.go, ws.go |
|
||||
| transport/ | 3 | conn_manager.go, relay.go, plugin.go |
|
||||
| plugins/wg/ | 1 | wgparse.go |
|
||||
| **总计** | **16** | |
|
||||
|
||||
**已删除的文件**:
|
||||
- ~~grpc_service.go~~ - 不再需要(改为直接函数调用)
|
||||
- ~~pool/connpool.go~~ - 不再需要(无连接池)
|
||||
- ~~proto/core.proto~~ - 不再需要 gRPC
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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-FakeTCP(UDP 封装 TCP 头部,欺骗防火墙)
|
||||
LayerFakeTCP
|
||||
|
||||
// LayerRealTCP Direct-RealTCP(P2P 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 DataChannel(DTLS 加密)
|
||||
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
|
||||
}
|
||||
@@ -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 服务器均查询失败")
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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 支持时再完善
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
+118
@@ -0,0 +1,118 @@
|
||||
package core
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"sync"
|
||||
|
||||
"go.uber.org/zap"
|
||||
)
|
||||
|
||||
// Core 进程入口 - 管理多个 Engine 实例
|
||||
type Core struct {
|
||||
engines map[string]*Engine // engineID -> Engine
|
||||
mu sync.RWMutex
|
||||
logger *zap.Logger
|
||||
}
|
||||
|
||||
// NewCore 创建 Core 实例(进程入口)
|
||||
func NewCore(logger *zap.Logger) *Core {
|
||||
return &Core{
|
||||
engines: make(map[string]*Engine),
|
||||
logger: logger,
|
||||
}
|
||||
}
|
||||
|
||||
// CreateEngine 创建 Engine 实例
|
||||
func (c *Core) CreateEngine(engineID string, metrics *Metrics) (*Engine, error) {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
|
||||
// 检查是否已存在
|
||||
if _, ok := c.engines[engineID]; ok {
|
||||
return nil, fmt.Errorf("engine %s already exists", engineID)
|
||||
}
|
||||
|
||||
// 创建新 Engine
|
||||
engine := NewEngine(c.logger, metrics)
|
||||
c.engines[engineID] = engine
|
||||
|
||||
c.logger.Info("创建 Engine 实例",
|
||||
zap.String("engine_id", engineID))
|
||||
|
||||
return engine, nil
|
||||
}
|
||||
|
||||
// GetEngine 获取 Engine 实例
|
||||
func (c *Core) GetEngine(engineID string) (*Engine, error) {
|
||||
c.mu.RLock()
|
||||
defer c.mu.RUnlock()
|
||||
|
||||
engine, ok := c.engines[engineID]
|
||||
if !ok {
|
||||
return nil, fmt.Errorf("engine %s not found", engineID)
|
||||
}
|
||||
|
||||
return engine, nil
|
||||
}
|
||||
|
||||
// RemoveEngine 移除 Engine 实例
|
||||
func (c *Core) RemoveEngine(engineID string) error {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
|
||||
engine, ok := c.engines[engineID]
|
||||
if !ok {
|
||||
return fmt.Errorf("engine %s not found", engineID)
|
||||
}
|
||||
|
||||
// 停止 Engine
|
||||
if err := engine.Stop(); err != nil {
|
||||
c.logger.Warn("停止 Engine 失败",
|
||||
zap.String("engine_id", engineID),
|
||||
zap.Error(err))
|
||||
}
|
||||
|
||||
delete(c.engines, engineID)
|
||||
c.logger.Info("移除 Engine 实例",
|
||||
zap.String("engine_id", engineID))
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// ListEngines 列出所有 Engine ID
|
||||
func (c *Core) ListEngines() []string {
|
||||
c.mu.RLock()
|
||||
defer c.mu.RUnlock()
|
||||
|
||||
ids := make([]string, 0, len(c.engines))
|
||||
for id := range c.engines {
|
||||
ids = append(ids, id)
|
||||
}
|
||||
return ids
|
||||
}
|
||||
|
||||
// Count 获取 Engine 数量
|
||||
func (c *Core) Count() int {
|
||||
c.mu.RLock()
|
||||
defer c.mu.RUnlock()
|
||||
return len(c.engines)
|
||||
}
|
||||
|
||||
// Close 关闭所有 Engine
|
||||
func (c *Core) Close() error {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
|
||||
for id, engine := range c.engines {
|
||||
if err := engine.Stop(); err != nil {
|
||||
c.logger.Warn("停止 Engine 失败",
|
||||
zap.String("engine_id", id),
|
||||
zap.Error(err))
|
||||
}
|
||||
}
|
||||
|
||||
c.engines = make(map[string]*Engine)
|
||||
c.logger.Info("关闭所有 Engine")
|
||||
|
||||
return nil
|
||||
}
|
||||
+333
@@ -0,0 +1,333 @@
|
||||
package core
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"net"
|
||||
"time"
|
||||
|
||||
"git.zkcoi.com/zkcoi/meshray/core/connect"
|
||||
"git.zkcoi.com/zkcoi/meshray/core/plugins/wg"
|
||||
"git.zkcoi.com/zkcoi/meshray/core/transport"
|
||||
"go.uber.org/zap"
|
||||
)
|
||||
|
||||
// Engine 引擎实例 - 一个组网的引擎实例
|
||||
type Engine struct {
|
||||
logger *zap.Logger
|
||||
scheduler *connect.StrategyScheduler
|
||||
connMgr *transport.ConnManager
|
||||
relay *transport.Relay
|
||||
plugin transport.ProtocolPlugin
|
||||
metrics *Metrics
|
||||
|
||||
// 候选地址存储(用于 NotifyPeerInfo)
|
||||
candidateStore map[string][]Candidate
|
||||
routeIDStore map[string]uint32
|
||||
}
|
||||
|
||||
// NewEngine 创建引擎实例
|
||||
func NewEngine(logger *zap.Logger, metrics *Metrics) *Engine {
|
||||
// 创建 WG 插件
|
||||
plugin := wg.NewWGPlugin()
|
||||
|
||||
// 创建连接管理器
|
||||
connMgr := transport.NewConnManager(logger)
|
||||
|
||||
// 创建数据转发器
|
||||
relay := transport.NewRelay(plugin, connMgr, logger)
|
||||
|
||||
// 创建策略调度器
|
||||
scheduler := connect.NewStrategyScheduler(logger)
|
||||
scheduler.OnConnectionUpdate = func(peerID string, conn net.Conn, err error) {
|
||||
if err != nil {
|
||||
logger.Warn("收到策略调度器连接错误更新", zap.String("peer_id", peerID), zap.Error(err))
|
||||
}
|
||||
if conn != nil {
|
||||
logger.Info("策略调度器连接已建立,开始双向转发", zap.String("peer_id", peerID))
|
||||
connMgr.Add(peerID, conn)
|
||||
// 启动远端接收协程
|
||||
relay.StartReadFromRemoteConn(context.Background(), peerID, conn)
|
||||
}
|
||||
}
|
||||
|
||||
scheduler.RegisterFactory(connect.NewDirectFactory(nil, logger)) // 1. Direct-UDP
|
||||
scheduler.RegisterFactory(connect.NewFakeTCPFactory(logger)) // 2. FakeTCP
|
||||
scheduler.RegisterFactory(connect.NewRealTCPFactory(logger)) // 3. RealTCP
|
||||
scheduler.RegisterFactory(connect.NewTURNFactory(connect.TURNProtocolUDP, nil, "", "", logger)) // 4. TURN-UDP
|
||||
// scheduler.RegisterFactory(connect.NewTURNFactoryQUIC(nil, "", "", logger)) // 5. TURN-QUIC (暂不启用)
|
||||
scheduler.RegisterFactory(connect.NewTURNFactory(connect.TURNProtocolTCP, nil, "", "", logger)) // 6. TURN-TCP
|
||||
scheduler.RegisterFactory(connect.NewTURNFactory(connect.TURNProtocolTLS, nil, "", "", logger)) // 7. TURN-TLS
|
||||
scheduler.RegisterFactory(connect.NewWebRTCFactory(&connect.ICEConfig{}, logger)) // 8. WebRTC
|
||||
scheduler.RegisterFactory(connect.NewWSFactory(nil, logger)) // 9. WS/WSS
|
||||
|
||||
engine := &Engine{
|
||||
logger: logger,
|
||||
scheduler: scheduler,
|
||||
connMgr: connMgr,
|
||||
relay: relay,
|
||||
plugin: plugin,
|
||||
metrics: metrics,
|
||||
candidateStore: make(map[string][]Candidate),
|
||||
routeIDStore: make(map[string]uint32),
|
||||
}
|
||||
|
||||
// 设置拨号触发器:当 WG 发包但没连接时自动 9 层拨号
|
||||
relay.OnDialTrigger = func(peerKey string) {
|
||||
go engine.initiateConnection(peerKey)
|
||||
}
|
||||
|
||||
return engine
|
||||
}
|
||||
|
||||
// Start 启动引擎
|
||||
func (e *Engine) Start() error {
|
||||
e.logger.Info("Core 引擎启动")
|
||||
return nil
|
||||
}
|
||||
|
||||
// Stop 停止引擎
|
||||
func (e *Engine) Stop() error {
|
||||
e.logger.Info("Core 引擎停止")
|
||||
|
||||
// 关闭所有连接
|
||||
if e.connMgr != nil {
|
||||
e.connMgr.CloseAll()
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// SetICEConfig 设置 ICE 配置(用于 WebRTC)
|
||||
func (e *Engine) SetICEConfig(config connect.ICEConfig) error {
|
||||
e.logger.Info("更新 ICE 配置",
|
||||
zap.Int("stun_servers", len(config.STUNServers)),
|
||||
zap.Int("turn_servers", len(config.TURNServers)))
|
||||
|
||||
// TODO: 实现 ICE 配置更新逻辑
|
||||
// 1. 找到 WebRTC 工厂
|
||||
// 2. 更新其 ICE 配置
|
||||
// 3. 重新注册工厂
|
||||
|
||||
// 目前先记录日志,P3 阶段实现
|
||||
e.logger.Warn("SetICEConfig 暂未实现,将在 P3 阶段完成")
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// GetScheduler 获取策略调度器
|
||||
func (e *Engine) GetScheduler() *connect.StrategyScheduler {
|
||||
return e.scheduler
|
||||
}
|
||||
|
||||
// GetConnMgr 获取连接管理器
|
||||
func (e *Engine) GetConnMgr() *transport.ConnManager {
|
||||
return e.connMgr
|
||||
}
|
||||
|
||||
// GetRelay 获取数据转发器
|
||||
func (e *Engine) GetRelay() *transport.Relay {
|
||||
return e.relay
|
||||
}
|
||||
|
||||
// GetMetrics 获取监控指标
|
||||
func (e *Engine) GetMetrics() *Metrics {
|
||||
return e.metrics
|
||||
}
|
||||
|
||||
// Bind 为指定 Peer 开启本地端口,开始建连
|
||||
// peerKey: 对端公钥哈希(8 字符)
|
||||
// localPort: 本地监听端口(传 0 表示系统自动分配)
|
||||
// 返回值:实际绑定的端口号
|
||||
func (e *Engine) Bind(peerKey string, localPort int) (int, error) {
|
||||
// 1. 在本地端口监听
|
||||
addr := &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1), Port: localPort}
|
||||
conn, err := net.ListenUDP("udp", addr)
|
||||
if err != nil {
|
||||
return 0, fmt.Errorf("监听本地端口失败:%w", err)
|
||||
}
|
||||
|
||||
actualPort := conn.LocalAddr().(*net.UDPAddr).Port
|
||||
|
||||
// 2. 提取 route_id(从 peerKey 派生)
|
||||
routeID := extractRouteID(peerKey)
|
||||
|
||||
// 3. 注册到 Relay
|
||||
e.relay.RegisterLocalPort(routeID, conn)
|
||||
|
||||
// 4. 注册到 ConnManager
|
||||
e.connMgr.Add(peerKey, nil) // conn 初始为 nil,建连后设置
|
||||
|
||||
// 5. 启动读取协程
|
||||
ctx := context.Background()
|
||||
e.relay.StartReadFromLocalPort(ctx, routeID, peerKey)
|
||||
|
||||
e.logger.Info("Bind 成功",
|
||||
zap.String("peer_key", peerKey),
|
||||
zap.Int("local_port", actualPort),
|
||||
zap.Uint32("route_id", routeID))
|
||||
|
||||
return actualPort, nil
|
||||
}
|
||||
|
||||
// Unbind 停止指定 Peer 的端口监听
|
||||
func (e *Engine) Unbind(peerKey string) error {
|
||||
// 1. 提取 route_id
|
||||
routeID := extractRouteID(peerKey)
|
||||
|
||||
// 2. 从 Relay 注销
|
||||
e.relay.UnregisterLocalPort(routeID)
|
||||
|
||||
// 3. 从 ConnManager 移除
|
||||
e.connMgr.Remove(peerKey)
|
||||
|
||||
// 4. 清理存储
|
||||
delete(e.candidateStore, peerKey)
|
||||
delete(e.routeIDStore, peerKey)
|
||||
|
||||
e.logger.Info("Unbind 成功",
|
||||
zap.String("peer_key", peerKey),
|
||||
zap.Uint32("route_id", routeID))
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// EngineStatus Engine 状态
|
||||
type EngineStatus struct {
|
||||
PeerCount int `json:"peer_count"`
|
||||
Peers map[string]*PeerStatus `json:"peers"`
|
||||
|
||||
// Metrics
|
||||
ActiveConnections int64 `json:"active_connections"`
|
||||
TotalConnections int64 `json:"total_connections"`
|
||||
BytesSent uint64 `json:"bytes_sent"`
|
||||
BytesReceived uint64 `json:"bytes_received"`
|
||||
StrategyFallbacks int64 `json:"strategy_fallbacks"`
|
||||
LastSwitchTime int64 `json:"last_switch_time"` // unix timestamp
|
||||
}
|
||||
|
||||
// PeerStatus Peer 状态
|
||||
type PeerStatus struct {
|
||||
PeerKey string `json:"peer_key"`
|
||||
Connected bool `json:"connected"`
|
||||
Layer string `json:"layer,omitempty"` // 当前传输层
|
||||
}
|
||||
|
||||
// GetStatus 查询 Engine 状态
|
||||
func (e *Engine) GetStatus() (*EngineStatus, error) {
|
||||
status := &EngineStatus{
|
||||
PeerCount: e.connMgr.Count(),
|
||||
Peers: make(map[string]*PeerStatus),
|
||||
}
|
||||
|
||||
// 收集 Metrics
|
||||
if e.metrics != nil {
|
||||
status.ActiveConnections = e.metrics.GetActiveConnections()
|
||||
status.TotalConnections = e.metrics.GetTotalConnections()
|
||||
status.BytesSent = e.metrics.GetBytesSent()
|
||||
status.BytesReceived = e.metrics.GetBytesReceived()
|
||||
status.StrategyFallbacks = e.metrics.GetStrategyFallbacks()
|
||||
status.LastSwitchTime = e.metrics.GetLastSwitchTime().Unix()
|
||||
}
|
||||
|
||||
// 收集所有 Peer 状态
|
||||
for peerKey, conn := range e.connMgr.GetAll() {
|
||||
peerStatus := &PeerStatus{
|
||||
PeerKey: peerKey,
|
||||
Connected: conn != nil,
|
||||
}
|
||||
if conn != nil {
|
||||
peerStatus.Layer = e.scheduler.GetPeerLayer(peerKey).String()
|
||||
}
|
||||
status.Peers[peerKey] = peerStatus
|
||||
}
|
||||
|
||||
return status, nil
|
||||
}
|
||||
|
||||
// Candidate 候选地址(与 connect.Candidate 对齐)
|
||||
type Candidate struct {
|
||||
Addr string `json:"addr"` // 候选地址(ip:port)
|
||||
Type string `json:"type"` // 候选类型:host/srflx/relay
|
||||
Priority int `json:"priority"` // 优先级
|
||||
Protocol string `json:"protocol"` // 协议:udp/tcp
|
||||
}
|
||||
|
||||
// NotifyPeerInfo 下发对端候选地址和 route_id
|
||||
// peerKey: 对端公钥哈希
|
||||
// candidates: 对端候选地址列表(由信使服务器转发)
|
||||
// routeID: 路由 ID(用于数据转发)
|
||||
func (e *Engine) NotifyPeerInfo(peerKey string, candidates []Candidate, routeID uint32) error {
|
||||
// 1. 存储候选地址(用于后续建连)
|
||||
e.candidateStore[peerKey] = candidates
|
||||
|
||||
// 2. 存储 route_id 映射
|
||||
e.routeIDStore[peerKey] = routeID
|
||||
|
||||
// 3. 触发建连流程
|
||||
go e.initiateConnection(peerKey)
|
||||
|
||||
e.logger.Info("NotifyPeerInfo 成功",
|
||||
zap.String("peer_key", peerKey),
|
||||
zap.Int("candidate_count", len(candidates)),
|
||||
zap.Uint32("route_id", routeID))
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// initiateConnection 触发建连流程
|
||||
func (e *Engine) initiateConnection(peerKey string) {
|
||||
// 1. 获取候选地址
|
||||
candidates := e.candidateStore[peerKey]
|
||||
if len(candidates) == 0 {
|
||||
return
|
||||
}
|
||||
|
||||
// 2. 检查是否已经在拨号或已连接
|
||||
if conn, ok := e.connMgr.Get(peerKey); ok && conn != nil {
|
||||
return
|
||||
}
|
||||
|
||||
// 2. 初始化 DialConfig
|
||||
config := &connect.DialConfig{
|
||||
PeerID: peerKey,
|
||||
Timeout: 10 * time.Second,
|
||||
Logger: e.logger,
|
||||
}
|
||||
|
||||
e.logger.Info("开始建立连接到对端",
|
||||
zap.String("peer_key", peerKey),
|
||||
zap.Int("candidate_count", len(candidates)))
|
||||
|
||||
// 3. 获取 RouteID
|
||||
_, ok := e.routeIDStore[peerKey]
|
||||
if !ok {
|
||||
e.logger.Warn("未找到 route_id",
|
||||
zap.String("peer_key", peerKey))
|
||||
return
|
||||
}
|
||||
|
||||
// 4. 使用策略调度器尝试建连
|
||||
conn, err := e.scheduler.Dial(config)
|
||||
if err != nil {
|
||||
e.logger.Error("所有策略层尝试连接均失败",
|
||||
zap.String("peer_key", peerKey),
|
||||
zap.Error(err))
|
||||
return
|
||||
}
|
||||
|
||||
// 5. 连接成功,更新到 ConnManager
|
||||
e.connMgr.Add(peerKey, conn)
|
||||
e.logger.Info("连接建立并更新成功", zap.String("peer_key", peerKey))
|
||||
}
|
||||
|
||||
// extractRouteID 从 peerKey 提取 route_id(简化版本)
|
||||
// 实际应该使用一致的哈希算法
|
||||
func extractRouteID(peerKey string) uint32 {
|
||||
// 简单哈希:取前 4 个字符的 ASCII 码和
|
||||
var sum uint32 = 0
|
||||
for i := 0; i < len(peerKey) && i < 4; i++ {
|
||||
sum += uint32(peerKey[i])
|
||||
}
|
||||
return sum
|
||||
}
|
||||
@@ -0,0 +1,83 @@
|
||||
package core
|
||||
|
||||
import (
|
||||
"sync/atomic"
|
||||
"time"
|
||||
)
|
||||
|
||||
// Metrics Core 监控指标
|
||||
type Metrics struct {
|
||||
// 连接统计
|
||||
activeConnections atomic.Int64
|
||||
totalConnections atomic.Int64
|
||||
|
||||
// 流量统计
|
||||
bytesSent atomic.Uint64
|
||||
bytesReceived atomic.Uint64
|
||||
|
||||
// 策略统计
|
||||
strategyFallbacks atomic.Int64
|
||||
lastSwitchTime atomic.Int64 // Unix timestamp
|
||||
}
|
||||
|
||||
// NewMetrics 创建监控指标
|
||||
func NewMetrics() *Metrics {
|
||||
return &Metrics{}
|
||||
}
|
||||
|
||||
// GetActiveConnections 获取活跃连接数
|
||||
func (m *Metrics) GetActiveConnections() int64 {
|
||||
return m.activeConnections.Load()
|
||||
}
|
||||
|
||||
// IncrActiveConnections 增加活跃连接数
|
||||
func (m *Metrics) IncrActiveConnections() {
|
||||
m.activeConnections.Add(1)
|
||||
m.totalConnections.Add(1)
|
||||
}
|
||||
|
||||
// DecrActiveConnections 减少活跃连接数
|
||||
func (m *Metrics) DecrActiveConnections() {
|
||||
m.activeConnections.Add(-1)
|
||||
}
|
||||
|
||||
// AddBytesSent 增加发送字节数
|
||||
func (m *Metrics) AddBytesSent(n uint64) {
|
||||
m.bytesSent.Add(n)
|
||||
}
|
||||
|
||||
// AddBytesReceived 增加接收字节数
|
||||
func (m *Metrics) AddBytesReceived(n uint64) {
|
||||
m.bytesReceived.Add(n)
|
||||
}
|
||||
|
||||
// IncrStrategyFallbacks 增加策略降级次数
|
||||
func (m *Metrics) IncrStrategyFallbacks() {
|
||||
m.strategyFallbacks.Add(1)
|
||||
m.lastSwitchTime.Store(time.Now().Unix())
|
||||
}
|
||||
|
||||
// GetTotalConnections 获取总连接数
|
||||
func (m *Metrics) GetTotalConnections() int64 {
|
||||
return m.totalConnections.Load()
|
||||
}
|
||||
|
||||
// GetBytesSent 获取发送字节数
|
||||
func (m *Metrics) GetBytesSent() uint64 {
|
||||
return m.bytesSent.Load()
|
||||
}
|
||||
|
||||
// GetBytesReceived 获取接收字节数
|
||||
func (m *Metrics) GetBytesReceived() uint64 {
|
||||
return m.bytesReceived.Load()
|
||||
}
|
||||
|
||||
// GetStrategyFallbacks 获取策略降级次数
|
||||
func (m *Metrics) GetStrategyFallbacks() int64 {
|
||||
return m.strategyFallbacks.Load()
|
||||
}
|
||||
|
||||
// GetLastSwitchTime 获取最后切换时间
|
||||
func (m *Metrics) GetLastSwitchTime() time.Time {
|
||||
return time.Unix(m.lastSwitchTime.Load(), 0)
|
||||
}
|
||||
@@ -0,0 +1,47 @@
|
||||
package wg
|
||||
|
||||
import (
|
||||
"encoding/binary"
|
||||
"fmt"
|
||||
)
|
||||
|
||||
// WGPlugin WireGuard 协议插件实现
|
||||
type WGPlugin struct{}
|
||||
|
||||
// NewWGPlugin 创建 WireGuard 协议插件
|
||||
func NewWGPlugin() *WGPlugin {
|
||||
return &WGPlugin{}
|
||||
}
|
||||
|
||||
// IsControlPacket 判断是否为控制包
|
||||
// WG 控制包类型:1 (Initiation), 2 (Response), 3 (CookieReply)
|
||||
func (p *WGPlugin) IsControlPacket(packet []byte) bool {
|
||||
if len(packet) < 1 {
|
||||
return false
|
||||
}
|
||||
|
||||
packetType := packet[0]
|
||||
return packetType == 1 || packetType == 2 || packetType == 3
|
||||
}
|
||||
|
||||
// IsDataPacket 判断是否为数据包
|
||||
// WG 数据包类型:4
|
||||
func (p *WGPlugin) IsDataPacket(packet []byte) bool {
|
||||
if len(packet) < 1 {
|
||||
return false
|
||||
}
|
||||
|
||||
return packet[0] == 4
|
||||
}
|
||||
|
||||
// ExtractRouteID 从数据包中提取路由标识(WG receiver index)
|
||||
// WG 数据包格式:[类型 (1 字节)][保留 (3 字节)][receiver index (4 字节)]...
|
||||
func (p *WGPlugin) ExtractRouteID(packet []byte) (uint32, error) {
|
||||
if len(packet) < 8 {
|
||||
return 0, fmt.Errorf("数据包过短:%d", len(packet))
|
||||
}
|
||||
|
||||
// 读取 packet[4:8],网络字节序解析为 uint32
|
||||
routeID := binary.BigEndian.Uint32(packet[4:8])
|
||||
return routeID, nil
|
||||
}
|
||||
@@ -0,0 +1,99 @@
|
||||
package transport
|
||||
|
||||
import (
|
||||
"net"
|
||||
"sync"
|
||||
|
||||
"go.uber.org/zap"
|
||||
)
|
||||
|
||||
// ConnManager 连接管理器
|
||||
// 维护 peer_key → net.Conn 的映射关系
|
||||
type ConnManager struct {
|
||||
conns map[string]net.Conn
|
||||
mu sync.RWMutex
|
||||
logger *zap.Logger
|
||||
}
|
||||
|
||||
// NewConnManager 创建连接管理器
|
||||
func NewConnManager(logger *zap.Logger) *ConnManager {
|
||||
return &ConnManager{
|
||||
conns: make(map[string]net.Conn),
|
||||
logger: logger,
|
||||
}
|
||||
}
|
||||
|
||||
// Add 添加连接
|
||||
func (m *ConnManager) Add(peerKey string, conn net.Conn) {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
|
||||
// 如果已存在,先关闭旧连接
|
||||
if oldConn, ok := m.conns[peerKey]; ok {
|
||||
oldConn.Close()
|
||||
m.logger.Debug("关闭旧连接", zap.String("peer_key", peerKey))
|
||||
}
|
||||
|
||||
m.conns[peerKey] = conn
|
||||
m.logger.Info("添加新连接",
|
||||
zap.String("peer_key", peerKey),
|
||||
zap.String("remote_addr", conn.RemoteAddr().String()))
|
||||
}
|
||||
|
||||
// Get 获取连接
|
||||
func (m *ConnManager) Get(peerKey string) (net.Conn, bool) {
|
||||
m.mu.RLock()
|
||||
defer m.mu.RUnlock()
|
||||
|
||||
conn, ok := m.conns[peerKey]
|
||||
return conn, ok
|
||||
}
|
||||
|
||||
// Remove 移除连接
|
||||
func (m *ConnManager) Remove(peerKey string) {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
|
||||
if conn, ok := m.conns[peerKey]; ok {
|
||||
conn.Close()
|
||||
delete(m.conns, peerKey)
|
||||
m.logger.Info("移除连接", zap.String("peer_key", peerKey))
|
||||
}
|
||||
}
|
||||
|
||||
// Count 获取连接数量
|
||||
func (m *ConnManager) Count() int {
|
||||
m.mu.RLock()
|
||||
defer m.mu.RUnlock()
|
||||
return len(m.conns)
|
||||
}
|
||||
|
||||
// CloseAll 关闭所有连接
|
||||
func (m *ConnManager) CloseAll() {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
|
||||
for peerKey, conn := range m.conns {
|
||||
conn.Close()
|
||||
m.logger.Debug("关闭连接", zap.String("peer_key", peerKey))
|
||||
}
|
||||
|
||||
m.conns = make(map[string]net.Conn)
|
||||
}
|
||||
|
||||
// List 列出所有连接(返回副本)
|
||||
func (m *ConnManager) List() map[string]net.Conn {
|
||||
m.mu.RLock()
|
||||
defer m.mu.RUnlock()
|
||||
|
||||
result := make(map[string]net.Conn)
|
||||
for k, v := range m.conns {
|
||||
result[k] = v
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
// GetAll 获取所有连接(同 List,为了兼容)
|
||||
func (m *ConnManager) GetAll() map[string]net.Conn {
|
||||
return m.List()
|
||||
}
|
||||
@@ -0,0 +1,17 @@
|
||||
package transport
|
||||
|
||||
// ProtocolPlugin 协议插件接口
|
||||
// relay.go 通过这个接口适配不同协议,不感知具体协议细节
|
||||
type ProtocolPlugin interface {
|
||||
// IsControlPacket 判断是否为控制包
|
||||
// 控制包用于建连协商,需要透传到对端
|
||||
IsControlPacket(packet []byte) bool
|
||||
|
||||
// IsDataPacket 判断是否为数据包
|
||||
// 数据包包含路由标识,需要查表转发
|
||||
IsDataPacket(packet []byte) bool
|
||||
|
||||
// ExtractRouteID 从数据包中提取路由标识
|
||||
// 返回的 route_id 用于查找对应的本地端口
|
||||
ExtractRouteID(packet []byte) (uint32, error)
|
||||
}
|
||||
@@ -0,0 +1,197 @@
|
||||
package transport
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net"
|
||||
"sync"
|
||||
|
||||
"go.uber.org/zap"
|
||||
)
|
||||
|
||||
// OnDialTrigger 当从本地端口收到包但没有远端连接时触发
|
||||
type OnDialTrigger func(peerKey string)
|
||||
|
||||
// Relay 数据转发器
|
||||
// 负责从本地端口收包 → 查路由 → 通过 conn 发送
|
||||
// 从 conn 收包 → 发到本地端口
|
||||
type Relay struct {
|
||||
plugin ProtocolPlugin
|
||||
connMgr *ConnManager
|
||||
localPorts map[uint32]net.PacketConn // route_id → local_port
|
||||
lastAddr map[uint32]net.Addr // route_id → last wg source addr
|
||||
portMu sync.RWMutex
|
||||
logger *zap.Logger
|
||||
OnDialTrigger OnDialTrigger // 拨号触发回调
|
||||
}
|
||||
|
||||
// NewRelay 创建数据转发器
|
||||
func NewRelay(plugin ProtocolPlugin, connMgr *ConnManager, logger *zap.Logger) *Relay {
|
||||
return &Relay{
|
||||
plugin: plugin,
|
||||
connMgr: connMgr,
|
||||
localPorts: make(map[uint32]net.PacketConn),
|
||||
lastAddr: make(map[uint32]net.Addr),
|
||||
logger: logger,
|
||||
}
|
||||
}
|
||||
|
||||
// RegisterLocalPort 注册本地端口(用于接收 WG 密文包)
|
||||
func (r *Relay) RegisterLocalPort(routeID uint32, port net.PacketConn) {
|
||||
r.portMu.Lock()
|
||||
defer r.portMu.Unlock()
|
||||
|
||||
r.localPorts[routeID] = port
|
||||
r.logger.Info("注册本地端口",
|
||||
zap.Uint32("route_id", routeID),
|
||||
zap.String("addr", port.LocalAddr().String()))
|
||||
}
|
||||
|
||||
// UnregisterLocalPort 注销本地端口
|
||||
func (r *Relay) UnregisterLocalPort(routeID uint32) {
|
||||
r.portMu.Lock()
|
||||
defer r.portMu.Unlock()
|
||||
|
||||
if port, ok := r.localPorts[routeID]; ok {
|
||||
port.Close()
|
||||
delete(r.localPorts, routeID)
|
||||
delete(r.lastAddr, routeID)
|
||||
r.logger.Info("注销本地端口", zap.Uint32("route_id", routeID))
|
||||
}
|
||||
}
|
||||
|
||||
// StartReadFromLocalPort 从本地端口读取 WG 密文包并转发(发送到远端)
|
||||
func (r *Relay) StartReadFromLocalPort(ctx context.Context, routeID uint32, peerKey string) {
|
||||
r.portMu.RLock()
|
||||
port, ok := r.localPorts[routeID]
|
||||
r.portMu.RUnlock()
|
||||
|
||||
if !ok {
|
||||
r.logger.Warn("本地端口未注册", zap.Uint32("route_id", routeID))
|
||||
return
|
||||
}
|
||||
|
||||
go func() {
|
||||
buf := make([]byte, 65535)
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return
|
||||
default:
|
||||
n, addr, err := port.ReadFrom(buf)
|
||||
if err != nil {
|
||||
// 检查是否是由于关闭引起的错误
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return
|
||||
default:
|
||||
}
|
||||
r.logger.Debug("读取本地端口失败",
|
||||
zap.Uint32("route_id", routeID),
|
||||
zap.Error(err))
|
||||
return
|
||||
}
|
||||
|
||||
// 记录 WG 的来源地址,以便后续把包发回去
|
||||
r.portMu.Lock()
|
||||
r.lastAddr[routeID] = addr
|
||||
r.portMu.Unlock()
|
||||
|
||||
packet := buf[:n]
|
||||
r.forwardOutgoing(ctx, packet, peerKey)
|
||||
}
|
||||
}
|
||||
}()
|
||||
|
||||
r.logger.Info("启动本地端口读取协程",
|
||||
zap.Uint32("route_id", routeID),
|
||||
zap.String("peer_key", peerKey))
|
||||
}
|
||||
|
||||
// StartReadFromRemoteConn 从远端连接读取数据并转发给本地监听端口(接收远端数据)
|
||||
func (r *Relay) StartReadFromRemoteConn(ctx context.Context, peerKey string, conn net.Conn) {
|
||||
if conn == nil {
|
||||
return
|
||||
}
|
||||
|
||||
go func() {
|
||||
buf := make([]byte, 65535)
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return
|
||||
default:
|
||||
n, err := conn.Read(buf)
|
||||
if err != nil {
|
||||
r.logger.Debug("读取远端连接失败,停止读取协程",
|
||||
zap.String("peer_key", peerKey),
|
||||
zap.Error(err))
|
||||
return
|
||||
}
|
||||
|
||||
packet := buf[:n]
|
||||
// 远端进来的包,需要根据 packet 里的索引转发给对应的 localPort
|
||||
r.forwardIncoming(packet, peerKey)
|
||||
}
|
||||
}
|
||||
}()
|
||||
|
||||
r.logger.Info("启动远端连接读取协程",
|
||||
zap.String("peer_key", peerKey),
|
||||
zap.String("addr", conn.RemoteAddr().String()))
|
||||
}
|
||||
|
||||
// forwardOutgoing 处理发出去的包(Local -> Remote)
|
||||
func (r *Relay) forwardOutgoing(ctx context.Context, packet []byte, peerKey string) {
|
||||
// 获取或触发建连
|
||||
conn, ok := r.connMgr.Get(peerKey)
|
||||
if !ok || conn == nil {
|
||||
// 没有连接,触发拨号
|
||||
if r.OnDialTrigger != nil {
|
||||
r.OnDialTrigger(peerKey)
|
||||
}
|
||||
r.logger.Debug("尚未建立连接,包已丢弃,触发静默拨号", zap.String("peer_key", peerKey))
|
||||
return
|
||||
}
|
||||
|
||||
// 转发给远端
|
||||
_, err := conn.Write(packet)
|
||||
if err != nil {
|
||||
r.logger.Debug("转发包到远端失败",
|
||||
zap.String("peer_key", peerKey),
|
||||
zap.Error(err))
|
||||
}
|
||||
}
|
||||
|
||||
// forwardIncoming 处理进来的包(Remote -> Local)
|
||||
func (r *Relay) forwardIncoming(packet []byte, _ string) {
|
||||
// 1. 判断是否为控制包/数据包并提取 routeID
|
||||
// 无论哪种 WG 包,前几位都是 routeID (receiver index)
|
||||
routeID, err := r.plugin.ExtractRouteID(packet)
|
||||
if err != nil {
|
||||
r.logger.Debug("提取包内索引失败", zap.Error(err))
|
||||
return
|
||||
}
|
||||
|
||||
// 2. 这里的 routeID 是我们 RegisterLocalPort 时用的 ID
|
||||
r.portMu.RLock()
|
||||
port, ok := r.localPorts[routeID]
|
||||
addr, addrOk := r.lastAddr[routeID]
|
||||
r.portMu.RUnlock()
|
||||
|
||||
if !ok || port == nil {
|
||||
r.logger.Debug("未找到转发目标的本地端口", zap.Uint32("route_id", routeID))
|
||||
return
|
||||
}
|
||||
|
||||
if !addrOk || addr == nil {
|
||||
// 如果还没收到过 WG 的包,尝试发给 127.0.0.1:0 (通常不会成功,但作为 fallback)
|
||||
// 实际上 WG 发送握手包后就会刷新 addr
|
||||
addr = &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1), Port: 0}
|
||||
}
|
||||
|
||||
// 3. 转发给本地 WG
|
||||
_, err = port.WriteTo(packet, addr)
|
||||
if err != nil {
|
||||
r.logger.Debug("转发给本地 WG 失败", zap.Uint32("route_id", routeID), zap.Error(err))
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user