Files
nl-im-service/internal/mediaserver/rtmp_protocol.go

615 lines
18 KiB
Go
Raw Permalink 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
*
* 自实现的 RTMP 协议处理
* 用于替代 github.com/yutopp/go-rtmp解决 SetChunkSize 导致的 panic 问题
*/
package mediaserver
import (
"bufio"
"encoding/binary"
"fmt"
"io"
"log"
"net"
"sync"
)
// RTMP 消息类型常量
const (
RTMP_MSG_CHUNK_SIZE = 1 // SetChunkSize
RTMP_MSG_ABORT = 2 // Abort Message
RTMP_MSG_ACK = 3 // Acknowledgement
RTMP_MSG_USER_CONTROL = 4 // User Control Message
RTMP_MSG_WIN_ACK_SIZE = 5 // Window Acknowledgement Size
RTMP_MSG_SET_PEER_BW = 6 // Set Peer Bandwidth
RTMP_MSG_AUDIO = 8 // Audio Message
RTMP_MSG_VIDEO = 9 // Video Message
RTMP_MSG_AMF3_DATA = 15 // AMF3 Data Message
RTMP_MSG_AMF3_SHARED_OBJ = 16 // AMF3 Shared Object Message
RTMP_MSG_AMF3_CMD = 17 // AMF3 Command Message
RTMP_MSG_AMF0_DATA = 18 // AMF0 Data Message (@setDataFrame)
RTMP_MSG_AMF0_SHARED_OBJ = 19 // AMF0 Shared Object Message
RTMP_MSG_AMF0_CMD = 20 // AMF0 Command Message (connect, publish, etc)
RTMP_MSG_AGGREGATE = 22 // Aggregate Message
)
// Chunk 格式类型
const (
CHUNK_FMT_0 = 0 // 11 bytes header
CHUNK_FMT_1 = 1 // 7 bytes header
CHUNK_FMT_2 = 2 // 3 bytes header
CHUNK_FMT_3 = 3 // 0 bytes header
)
// 默认值
const (
DEFAULT_CHUNK_SIZE = 128
MAX_CHUNK_SIZE = 65536
DEFAULT_WINDOW_SIZE = 2500000
RTMP_PROTOCOL_VERSION = 3
)
// RTMPMessage 表示一个完整的 RTMP 消息
type RTMPMessage struct {
ChunkStreamID uint32
Timestamp uint32
TypeID uint8
StreamID uint32
Data []byte
}
// ChunkHeader 表示 chunk 的头部信息
type ChunkHeader struct {
Format uint8 // 0-3
ChunkStreamID uint32 // 2-65599
Timestamp uint32 // 24-bit or 32-bit (extended)
MessageLength uint32 // 24-bit
MessageTypeID uint8
MessageSID uint32 // 32-bit little-endian
ExtendedTS bool // Timestamp >= 0xFFFFFF
}
// ChunkReader RTMP Chunk 读取器
type ChunkReader struct {
conn net.Conn
reader *bufio.Reader
chunkSize uint32
mu sync.RWMutex
// 缓存每个 chunk stream 的头部信息(用于 fmt 1/2/3
prevHeaders map[uint32]*ChunkHeader
// 缓存不完整的消息数据
messageBuffer map[uint32]*messageState
// 追踪已经收到过有效头部的 csid用于 fmt 2/3 验证)
knownCSIDs map[uint32]bool
}
// messageState 追踪消息的读取状态
type messageState struct {
header *ChunkHeader
data []byte
bytesRead uint32
}
// ChunkWriter RTMP Chunk 写入器
type ChunkWriter struct {
conn net.Conn
writer *bufio.Writer
chunkSize uint32
mu sync.Mutex
}
// NewChunkReader 创建新的 Chunk 读取器
func NewChunkReader(conn net.Conn) *ChunkReader {
return &ChunkReader{
conn: conn,
reader: bufio.NewReaderSize(conn, 4096),
chunkSize: DEFAULT_CHUNK_SIZE,
prevHeaders: make(map[uint32]*ChunkHeader),
messageBuffer: make(map[uint32]*messageState),
knownCSIDs: make(map[uint32]bool),
}
}
// NewChunkWriter 创建新的 Chunk 写入器
func NewChunkWriter(conn net.Conn) *ChunkWriter {
return &ChunkWriter{
conn: conn,
writer: bufio.NewWriterSize(conn, 4096),
chunkSize: DEFAULT_CHUNK_SIZE,
}
}
// SetChunkSize 设置读取的 chunk 大小
func (r *ChunkReader) SetChunkSize(size uint32) {
r.mu.Lock()
defer r.mu.Unlock()
if size > 0 && size <= MAX_CHUNK_SIZE {
log.Printf("📝 [RTMP Protocol] ChunkReader: 更新 chunkSize %d -> %d", r.chunkSize, size)
r.chunkSize = size
}
}
// GetChunkSize 获取当前 chunk 大小
func (r *ChunkReader) GetChunkSize() uint32 {
r.mu.RLock()
defer r.mu.RUnlock()
return r.chunkSize
}
// SetChunkSize 设置写入的 chunk 大小
func (w *ChunkWriter) SetChunkSize(size uint32) {
w.mu.Lock()
defer w.mu.Unlock()
if size > 0 && size <= MAX_CHUNK_SIZE {
log.Printf("📝 [RTMP Protocol] ChunkWriter: 更新 chunkSize %d -> %d", w.chunkSize, size)
w.chunkSize = size
}
}
// ReadMessage 读取一个完整的 RTMP 消息
func (r *ChunkReader) ReadMessage() (*RTMPMessage, error) {
for {
// 读取 chunk header
header, err := r.readChunkHeader()
if err != nil {
return nil, err
}
// 获取或创建消息状态
// 关键修复:检查消息是否已完成,如果是则为新消息创建新状态
state, exists := r.messageBuffer[header.ChunkStreamID]
if !exists || state.bytesRead >= state.header.MessageLength {
// 新消息或上一条消息已完成,创建新状态
// 注意:如果 header.MessageLength 为 0fmt=3 无有效 prevHeader
// 这里会创建一个容量为 0 的 slice后续 toRead 检查会捕获这种情况
state = &messageState{
header: header,
data: make([]byte, 0, header.MessageLength),
bytesRead: 0,
}
r.messageBuffer[header.ChunkStreamID] = state
} else {
// 消息续传:更新时间戳(如果 header 有新的时间戳)
if header.Timestamp > 0 {
state.header.Timestamp = header.Timestamp
}
}
// 计算本次要读取的字节数
// 关键修复:使用 state.header.MessageLength 而非 header.MessageLength
// 因为后续 chunk (fmt 1/2/3) 的 header 可能从前一个 header 继承值
remaining := state.header.MessageLength - state.bytesRead
r.mu.RLock()
toRead := r.chunkSize
r.mu.RUnlock()
if remaining < toRead {
toRead = remaining
}
// toRead=0 说明 MessageLength=0可能是 fmt=2/3 没有有效的 prevHeader
// 尝试通过搜索下一个 fmt=0 chunk 来重新同步
if toRead == 0 {
log.Printf("⚠️ [RTMP Protocol] csid=%d fmt=%d msgLen=0, 尝试重新同步...",
header.ChunkStreamID, header.Format)
// 清除此 csid 的状态
delete(r.messageBuffer, header.ChunkStreamID)
// 尝试寻找下一个有效的 chunk 起始位置
// 读取并丢弃字节,直到找到可能的 fmt=0 chunk header
maxScanBytes := 4096 // 最多扫描 4KB
scanned := 0
foundSync := false
for scanned < maxScanBytes {
b, err := r.reader.ReadByte()
if err != nil {
return nil, fmt.Errorf("重新同步时读取失败: %w", err)
}
scanned++
// 检查是否可能是 fmt=0 的 chunk header
// fmt=0 的 basic header 第一个字节: 高 2 位为 00
potentialFmt := (b >> 6) & 0x03
potentialCsid := uint32(b & 0x3F)
if potentialFmt == 0 && potentialCsid >= 2 && potentialCsid < 64 {
// 可能找到了 fmt=0 chunk尝试验证
// 先放回这个字节
if err := r.reader.UnreadByte(); err != nil {
return nil, fmt.Errorf("UnreadByte 失败: %w", err)
}
log.Printf(" [RTMP Protocol] 可能找到同步点: 跳过了 %d 字节, 潜在 csid=%d",
scanned-1, potentialCsid)
foundSync = true
break
}
}
if !foundSync {
log.Printf("❌ [RTMP Protocol] 扫描 %d 字节后仍未找到同步点", scanned)
return nil, fmt.Errorf("RTMP 协议错误: 无法重新同步")
}
continue // 重新尝试读取 chunk header
}
// 读取 chunk 数据
chunkData := make([]byte, toRead)
if _, err := io.ReadFull(r.reader, chunkData); err != nil {
return nil, fmt.Errorf("读取 chunk 数据失败: %w", err)
}
state.data = append(state.data, chunkData...)
state.bytesRead += toRead
// 检查消息是否完整
if state.bytesRead >= state.header.MessageLength {
msg := &RTMPMessage{
ChunkStreamID: header.ChunkStreamID,
Timestamp: state.header.Timestamp,
TypeID: state.header.MessageTypeID,
StreamID: state.header.MessageSID,
Data: state.data,
}
// 调试日志:记录消息完成信息
if msg.TypeID == RTMP_MSG_AMF0_DATA || msg.TypeID == RTMP_MSG_AUDIO || msg.TypeID == RTMP_MSG_VIDEO {
log.Printf("✅ [RTMP Debug] 消息完成: csid=%d type=%d len=%d",
header.ChunkStreamID, msg.TypeID, len(msg.Data))
}
// 清除消息缓冲
delete(r.messageBuffer, header.ChunkStreamID)
return msg, nil
}
}
}
// readChunkHeader 读取 chunk 头部
func (r *ChunkReader) readChunkHeader() (*ChunkHeader, error) {
// 读取第一个字节Basic Header (1-3 bytes)
firstByte, err := r.reader.ReadByte()
if err != nil {
return nil, err
}
format := (firstByte >> 6) & 0x03
csid := uint32(firstByte & 0x3F)
// 调试日志:记录每个 chunk header 的原始字节
// 注意csid > 64 或异常 csid 值可能表示流同步丢失
if csid > 20 || (format != 0 && format != 3) {
log.Printf("🔍 [RTMP Debug] firstByte=0x%02X fmt=%d csid=%d", firstByte, format, csid)
}
// 扩展 chunk stream ID
if csid == 0 {
// 2 byte header
secondByte, err := r.reader.ReadByte()
if err != nil {
return nil, err
}
csid = uint32(secondByte) + 64
} else if csid == 1 {
// 3 byte header
bytes := make([]byte, 2)
if _, err := io.ReadFull(r.reader, bytes); err != nil {
return nil, err
}
csid = uint32(bytes[0]) + uint32(bytes[1])*256 + 64
}
// 获取上一个头部(用于 fmt 1/2/3
prevHeader := r.prevHeaders[csid]
if prevHeader == nil {
prevHeader = &ChunkHeader{
ChunkStreamID: csid,
}
}
header := &ChunkHeader{
Format: format,
ChunkStreamID: csid,
Timestamp: prevHeader.Timestamp,
MessageLength: prevHeader.MessageLength,
MessageTypeID: prevHeader.MessageTypeID,
MessageSID: prevHeader.MessageSID,
}
// 根据 format 读取 Message Header
switch format {
case CHUNK_FMT_0:
// 11 bytes: timestamp(3) + length(3) + typeID(1) + streamID(4)
data := make([]byte, 11)
if _, err := io.ReadFull(r.reader, data); err != nil {
return nil, fmt.Errorf("读取 fmt0 header 失败: %w", err)
}
header.Timestamp = uint32(data[0])<<16 | uint32(data[1])<<8 | uint32(data[2])
header.MessageLength = uint32(data[3])<<16 | uint32(data[4])<<8 | uint32(data[5])
header.MessageTypeID = data[6]
header.MessageSID = binary.LittleEndian.Uint32(data[7:11])
// 标记此 csid 已收到有效头部
r.knownCSIDs[csid] = true
// 调试日志:记录 fmt=0 的完整信息
log.Printf("🔍 [RTMP Debug] fmt=0 csid=%d ts=%d msgLen=%d typeID=%d streamID=%d",
csid, header.Timestamp, header.MessageLength, header.MessageTypeID, header.MessageSID)
case CHUNK_FMT_1:
// 7 bytes: timestamp delta(3) + length(3) + typeID(1)
data := make([]byte, 7)
if _, err := io.ReadFull(r.reader, data); err != nil {
return nil, fmt.Errorf("读取 fmt1 header 失败: %w", err)
}
delta := uint32(data[0])<<16 | uint32(data[1])<<8 | uint32(data[2])
header.Timestamp = prevHeader.Timestamp + delta
header.MessageLength = uint32(data[3])<<16 | uint32(data[4])<<8 | uint32(data[5])
header.MessageTypeID = data[6]
// 标记此 csid 已收到有效头部
r.knownCSIDs[csid] = true
case CHUNK_FMT_2:
// 3 bytes: timestamp delta(3)
data := make([]byte, 3)
if _, err := io.ReadFull(r.reader, data); err != nil {
return nil, fmt.Errorf("读取 fmt2 header 失败: %w", err)
}
delta := uint32(data[0])<<16 | uint32(data[1])<<8 | uint32(data[2])
header.Timestamp = prevHeader.Timestamp + delta
// 关键修复:如果没有有效的 prevHeader尝试根据 csid 推断消息类型
// 微信小程序 live-pusher 可能首次发送音视频时就使用 fmt=2
if prevHeader.MessageLength == 0 && !r.knownCSIDs[csid] {
// 尝试推断消息类型(基于常见的 RTMP 实现)
// csid >= 4 通常用于音视频数据
if csid >= 4 {
// 设置默认消息长度为 chunkSize会在后续 chunk 中累积)
// 这是一个启发式处理,让数据能够被读取
r.mu.RLock()
header.MessageLength = r.chunkSize
r.mu.RUnlock()
// 猜测消息类型:偶数 csid 可能是音频,奇数是视频(这是一个启发式)
// 实际上我们需要从数据本身来判断
log.Printf("⚠️ [RTMP Protocol] fmt=2 csid=%d 无 prevHeader设置临时 msgLen=%d",
csid, header.MessageLength)
} else {
log.Printf("⚠️ [RTMP Protocol] fmt=2 on csid %d without valid prevHeader", csid)
}
}
case CHUNK_FMT_3:
// 0 bytes: 使用上一个头部的所有字段(已经复制了 prevHeader 的值)
// 关键修复:检查消息续传或尝试推断
if prevHeader.MessageLength == 0 && !r.knownCSIDs[csid] {
// 检查是否是消息续传
if state, exists := r.messageBuffer[csid]; exists && state.bytesRead < state.header.MessageLength {
// 是消息续传,使用已保存的消息头信息
header.MessageLength = state.header.MessageLength
header.MessageTypeID = state.header.MessageTypeID
header.MessageSID = state.header.MessageSID
header.Timestamp = state.header.Timestamp
log.Printf(" [RTMP Protocol] fmt=3 续传 csid=%d msgLen=%d", csid, header.MessageLength)
} else if csid >= 4 {
// 对于可能的音视频 csid设置默认消息长度
r.mu.RLock()
header.MessageLength = r.chunkSize
r.mu.RUnlock()
log.Printf("⚠️ [RTMP Protocol] fmt=3 csid=%d 无 prevHeader设置临时 msgLen=%d",
csid, header.MessageLength)
} else {
log.Printf("⚠️ [RTMP Protocol] fmt=3 on csid %d without valid prevHeader (首次使用), msgLen=0", csid)
}
}
}
// 检查是否有扩展时间戳
if header.Timestamp == 0xFFFFFF {
header.ExtendedTS = true
extTS := make([]byte, 4)
if _, err := io.ReadFull(r.reader, extTS); err != nil {
return nil, fmt.Errorf("读取扩展时间戳失败: %w", err)
}
header.Timestamp = binary.BigEndian.Uint32(extTS)
}
// 保存当前头部供后续 chunk 使用
r.prevHeaders[csid] = header
return header, nil
}
// WriteMessage 写入一个完整的 RTMP 消息
func (w *ChunkWriter) WriteMessage(msg *RTMPMessage) error {
w.mu.Lock()
defer w.mu.Unlock()
data := msg.Data
dataLen := uint32(len(data))
offset := uint32(0)
firstChunk := true
for offset < dataLen {
// 计算本次写入的字节数
remaining := dataLen - offset
toWrite := w.chunkSize
if remaining < toWrite {
toWrite = remaining
}
// 写入 chunk header
if firstChunk {
// fmt 0: 完整头部
if err := w.writeChunkHeader0(msg, dataLen); err != nil {
return err
}
firstChunk = false
} else {
// fmt 3: 无头部
if err := w.writeChunkHeader3(msg.ChunkStreamID); err != nil {
return err
}
}
// 写入数据
if _, err := w.writer.Write(data[offset : offset+toWrite]); err != nil {
return err
}
offset += toWrite
}
return w.writer.Flush()
}
// writeChunkHeader0 写入 fmt 0 头部
func (w *ChunkWriter) writeChunkHeader0(msg *RTMPMessage, dataLen uint32) error {
// Basic Header
csid := msg.ChunkStreamID
if csid < 64 {
if err := w.writer.WriteByte(byte(csid)); err != nil {
return err
}
} else if csid < 320 {
if _, err := w.writer.Write([]byte{0, byte(csid - 64)}); err != nil {
return err
}
} else {
csid -= 64
if _, err := w.writer.Write([]byte{1, byte(csid & 0xFF), byte(csid >> 8)}); err != nil {
return err
}
}
// Message Header (11 bytes)
header := make([]byte, 11)
ts := msg.Timestamp
if ts >= 0xFFFFFF {
ts = 0xFFFFFF
}
header[0] = byte(ts >> 16)
header[1] = byte(ts >> 8)
header[2] = byte(ts)
header[3] = byte(dataLen >> 16)
header[4] = byte(dataLen >> 8)
header[5] = byte(dataLen)
header[6] = msg.TypeID
binary.LittleEndian.PutUint32(header[7:11], msg.StreamID)
if _, err := w.writer.Write(header); err != nil {
return err
}
// Extended Timestamp
if msg.Timestamp >= 0xFFFFFF {
extTS := make([]byte, 4)
binary.BigEndian.PutUint32(extTS, msg.Timestamp)
if _, err := w.writer.Write(extTS); err != nil {
return err
}
}
return nil
}
// writeChunkHeader3 写入 fmt 3 头部
func (w *ChunkWriter) writeChunkHeader3(csid uint32) error {
// Basic Header with fmt = 3
if csid < 64 {
return w.writer.WriteByte(byte(0xC0 | csid))
} else if csid < 320 {
_, err := w.writer.Write([]byte{0xC0, byte(csid - 64)})
return err
} else {
csid -= 64
_, err := w.writer.Write([]byte{0xC1, byte(csid & 0xFF), byte(csid >> 8)})
return err
}
}
// WriteSetChunkSize 发送 SetChunkSize 消息
func (w *ChunkWriter) WriteSetChunkSize(size uint32) error {
data := make([]byte, 4)
binary.BigEndian.PutUint32(data, size)
msg := &RTMPMessage{
ChunkStreamID: 2,
Timestamp: 0,
TypeID: RTMP_MSG_CHUNK_SIZE,
StreamID: 0,
Data: data,
}
if err := w.WriteMessage(msg); err != nil {
return err
}
w.SetChunkSize(size)
return nil
}
// WriteWindowAckSize 发送 Window Acknowledgement Size 消息
func (w *ChunkWriter) WriteWindowAckSize(size uint32) error {
data := make([]byte, 4)
binary.BigEndian.PutUint32(data, size)
return w.WriteMessage(&RTMPMessage{
ChunkStreamID: 2,
Timestamp: 0,
TypeID: RTMP_MSG_WIN_ACK_SIZE,
StreamID: 0,
Data: data,
})
}
// WriteSetPeerBandwidth 发送 Set Peer Bandwidth 消息
func (w *ChunkWriter) WriteSetPeerBandwidth(size uint32, limitType uint8) error {
data := make([]byte, 5)
binary.BigEndian.PutUint32(data, size)
data[4] = limitType
return w.WriteMessage(&RTMPMessage{
ChunkStreamID: 2,
Timestamp: 0,
TypeID: RTMP_MSG_SET_PEER_BW,
StreamID: 0,
Data: data,
})
}
// WriteUserControl 发送 User Control Message
func (w *ChunkWriter) WriteUserControl(eventType uint16, data []byte) error {
payload := make([]byte, 2+len(data))
binary.BigEndian.PutUint16(payload, eventType)
copy(payload[2:], data)
return w.WriteMessage(&RTMPMessage{
ChunkStreamID: 2,
Timestamp: 0,
TypeID: RTMP_MSG_USER_CONTROL,
StreamID: 0,
Data: payload,
})
}
// WriteStreamBegin 发送 Stream Begin 事件
func (w *ChunkWriter) WriteStreamBegin(streamID uint32) error {
data := make([]byte, 4)
binary.BigEndian.PutUint32(data, streamID)
return w.WriteUserControl(0, data) // 0 = StreamBegin
}
// WriteCommand 发送 AMF0 命令消息
func (w *ChunkWriter) WriteCommand(csid uint32, streamID uint32, data []byte) error {
return w.WriteMessage(&RTMPMessage{
ChunkStreamID: csid,
Timestamp: 0,
TypeID: RTMP_MSG_AMF0_CMD,
StreamID: streamID,
Data: data,
})
}