Initial commit
This commit is contained in:
@@ -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
|
||||
}
|
||||
Reference in New Issue
Block a user