586 lines
13 KiB
Go
586 lines
13 KiB
Go
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
|
||
}
|