359 lines
8.6 KiB
Go
359 lines
8.6 KiB
Go
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)
|
||
}
|