/** * package mediaserver * * WebSocket to RTMP 代理服务 * 接收 Web端 通过 WebSocket 发送的 WebM 数据,转换为 FLV 并推送到 RTMP * * 技术栈: * - WebSocket: gorilla/websocket * - WebM 解析: github.com/at-wat/ebml-go * - FLV 封装: github.com/yutopp/go-flv * - RTMP 推送: 通过内部流管道 */ package mediaserver import ( "bytes" "encoding/binary" "fmt" "io" "log" "net/http" "sync" "time" "github.com/at-wat/ebml-go" "github.com/gorilla/websocket" ) // WebMToRTMPProxy WebSocket to RTMP 代理 type WebMToRTMPProxy struct { rtmpServer *RTMPServer upgrader websocket.Upgrader sessions map[string]*ProxySession sessionsMu sync.RWMutex } // ProxySession 代理会话 type ProxySession struct { ID string StreamID string UserID string RoomID string Conn *websocket.Conn Stream *RTMPStream StopChan chan struct{} StartTime time.Time // WebM 解析状态 webmBuffer *bytes.Buffer headerParsed bool videoTrackNum uint64 audioTrackNum uint64 // 时间戳 baseTimestamp uint32 lastTimestamp uint32 } // WebMHeader WebM 文件头信息 type WebMHeader struct { EBMLVersion uint64 `ebml:"EBMLVersion"` EBMLReadVersion uint64 `ebml:"EBMLReadVersion"` EBMLMaxIDLength uint64 `ebml:"EBMLMaxIDLength"` EBMLMaxSizeLength uint64 `ebml:"EBMLMaxSizeLength"` DocType string `ebml:"DocType"` DocTypeVersion uint64 `ebml:"DocTypeVersion"` DocTypeReadVersion uint64 `ebml:"DocTypeReadVersion"` } // WebMSegment WebM Segment type WebMSegment struct { Info WebMSegmentInfo `ebml:"Info"` Tracks WebMTracks `ebml:"Tracks"` Cluster []WebMCluster `ebml:"Cluster"` } // WebMSegmentInfo Segment 信息 type WebMSegmentInfo struct { TimecodeScale uint64 `ebml:"TimecodeScale"` Duration float64 `ebml:"Duration,omitempty"` MuxingApp string `ebml:"MuxingApp,omitempty"` WritingApp string `ebml:"WritingApp,omitempty"` } // WebMTracks 轨道信息 type WebMTracks struct { TrackEntry []WebMTrackEntry `ebml:"TrackEntry"` } // WebMTrackEntry 轨道条目 type WebMTrackEntry struct { TrackNumber uint64 `ebml:"TrackNumber"` TrackType uint64 `ebml:"TrackType"` // 1=video, 2=audio CodecID string `ebml:"CodecID"` Video *WebMVideoTrack `ebml:"Video,omitempty"` Audio *WebMAudioTrack `ebml:"Audio,omitempty"` } // WebMVideoTrack 视频轨道 type WebMVideoTrack struct { PixelWidth uint64 `ebml:"PixelWidth"` PixelHeight uint64 `ebml:"PixelHeight"` } // WebMAudioTrack 音频轨道 type WebMAudioTrack struct { SamplingFrequency float64 `ebml:"SamplingFrequency"` Channels uint64 `ebml:"Channels"` BitDepth uint64 `ebml:"BitDepth,omitempty"` } // WebMCluster WebM Cluster type WebMCluster struct { Timecode uint64 `ebml:"Timecode"` SimpleBlock []ebml.Block `ebml:"SimpleBlock,omitempty"` } // NewWebMToRTMPProxy 创建代理服务 func NewWebMToRTMPProxy(rtmpServer *RTMPServer) *WebMToRTMPProxy { return &WebMToRTMPProxy{ rtmpServer: rtmpServer, upgrader: websocket.Upgrader{ CheckOrigin: func(r *http.Request) bool { return true // 允许跨域 }, ReadBufferSize: 1024 * 1024, WriteBufferSize: 1024 * 1024, }, sessions: make(map[string]*ProxySession), } } // HandleWebSocket 处理 WebSocket 连接 func (p *WebMToRTMPProxy) HandleWebSocket(w http.ResponseWriter, r *http.Request) { log.Printf("🔌 [WSProxy] 收到 WebSocket 连接请求: %s from %s", r.URL.String(), r.RemoteAddr) // 检查 WebSocket 升级头 if r.Header.Get("Upgrade") != "websocket" { log.Printf("❌ [WSProxy] 不是 WebSocket 请求: Upgrade=%s", r.Header.Get("Upgrade")) http.Error(w, "Not a WebSocket request", http.StatusBadRequest) return } // 获取参数 streamID := r.URL.Query().Get("stream_id") userID := r.URL.Query().Get("user_id") roomID := r.URL.Query().Get("room_id") token := r.URL.Query().Get("token") log.Printf("📋 [WSProxy] 参数: stream_id=%s user_id=%s room_id=%s token=%v", streamID, userID, roomID, token != "") if streamID == "" || userID == "" { log.Printf("❌ [WSProxy] 缺少必要参数") http.Error(w, "Missing stream_id or user_id", http.StatusBadRequest) return } // 验证 token(简化版,生产环境需要更严格的验证) if token == "" { log.Printf("⚠️ [WSProxy] 缺少 token: stream=%s", streamID) } // 升级到 WebSocket log.Printf("🔄 [WSProxy] 尝试 WebSocket 升级...") conn, err := p.upgrader.Upgrade(w, r, nil) if err != nil { log.Printf("❌ [WSProxy] WebSocket 升级失败: %v (可能原因: 中间件干扰、响应已写入)", err) return } log.Printf("✅ [WSProxy] WebSocket 升级成功 | Stream:%s User:%s Room:%s", streamID, userID, roomID) // 获取或创建 RTMP 流 stream := p.rtmpServer.GetStream(streamID) if stream == nil { // 自动创建流 var err error stream, err = p.rtmpServer.GenerateStreamURLs(roomID, userID) if err != nil { log.Printf("❌ [WSProxy] 创建流失败: %v", err) conn.Close() return } // 更新 streamID 为实际生成的 streamID = stream.ID } // 创建会话 session := &ProxySession{ ID: fmt.Sprintf("%s_%d", userID, time.Now().UnixNano()), StreamID: streamID, UserID: userID, RoomID: roomID, Conn: conn, Stream: stream, StopChan: make(chan struct{}), StartTime: time.Now(), webmBuffer: bytes.NewBuffer(nil), headerParsed: false, } // 注册会话 p.sessionsMu.Lock() p.sessions[session.ID] = session p.sessionsMu.Unlock() // 激活流 stream.mu.Lock() stream.IsActive = true stream.mu.Unlock() // 处理连接 go p.handleSession(session) } // handleSession 处理会话 func (p *WebMToRTMPProxy) handleSession(session *ProxySession) { log.Printf("▶️ [WSProxy] 开始处理会话 | Stream:%s User:%s", session.StreamID, session.UserID) messageCount := 0 totalBytes := int64(0) defer func() { // 清理 p.sessionsMu.Lock() delete(p.sessions, session.ID) p.sessionsMu.Unlock() session.Conn.Close() close(session.StopChan) // 标记流为非活动 if session.Stream != nil { session.Stream.mu.Lock() session.Stream.IsActive = false session.Stream.mu.Unlock() } log.Printf("🔌 [WSProxy] 连接关闭 | Stream:%s User:%s | 收到消息:%d 总字节:%d", session.StreamID, session.UserID, messageCount, totalBytes) }() for { select { case <-session.StopChan: log.Printf("⏹️ [WSProxy] 收到停止信号 | Stream:%s", session.StreamID) return default: } // 读取 WebSocket 消息 messageType, data, err := session.Conn.ReadMessage() if err != nil { if websocket.IsUnexpectedCloseError(err, websocket.CloseGoingAway, websocket.CloseAbnormalClosure) { log.Printf("⚠️ [WSProxy] 读取错误: %v | Stream:%s", err, session.StreamID) } else { log.Printf("ℹ️ [WSProxy] 连接关闭: %v | Stream:%s", err, session.StreamID) } return } messageCount++ totalBytes += int64(len(data)) // 记录首次收到消息 if messageCount == 1 { log.Printf("📥 [WSProxy] 首次收到消息 | Stream:%s Type:%d Size:%d", session.StreamID, messageType, len(data)) } // 每 100 条消息记录一次统计 if messageCount%100 == 0 { log.Printf("📊 [WSProxy] 消息统计 | Stream:%s Count:%d TotalBytes:%d", session.StreamID, messageCount, totalBytes) } if messageType != websocket.BinaryMessage { log.Printf("⚠️ [WSProxy] 忽略非二进制消息 | Type:%d", messageType) continue } // 处理 WebM 数据 if err := p.processWebMData(session, data); err != nil { log.Printf("⚠️ [WSProxy] 处理 WebM 数据失败: %v", err) } } } // processWebMData 处理 WebM 数据 func (p *WebMToRTMPProxy) processWebMData(session *ProxySession, data []byte) error { // 将数据追加到缓冲区 session.webmBuffer.Write(data) // 尝试解析 WebM 数据 return p.parseAndConvert(session) } // parseAndConvert 解析 WebM 并转换为 FLV func (p *WebMToRTMPProxy) parseAndConvert(session *ProxySession) error { bufData := session.webmBuffer.Bytes() if len(bufData) < 4 { return nil // 数据不足 } // 检查 EBML 头部 if !session.headerParsed { // 尝试解析头部 if err := p.parseWebMHeader(session, bufData); err != nil { // 头部不完整,等待更多数据 return nil } session.headerParsed = true log.Printf("📦 [WSProxy] WebM 头部解析完成 | Stream:%s", session.StreamID) } // 解析 Cluster 并转换 return p.parseWebMClusters(session) } // parseWebMHeader 解析 WebM 头部 func (p *WebMToRTMPProxy) parseWebMHeader(session *ProxySession, data []byte) error { reader := bytes.NewReader(data) // 解析 EBML 头 var header struct { EBML struct { EBMLVersion uint64 `ebml:"EBMLVersion"` DocType string `ebml:"DocType"` } `ebml:"EBML"` } if err := ebml.Unmarshal(reader, &header); err != nil { return err } log.Printf("📦 [WSProxy] WebM DocType: %s", header.EBML.DocType) return nil } // parseWebMClusters 解析 WebM Clusters func (p *WebMToRTMPProxy) parseWebMClusters(session *ProxySession) error { // 简化处理:直接将 WebM 数据转换为 FLV // 实际实现需要完整解析 WebM 的 Cluster/SimpleBlock bufData := session.webmBuffer.Bytes() if len(bufData) < 100 { return nil } // 寻找 Cluster 标记 (0x1F43B675) clusterMarker := []byte{0x1F, 0x43, 0xB6, 0x75} for { idx := bytes.Index(bufData, clusterMarker) if idx == -1 || idx+12 > len(bufData) { break } // 解析 Cluster 大小 sizeStart := idx + 4 clusterSize, bytesRead := readVarInt(bufData[sizeStart:]) if bytesRead == 0 { break } totalSize := idx + 4 + bytesRead + int(clusterSize) if totalSize > len(bufData) { // Cluster 不完整 break } // 提取 Cluster 数据 clusterData := bufData[idx:totalSize] // 转换为 FLV 并广播 if err := p.convertClusterToFLV(session, clusterData); err != nil { log.Printf("⚠️ [WSProxy] 转换 Cluster 失败: %v", err) } // 从缓冲区移除已处理的数据 bufData = bufData[totalSize:] } // 更新缓冲区 session.webmBuffer.Reset() session.webmBuffer.Write(bufData) return nil } // readVarInt 读取 EBML 变长整数 func readVarInt(data []byte) (uint64, int) { if len(data) == 0 { return 0, 0 } first := data[0] var length int var mask byte switch { case first&0x80 != 0: length = 1 mask = 0x7F case first&0x40 != 0: length = 2 mask = 0x3F case first&0x20 != 0: length = 3 mask = 0x1F case first&0x10 != 0: length = 4 mask = 0x0F case first&0x08 != 0: length = 5 mask = 0x07 case first&0x04 != 0: length = 6 mask = 0x03 case first&0x02 != 0: length = 7 mask = 0x01 case first&0x01 != 0: length = 8 mask = 0x00 default: return 0, 0 } if len(data) < length { return 0, 0 } value := uint64(data[0] & mask) for i := 1; i < length; i++ { value = (value << 8) | uint64(data[i]) } return value, length } // convertClusterToFLV 将 WebM Cluster 转换为 FLV 并广播 func (p *WebMToRTMPProxy) convertClusterToFLV(session *ProxySession, clusterData []byte) error { // 解析 Cluster 时间戳 timecodeMarker := []byte{0xE7} // Timecode element ID timecodeIdx := bytes.Index(clusterData, timecodeMarker) var timestamp uint32 = session.lastTimestamp if timecodeIdx != -1 && timecodeIdx+1 < len(clusterData) { tcSize, bytesRead := readVarInt(clusterData[timecodeIdx+1:]) if bytesRead > 0 && timecodeIdx+1+bytesRead+int(tcSize) <= len(clusterData) { tcData := clusterData[timecodeIdx+1+bytesRead : timecodeIdx+1+bytesRead+int(tcSize)] for _, b := range tcData { timestamp = (timestamp << 8) | uint32(b) } } } // 解析 SimpleBlock simpleBlockMarker := []byte{0xA3} // SimpleBlock element ID blockData := clusterData for { blockIdx := bytes.Index(blockData, simpleBlockMarker) if blockIdx == -1 || blockIdx+1 >= len(blockData) { break } // 解析 block 大小 sizeStart := blockIdx + 1 blockSize, bytesRead := readVarInt(blockData[sizeStart:]) if bytesRead == 0 { break } dataStart := sizeStart + bytesRead dataEnd := dataStart + int(blockSize) if dataEnd > len(blockData) { break } // 提取 block 内容 block := blockData[dataStart:dataEnd] if len(block) < 4 { blockData = blockData[dataEnd:] continue } // 解析 track number(变长) trackNum, trackBytes := readVarInt(block) if trackBytes == 0 { blockData = blockData[dataEnd:] continue } // 解析相对时间戳(2 bytes, big-endian) if trackBytes+2 > len(block) { blockData = blockData[dataEnd:] continue } relativeTimestamp := binary.BigEndian.Uint16(block[trackBytes : trackBytes+2]) // 解析 flags (1 byte) flagsIdx := trackBytes + 2 if flagsIdx >= len(block) { blockData = blockData[dataEnd:] continue } flags := block[flagsIdx] // 帧数据 frameData := block[flagsIdx+1:] if len(frameData) == 0 { blockData = blockData[dataEnd:] continue } // 计算绝对时间戳 absoluteTimestamp := timestamp + uint32(relativeTimestamp) session.lastTimestamp = absoluteTimestamp // 判断是视频还是音频(简化:假设 track 1 是视频,track 2 是音频) isVideo := trackNum == 1 isKeyframe := (flags & 0x80) != 0 // 创建 FLV tag 并广播 var flvTag []byte if isVideo { flvTag = p.createFLVVideoTag(absoluteTimestamp, frameData, isKeyframe) } else { flvTag = p.createFLVAudioTag(absoluteTimestamp, frameData) } if flvTag != nil { p.broadcastFLVTag(session, flvTag) } blockData = blockData[dataEnd:] } return nil } // createFLVVideoTag 创建 FLV 视频 tag // VP8 -> FLV (需要转换为 H.264,这里简化处理) func (p *WebMToRTMPProxy) createFLVVideoTag(timestamp uint32, data []byte, isKeyframe bool) []byte { // 注意:VP8 不能直接封装到 FLV,需要转码为 H.264 // 这里使用一个简化的方案:将 VP8 数据作为自定义格式封装 // 实际生产环境需要使用 FFmpeg 或硬件编码器进行转码 // FLV Video Tag Header: // FrameType (4 bits): 1=keyframe, 2=inter frame // CodecID (4 bits): 7=AVC (H.264) // 由于 VP8 无法直接放入 FLV,这里使用一个变通方案 // 将 VP8 数据标记为私有编码格式 frameType := byte(2) // inter frame if isKeyframe { frameType = 1 // keyframe } // 使用 CodecID=12 (VP8 - 非标准,仅用于内部传输) // 或者可以考虑在服务端进行实时转码 codecID := byte(12) // 自定义:VP8 header := (frameType << 4) | codecID // 构建完整数据 videoData := make([]byte, 1+len(data)) videoData[0] = header copy(videoData[1:], data) return p.createFLVTag(9, timestamp, videoData) // 9 = video } // createFLVAudioTag 创建 FLV 音频 tag // Opus -> FLV (需要转换为 AAC,这里简化处理) func (p *WebMToRTMPProxy) createFLVAudioTag(timestamp uint32, data []byte) []byte { // 注意:Opus 不能直接封装到 FLV,需要转码为 AAC // 这里使用简化方案 // FLV Audio Tag Header: // SoundFormat (4 bits): 10=AAC, 13=Opus (非标准) // SoundRate (2 bits): 3=44kHz // SoundSize (1 bit): 1=16-bit // SoundType (1 bit): 1=stereo // 使用自定义格式标记 Opus soundFormat := byte(13) // 自定义:Opus soundRate := byte(3) // 44kHz soundSize := byte(1) // 16-bit soundType := byte(1) // stereo header := (soundFormat << 4) | (soundRate << 2) | (soundSize << 1) | soundType // 构建完整数据 audioData := make([]byte, 1+len(data)) audioData[0] = header copy(audioData[1:], data) return p.createFLVTag(8, timestamp, audioData) // 8 = audio } // createFLVTag 创建 FLV tag func (p *WebMToRTMPProxy) createFLVTag(tagType byte, timestamp uint32, data []byte) []byte { dataSize := len(data) tagSize := 11 + dataSize + 4 tag := make([]byte, tagSize) // Tag type tag[0] = tagType // Data size (24-bit big-endian) tag[1] = byte((dataSize >> 16) & 0xff) tag[2] = byte((dataSize >> 8) & 0xff) tag[3] = byte(dataSize & 0xff) // Timestamp (24-bit big-endian) tag[4] = byte((timestamp >> 16) & 0xff) tag[5] = byte((timestamp >> 8) & 0xff) tag[6] = byte(timestamp & 0xff) // Timestamp extended tag[7] = byte((timestamp >> 24) & 0xff) // Stream ID (always 0) tag[8] = 0 tag[9] = 0 tag[10] = 0 // Data copy(tag[11:], data) // Previous tag size prevTagSize := 11 + dataSize tag[11+dataSize] = byte((prevTagSize >> 24) & 0xff) tag[11+dataSize+1] = byte((prevTagSize >> 16) & 0xff) tag[11+dataSize+2] = byte((prevTagSize >> 8) & 0xff) tag[11+dataSize+3] = byte(prevTagSize & 0xff) return tag } // broadcastFLVTag 广播 FLV tag 到订阅者 func (p *WebMToRTMPProxy) broadcastFLVTag(session *ProxySession, tag []byte) { if session.Stream == nil { return } // 缓存到 GOP session.Stream.mu.Lock() session.Stream.gopCache = append(session.Stream.gopCache, tag) if len(session.Stream.gopCache) > 300 { session.Stream.gopCache = session.Stream.gopCache[len(session.Stream.gopCache)-300:] } session.Stream.mu.Unlock() // 广播到 HTTP-FLV 订阅者 p.rtmpServer.subMu.RLock() subs, exists := p.rtmpServer.subscribers[session.StreamID] if exists { for _, sub := range subs { select { case sub.DataChan <- tag: default: // 缓冲区满,跳过 } } } p.rtmpServer.subMu.RUnlock() } // GetSession 获取会话 func (p *WebMToRTMPProxy) GetSession(sessionID string) *ProxySession { p.sessionsMu.RLock() defer p.sessionsMu.RUnlock() return p.sessions[sessionID] } // GetSessionsByStream 获取流的所有会话 func (p *WebMToRTMPProxy) GetSessionsByStream(streamID string) []*ProxySession { p.sessionsMu.RLock() defer p.sessionsMu.RUnlock() sessions := make([]*ProxySession, 0) for _, s := range p.sessions { if s.StreamID == streamID { sessions = append(sessions, s) } } return sessions } // CloseSession 关闭会话 func (p *WebMToRTMPProxy) CloseSession(sessionID string) { p.sessionsMu.Lock() session, exists := p.sessions[sessionID] if exists { delete(p.sessions, sessionID) } p.sessionsMu.Unlock() if session != nil { select { case <-session.StopChan: default: close(session.StopChan) } session.Conn.Close() } } // CloseAllSessions 关闭所有会话 func (p *WebMToRTMPProxy) CloseAllSessions() { p.sessionsMu.Lock() sessions := make([]*ProxySession, 0, len(p.sessions)) for _, s := range p.sessions { sessions = append(sessions, s) } p.sessions = make(map[string]*ProxySession) p.sessionsMu.Unlock() for _, s := range sessions { select { case <-s.StopChan: default: close(s.StopChan) } s.Conn.Close() } } // 确保导入被使用 var _ io.Reader = (*bytes.Reader)(nil)