Files
nl-im-service/internal/mediaserver/ws_rtmp_proxy.go
2025-12-15 21:57:21 +08:00

737 lines
18 KiB
Go
Raw Blame History

This file contains invisible Unicode characters
This file contains invisible Unicode characters that are indistinguishable to humans but may be processed differently by a computer. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
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 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)