396 lines
9.9 KiB
Go
396 lines
9.9 KiB
Go
/**
|
|
* package mediaserver
|
|
*
|
|
* WebRTC SFU (Selective Forwarding Unit) 服务
|
|
* 功能:
|
|
* 1. 接收 WebRTC 流
|
|
* 2. 转发给同房间的其他用户
|
|
* 3. 支持音视频分离转发
|
|
*/
|
|
package mediaserver
|
|
|
|
import (
|
|
"crypto/hmac"
|
|
"crypto/sha1"
|
|
"encoding/base64"
|
|
"encoding/json"
|
|
"fmt"
|
|
"log"
|
|
"strings"
|
|
"sync"
|
|
"time"
|
|
|
|
"github.com/pion/interceptor"
|
|
"github.com/pion/interceptor/pkg/intervalpli"
|
|
"github.com/pion/webrtc/v3"
|
|
"github.com/spf13/viper"
|
|
)
|
|
|
|
// SFUServer WebRTC SFU 服务器
|
|
type SFUServer struct {
|
|
config *MediaServerConfig
|
|
api *webrtc.API
|
|
mu sync.RWMutex
|
|
running bool
|
|
|
|
// Track 管理
|
|
trackLocals map[string]*webrtc.TrackLocalStaticRTP // trackID -> localTrack
|
|
}
|
|
|
|
// NewSFUServer 创建 SFU 服务器
|
|
func NewSFUServer(config *MediaServerConfig) *SFUServer {
|
|
return &SFUServer{
|
|
config: config,
|
|
trackLocals: make(map[string]*webrtc.TrackLocalStaticRTP),
|
|
}
|
|
}
|
|
|
|
// Start 启动 SFU 服务器
|
|
func (s *SFUServer) Start() {
|
|
s.mu.Lock()
|
|
if s.running {
|
|
s.mu.Unlock()
|
|
return
|
|
}
|
|
s.running = true
|
|
s.mu.Unlock()
|
|
|
|
// 创建 MediaEngine
|
|
m := &webrtc.MediaEngine{}
|
|
|
|
// 注册默认编解码器
|
|
if err := m.RegisterDefaultCodecs(); err != nil {
|
|
log.Printf("❌ [SFU] 注册编解码器失败: %v", err)
|
|
return
|
|
}
|
|
|
|
// 创建拦截器注册表
|
|
i := &interceptor.Registry{}
|
|
|
|
// 注册 PLI 拦截器(用于请求关键帧)
|
|
intervalPliFactory, err := intervalpli.NewReceiverInterceptor()
|
|
if err != nil {
|
|
log.Printf("❌ [SFU] 创建 PLI 拦截器失败: %v", err)
|
|
return
|
|
}
|
|
i.Add(intervalPliFactory)
|
|
|
|
// 使用拦截器
|
|
if err := webrtc.RegisterDefaultInterceptors(m, i); err != nil {
|
|
log.Printf("❌ [SFU] 注册拦截器失败: %v", err)
|
|
return
|
|
}
|
|
|
|
// 创建 API
|
|
s.api = webrtc.NewAPI(webrtc.WithMediaEngine(m), webrtc.WithInterceptorRegistry(i))
|
|
|
|
log.Printf("🚀 [SFU] WebRTC SFU 已启动 | 端口: %d", s.config.WebRTCPort)
|
|
}
|
|
|
|
// Stop 停止 SFU 服务器
|
|
func (s *SFUServer) Stop() {
|
|
s.mu.Lock()
|
|
defer s.mu.Unlock()
|
|
|
|
s.running = false
|
|
s.trackLocals = make(map[string]*webrtc.TrackLocalStaticRTP)
|
|
log.Println("🛑 [SFU] 已停止")
|
|
}
|
|
|
|
// CreatePeerConnection 为用户创建 PeerConnection
|
|
func (s *SFUServer) CreatePeerConnection(roomID, userID string) (*webrtc.PeerConnection, error) {
|
|
s.mu.RLock()
|
|
if !s.running || s.api == nil {
|
|
s.mu.RUnlock()
|
|
return nil, ErrSFUNotReady
|
|
}
|
|
s.mu.RUnlock()
|
|
|
|
// 获取 ICE 服务器配置
|
|
iceServers := s.getICEServers(userID)
|
|
|
|
// 创建 PeerConnection 配置
|
|
config := webrtc.Configuration{
|
|
ICEServers: iceServers,
|
|
}
|
|
|
|
// 创建 PeerConnection
|
|
pc, err := s.api.NewPeerConnection(config)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("create peer connection failed: %w", err)
|
|
}
|
|
|
|
// 监听 ICE 连接状态
|
|
pc.OnICEConnectionStateChange(func(state webrtc.ICEConnectionState) {
|
|
log.Printf("🔗 [SFU] Room:%s User:%s ICE状态: %s", roomID, userID, state.String())
|
|
|
|
if state == webrtc.ICEConnectionStateFailed || state == webrtc.ICEConnectionStateDisconnected {
|
|
// 连接断开,清理资源
|
|
s.handleDisconnect(roomID, userID)
|
|
}
|
|
})
|
|
|
|
// 监听轨道
|
|
pc.OnTrack(func(remoteTrack *webrtc.TrackRemote, receiver *webrtc.RTPReceiver) {
|
|
s.handleTrack(roomID, userID, remoteTrack, receiver)
|
|
})
|
|
|
|
return pc, nil
|
|
}
|
|
|
|
// getICEServers 获取 ICE 服务器配置
|
|
func (s *SFUServer) getICEServers(userID string) []webrtc.ICEServer {
|
|
servers := []webrtc.ICEServer{}
|
|
|
|
// 添加 STUN 服务器
|
|
turnPublicIP := viper.GetString("turn.public_ip")
|
|
turnPort := viper.GetInt("turn.listen_port")
|
|
if turnPublicIP != "" && turnPort > 0 {
|
|
stunURL := fmt.Sprintf("stun:%s:%d", turnPublicIP, turnPort)
|
|
servers = append(servers, webrtc.ICEServer{
|
|
URLs: []string{stunURL},
|
|
})
|
|
|
|
// 添加 TURN 服务器(如果启用)
|
|
if viper.GetBool("turn.enabled") {
|
|
turnURL := fmt.Sprintf("turn:%s:%d", turnPublicIP, turnPort)
|
|
// 生成 TURN 凭证
|
|
username, credential := generateTURNCredentials(userID)
|
|
servers = append(servers, webrtc.ICEServer{
|
|
URLs: []string{turnURL},
|
|
Username: username,
|
|
Credential: credential,
|
|
})
|
|
}
|
|
}
|
|
|
|
return servers
|
|
}
|
|
|
|
// handleTrack 处理远程轨道
|
|
func (s *SFUServer) handleTrack(roomID, userID string, remoteTrack *webrtc.TrackRemote, receiver *webrtc.RTPReceiver) {
|
|
trackID := fmt.Sprintf("%s_%s_%s", roomID, userID, remoteTrack.Kind().String())
|
|
|
|
log.Printf("🎬 [SFU] 收到轨道 | Room:%s User:%s Kind:%s ID:%s",
|
|
roomID, userID, remoteTrack.Kind().String(), trackID)
|
|
|
|
// 创建本地轨道用于转发
|
|
localTrack, err := webrtc.NewTrackLocalStaticRTP(
|
|
remoteTrack.Codec().RTPCodecCapability,
|
|
trackID,
|
|
fmt.Sprintf("stream_%s", userID),
|
|
)
|
|
if err != nil {
|
|
log.Printf("❌ [SFU] 创建本地轨道失败: %v", err)
|
|
return
|
|
}
|
|
|
|
// 保存轨道
|
|
s.mu.Lock()
|
|
s.trackLocals[trackID] = localTrack
|
|
s.mu.Unlock()
|
|
|
|
// 转发 RTP 包
|
|
go func() {
|
|
buf := make([]byte, 1500)
|
|
for {
|
|
n, _, readErr := remoteTrack.Read(buf)
|
|
if readErr != nil {
|
|
log.Printf("⚠️ [SFU] 读取轨道失败 | Track:%s Error:%v", trackID, readErr)
|
|
break
|
|
}
|
|
|
|
// 写入本地轨道(会自动转发给所有订阅者)
|
|
if _, writeErr := localTrack.Write(buf[:n]); writeErr != nil {
|
|
log.Printf("⚠️ [SFU] 写入轨道失败 | Track:%s Error:%v", trackID, writeErr)
|
|
break
|
|
}
|
|
}
|
|
|
|
// 清理轨道
|
|
s.mu.Lock()
|
|
delete(s.trackLocals, trackID)
|
|
s.mu.Unlock()
|
|
}()
|
|
|
|
// 通知房间内其他用户有新轨道
|
|
s.notifyNewTrack(roomID, userID, localTrack)
|
|
}
|
|
|
|
// notifyNewTrack 通知房间内其他用户有新轨道
|
|
func (s *SFUServer) notifyNewTrack(roomID, senderUserID string, track *webrtc.TrackLocalStaticRTP) {
|
|
ms := GetServer()
|
|
room := ms.GetRoom(roomID)
|
|
if room == nil {
|
|
return
|
|
}
|
|
|
|
// 获取其他参与者
|
|
for _, p := range room.GetOtherParticipants(senderUserID) {
|
|
if p.Type != ParticipantTypeWebRTC {
|
|
continue // 只处理 WebRTC 用户
|
|
}
|
|
|
|
if p.PeerConnection == nil {
|
|
continue
|
|
}
|
|
|
|
pc, ok := p.PeerConnection.(*webrtc.PeerConnection)
|
|
if !ok {
|
|
continue
|
|
}
|
|
|
|
// 添加轨道到对方的 PeerConnection
|
|
if _, err := pc.AddTrack(track); err != nil {
|
|
log.Printf("⚠️ [SFU] 添加轨道到用户 %s 失败: %v", p.UserID, err)
|
|
}
|
|
}
|
|
}
|
|
|
|
// handleDisconnect 处理断开连接
|
|
func (s *SFUServer) handleDisconnect(roomID, userID string) {
|
|
ms := GetServer()
|
|
room := ms.GetRoom(roomID)
|
|
if room == nil {
|
|
return
|
|
}
|
|
|
|
room.RemoveParticipant(userID)
|
|
|
|
// 清理该用户的所有轨道
|
|
s.mu.Lock()
|
|
for trackID := range s.trackLocals {
|
|
// 检查 trackID 是否属于该用户
|
|
prefix := fmt.Sprintf("%s_%s_", roomID, userID)
|
|
if strings.HasPrefix(trackID, prefix) {
|
|
delete(s.trackLocals, trackID)
|
|
}
|
|
}
|
|
s.mu.Unlock()
|
|
|
|
// 如果房间空了,移除房间
|
|
if room.IsEmpty() {
|
|
ms.RemoveRoom(roomID)
|
|
}
|
|
}
|
|
|
|
// GetTrackCount 获取轨道数量
|
|
func (s *SFUServer) GetTrackCount() int {
|
|
s.mu.RLock()
|
|
defer s.mu.RUnlock()
|
|
return len(s.trackLocals)
|
|
}
|
|
|
|
// HandleOffer 处理 SDP Offer
|
|
func (s *SFUServer) HandleOffer(roomID, userID string, offerSDP string) (string, error) {
|
|
ms := GetServer()
|
|
room := ms.GetOrCreateRoom(roomID)
|
|
|
|
// 添加参与者
|
|
participant, err := room.AddParticipant(userID, ParticipantTypeWebRTC)
|
|
if err != nil {
|
|
return "", err
|
|
}
|
|
|
|
// 创建 PeerConnection
|
|
pc, err := s.CreatePeerConnection(roomID, userID)
|
|
if err != nil {
|
|
return "", err
|
|
}
|
|
participant.PeerConnection = pc
|
|
|
|
// 设置远程描述
|
|
offer := webrtc.SessionDescription{
|
|
Type: webrtc.SDPTypeOffer,
|
|
SDP: offerSDP,
|
|
}
|
|
if err := pc.SetRemoteDescription(offer); err != nil {
|
|
return "", fmt.Errorf("set remote description failed: %w", err)
|
|
}
|
|
|
|
// 添加房间内其他用户的轨道
|
|
s.addExistingTracks(pc, roomID, userID)
|
|
|
|
// 创建 Answer
|
|
answer, err := pc.CreateAnswer(nil)
|
|
if err != nil {
|
|
return "", fmt.Errorf("create answer failed: %w", err)
|
|
}
|
|
|
|
// 设置本地描述
|
|
if err := pc.SetLocalDescription(answer); err != nil {
|
|
return "", fmt.Errorf("set local description failed: %w", err)
|
|
}
|
|
|
|
// 等待 ICE 收集完成
|
|
<-webrtc.GatheringCompletePromise(pc)
|
|
|
|
return pc.LocalDescription().SDP, nil
|
|
}
|
|
|
|
// addExistingTracks 添加房间内已有的轨道
|
|
func (s *SFUServer) addExistingTracks(pc *webrtc.PeerConnection, roomID, excludeUserID string) {
|
|
s.mu.RLock()
|
|
defer s.mu.RUnlock()
|
|
|
|
prefix := fmt.Sprintf("%s_", roomID)
|
|
excludePrefix := fmt.Sprintf("%s_%s_", roomID, excludeUserID)
|
|
|
|
for trackID, track := range s.trackLocals {
|
|
// 检查是否是同房间的轨道,且不是自己的
|
|
if strings.HasPrefix(trackID, prefix) && !strings.HasPrefix(trackID, excludePrefix) {
|
|
if _, err := pc.AddTrack(track); err != nil {
|
|
log.Printf("⚠️ [SFU] 添加已有轨道失败: %v", err)
|
|
}
|
|
}
|
|
}
|
|
}
|
|
|
|
// HandleICECandidate 处理 ICE 候选
|
|
func (s *SFUServer) HandleICECandidate(roomID, userID string, candidateJSON string) error {
|
|
ms := GetServer()
|
|
room := ms.GetRoom(roomID)
|
|
if room == nil {
|
|
return ErrRoomNotFound
|
|
}
|
|
|
|
participant := room.GetParticipant(userID)
|
|
if participant == nil {
|
|
return ErrParticipantNotFound
|
|
}
|
|
|
|
pc, ok := participant.PeerConnection.(*webrtc.PeerConnection)
|
|
if !ok || pc == nil {
|
|
return ErrSFUNotReady
|
|
}
|
|
|
|
var candidate webrtc.ICECandidateInit
|
|
if err := json.Unmarshal([]byte(candidateJSON), &candidate); err != nil {
|
|
return fmt.Errorf("parse ICE candidate failed: %w", err)
|
|
}
|
|
|
|
if err := pc.AddICECandidate(candidate); err != nil {
|
|
return fmt.Errorf("add ICE candidate failed: %w", err)
|
|
}
|
|
|
|
return nil
|
|
}
|
|
|
|
// generateTURNCredentials 生成 TURN 凭证
|
|
func generateTURNCredentials(userID string) (string, string) {
|
|
timestamp := time.Now().Add(24 * time.Hour).Unix()
|
|
username := fmt.Sprintf("%d:%s", timestamp, userID)
|
|
secret := viper.GetString("turn.shared_secret")
|
|
|
|
mac := hmac.New(sha1.New, []byte(secret))
|
|
mac.Write([]byte(username))
|
|
password := base64.StdEncoding.EncodeToString(mac.Sum(nil))
|
|
|
|
return username, password
|
|
}
|
|
|
|
// GetTURNCredentials 获取 TURN 凭证(供外部调用)
|
|
func GetTURNCredentials(userID string) (string, string) {
|
|
return generateTURNCredentials(userID)
|
|
}
|