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

586 lines
13 KiB
Go
Raw Blame History

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