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

540 lines
15 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
*
* 自实现的 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=3 没有有效的 prevHeader
// 记录详细日志以便调试,然后跳过此 chunk 继续尝试读取下一个
if toRead == 0 {
log.Printf("⚠️ [RTMP Protocol] toRead=0 on csid=%d fmt=%d msgLen=%d bytesRead=%d, 跳过此 chunk",
header.ChunkStreamID, header.Format, state.header.MessageLength, state.bytesRead)
// 清除此 csid 的无效状态,等待有效的 fmt=0 重新开始
delete(r.messageBuffer, header.ChunkStreamID)
continue // 继续尝试读取下一个 chunk而不是返回错误
}
// 读取 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,
}
// 清除消息缓冲
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 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
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记录警告但不立即失败
// 依赖后续 ReadMessage 中的 toRead==0 检查来捕获真正的无效数据
if prevHeader.MessageLength == 0 && !r.knownCSIDs[csid] {
log.Printf("⚠️ [RTMP Protocol] fmt=2 on csid %d without valid prevHeader, may cause issues", csid)
}
case CHUNK_FMT_3:
// 0 bytes: 使用上一个头部的所有字段(已经复制了 prevHeader 的值)
// 放宽验证:如果没有有效的 prevHeader记录警告但不立即失败
// 某些 RTMP 客户端(如小程序 live-pusher可能在首次使用某 csid 时就用 fmt=3
// 检查是否在 messageBuffer 中有未完成的消息(消息续传场景)
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 {
log.Printf("⚠️ [RTMP Protocol] fmt=3 on csid %d without valid prevHeader (首次使用), msgLen=0", csid)
// 不返回错误,让 ReadMessage 中的 toRead==0 检查来处理
}
}
}
// 检查是否有扩展时间戳
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,
})
}