微信小程序\推拉流
This commit is contained in:
@@ -11,6 +11,7 @@ package api
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"log"
|
||||
|
||||
"xk-websocket-v2/internal/mediaserver"
|
||||
"xk-websocket-v2/internal/utils"
|
||||
@@ -59,10 +60,13 @@ type WebRTCICERequest struct {
|
||||
type JoinCallRoomResponse struct {
|
||||
RoomID string `json:"room_id"`
|
||||
Platform string `json:"platform"`
|
||||
ICEServers []ICEServerConfig `json:"ice_servers,omitempty"` // H5/App 用
|
||||
ICEServers []ICEServerConfig `json:"ice_servers,omitempty"` // H5/App 用 (WebRTC 模式)
|
||||
WSPushURL string `json:"ws_push_url,omitempty"` // H5/App 用 (RTMP 模式 WebSocket 推流地址)
|
||||
SelfFLVURL string `json:"self_flv_url,omitempty"` // H5/App 用 - 自己的 HTTP-FLV 地址(供小程序拉流)
|
||||
FlvPullURLs []PullURLInfo `json:"flv_pull_urls,omitempty"` // H5/App 拉取小程序流的 FLV 地址
|
||||
PushURL string `json:"push_url,omitempty"` // 小程序用
|
||||
PullURLs []PullURLInfo `json:"pull_urls,omitempty"` // 小程序用
|
||||
PushURL string `json:"push_url,omitempty"` // 小程序用 RTMP 推流地址
|
||||
FLVURL string `json:"flv_url,omitempty"` // 小程序用 - 自己的 HTTP-FLV 地址(供其他端拉流)
|
||||
PullURLs []PullURLInfo `json:"pull_urls,omitempty"` // 小程序用 RTMP 拉流地址
|
||||
Participants []ParticipantInfo `json:"participants"`
|
||||
}
|
||||
|
||||
@@ -154,17 +158,40 @@ func JoinCallRoomHandler(c *gin.Context) {
|
||||
|
||||
switch req.Platform {
|
||||
case "h5", "app":
|
||||
// H5/App 使用 WebRTC
|
||||
// H5/App 使用 WebRTC 或 RTMP 模式
|
||||
pType := mediaserver.ParticipantTypeWebRTC
|
||||
_, err := room.AddParticipant(req.UserID, pType)
|
||||
participant, err := room.AddParticipant(req.UserID, pType)
|
||||
if err != nil {
|
||||
utils.BadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
// 返回 ICE 服务器配置
|
||||
// 返回 ICE 服务器配置 (WebRTC 模式)
|
||||
response.ICEServers = getICEServers(req.UserID)
|
||||
|
||||
// 生成 WebSocket 推流地址 (RTMP 模式)
|
||||
// 格式: ws://host:port/api/call/ws-push?stream_id=xxx&user_id=xxx&room_id=xxx&token=xxx
|
||||
rtmpServer := ms.GetRTMP()
|
||||
if rtmpServer != nil {
|
||||
// 为 H5/App 生成流信息
|
||||
stream, err := rtmpServer.GenerateStreamURLs(req.RoomID, req.UserID)
|
||||
if err == nil {
|
||||
// 更新参与者的流地址
|
||||
room.SetParticipantRTMPURLs(req.UserID, stream.PushURL, stream.PullURL, stream.FLVURL, stream.ID)
|
||||
participant.PushURL = stream.PushURL
|
||||
participant.PullURL = stream.PullURL
|
||||
participant.FLVURL = stream.FLVURL
|
||||
participant.StreamID = stream.ID
|
||||
|
||||
// 构建 WebSocket 推流 URL
|
||||
token := mediaserver.GenerateStreamToken(stream.ID)
|
||||
response.WSPushURL = fmt.Sprintf("wss://g-ws.nailaoyun.cn/api/call/ws-push?stream_id=%s&user_id=%s&room_id=%s&token=%s",
|
||||
stream.ID, req.UserID, req.RoomID, token)
|
||||
// 返回 H5/App 自己的 FLV 地址(供小程序拉流)
|
||||
response.SelfFLVURL = stream.FLVURL
|
||||
}
|
||||
}
|
||||
|
||||
// 获取房间内小程序用户的 FLV 拉流地址(用于 Web 端播放小程序流)
|
||||
flvPullURLs := make([]PullURLInfo, 0)
|
||||
for _, p := range room.GetOtherParticipants(req.UserID) {
|
||||
@@ -209,6 +236,7 @@ func JoinCallRoomHandler(c *gin.Context) {
|
||||
participant.StreamID = stream.ID
|
||||
|
||||
response.PushURL = stream.PushURL
|
||||
response.FLVURL = stream.FLVURL // 返回自己的 FLV 地址,供客户端发送给其他端
|
||||
|
||||
// 获取房间内其他用户的拉流地址
|
||||
pullURLs := make([]PullURLInfo, 0)
|
||||
@@ -432,5 +460,33 @@ func RegisterCallRoutes(router *gin.RouterGroup) {
|
||||
callGroup.POST("/offer", WebRTCOfferHandler)
|
||||
callGroup.POST("/ice", WebRTCICEHandler)
|
||||
callGroup.GET("/ice-servers", GetICEServersHandler)
|
||||
// 注意:ws-push 路由已移到全局(main.go),避免中间件干扰 WebSocket 升级
|
||||
}
|
||||
}
|
||||
|
||||
// WSPushHandler WebSocket 推流处理
|
||||
// GET /api/call/ws-push?stream_id=xxx&user_id=xxx&room_id=xxx&token=xxx
|
||||
func WSPushHandler(c *gin.Context) {
|
||||
// 详细日志:确认请求到达
|
||||
log.Printf("🔌 [WSPush] 收到请求: %s from %s", c.Request.URL.String(), c.ClientIP())
|
||||
log.Printf("🔌 [WSPush] Headers: Upgrade=%s Connection=%s Origin=%s",
|
||||
c.GetHeader("Upgrade"), c.GetHeader("Connection"), c.GetHeader("Origin"))
|
||||
|
||||
ms := mediaserver.GetServer()
|
||||
if ms.GetConfig() == nil || !ms.GetConfig().Enabled {
|
||||
log.Printf("❌ [WSPush] 媒体服务未启用")
|
||||
utils.Error(c, 503, "媒体服务未启用")
|
||||
return
|
||||
}
|
||||
|
||||
wsProxy := ms.GetWSProxy()
|
||||
if wsProxy == nil {
|
||||
log.Printf("❌ [WSPush] WebSocket代理未启用")
|
||||
utils.Error(c, 503, "WebSocket代理未启用")
|
||||
return
|
||||
}
|
||||
|
||||
log.Printf("✅ [WSPush] 转交给 WebSocket 代理处理")
|
||||
// 交给 WebSocket 代理处理
|
||||
wsProxy.HandleWebSocket(c.Writer, c.Request)
|
||||
}
|
||||
|
||||
@@ -107,6 +107,15 @@ func (r *Room) AddParticipant(userID string, pType ParticipantType) (*Participan
|
||||
|
||||
// 检查是否已存在
|
||||
if p, exists := r.Participants[userID]; exists {
|
||||
// 用户重新加入,清除旧的流信息以便重新生成
|
||||
oldStreamID := p.StreamID
|
||||
p.PushURL = ""
|
||||
p.PullURL = ""
|
||||
p.FLVURL = ""
|
||||
p.StreamID = ""
|
||||
p.Type = pType
|
||||
p.JoinedAt = time.Now()
|
||||
log.Printf("🔄 [Room:%s] 用户 %s 重新加入,已清除旧流信息 (旧StreamID: %s)", r.ID, userID, oldStreamID)
|
||||
return p, nil
|
||||
}
|
||||
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
427
internal/mediaserver/rtmp_amf0.go
Normal file
427
internal/mediaserver/rtmp_amf0.go
Normal file
@@ -0,0 +1,427 @@
|
||||
/**
|
||||
* package mediaserver
|
||||
*
|
||||
* AMF0 (Action Message Format 0) 编解码实现
|
||||
* 用于 RTMP 命令消息的序列化和反序列化
|
||||
*/
|
||||
package mediaserver
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/binary"
|
||||
"fmt"
|
||||
"math"
|
||||
)
|
||||
|
||||
// AMF0 数据类型标记
|
||||
const (
|
||||
AMF0_NUMBER = 0x00 // 8 bytes double
|
||||
AMF0_BOOLEAN = 0x01 // 1 byte
|
||||
AMF0_STRING = 0x02 // 2 bytes length + data
|
||||
AMF0_OBJECT = 0x03 // key-value pairs
|
||||
AMF0_MOVIECLIP = 0x04 // reserved
|
||||
AMF0_NULL = 0x05 // no data
|
||||
AMF0_UNDEFINED = 0x06 // no data
|
||||
AMF0_REFERENCE = 0x07 // 2 bytes
|
||||
AMF0_ECMA_ARRAY = 0x08 // associative array
|
||||
AMF0_OBJECT_END = 0x09 // object end marker
|
||||
AMF0_STRICT_ARRAY = 0x0A // strict array
|
||||
AMF0_DATE = 0x0B // 8 bytes double + 2 bytes timezone
|
||||
AMF0_LONG_STRING = 0x0C // 4 bytes length + data
|
||||
AMF0_UNSUPPORTED = 0x0D
|
||||
AMF0_RECORDSET = 0x0E // reserved
|
||||
AMF0_XML_DOCUMENT = 0x0F
|
||||
AMF0_TYPED_OBJECT = 0x10
|
||||
AMF0_AVMPLUS = 0x11 // switch to AMF3
|
||||
)
|
||||
|
||||
// AMF0Object 表示一个 AMF0 对象
|
||||
type AMF0Object map[string]interface{}
|
||||
|
||||
// AMF0Encoder AMF0 编码器
|
||||
type AMF0Encoder struct {
|
||||
buf *bytes.Buffer
|
||||
}
|
||||
|
||||
// AMF0Decoder AMF0 解码器
|
||||
type AMF0Decoder struct {
|
||||
data []byte
|
||||
offset int
|
||||
}
|
||||
|
||||
// NewAMF0Encoder 创建 AMF0 编码器
|
||||
func NewAMF0Encoder() *AMF0Encoder {
|
||||
return &AMF0Encoder{
|
||||
buf: new(bytes.Buffer),
|
||||
}
|
||||
}
|
||||
|
||||
// NewAMF0Decoder 创建 AMF0 解码器
|
||||
func NewAMF0Decoder(data []byte) *AMF0Decoder {
|
||||
return &AMF0Decoder{
|
||||
data: data,
|
||||
offset: 0,
|
||||
}
|
||||
}
|
||||
|
||||
// Bytes 获取编码后的字节
|
||||
func (e *AMF0Encoder) Bytes() []byte {
|
||||
return e.buf.Bytes()
|
||||
}
|
||||
|
||||
// Reset 重置编码器
|
||||
func (e *AMF0Encoder) Reset() {
|
||||
e.buf.Reset()
|
||||
}
|
||||
|
||||
// EncodeNumber 编码数字
|
||||
func (e *AMF0Encoder) EncodeNumber(val float64) {
|
||||
e.buf.WriteByte(AMF0_NUMBER)
|
||||
bits := math.Float64bits(val)
|
||||
binary.Write(e.buf, binary.BigEndian, bits)
|
||||
}
|
||||
|
||||
// EncodeBoolean 编码布尔值
|
||||
func (e *AMF0Encoder) EncodeBoolean(val bool) {
|
||||
e.buf.WriteByte(AMF0_BOOLEAN)
|
||||
if val {
|
||||
e.buf.WriteByte(1)
|
||||
} else {
|
||||
e.buf.WriteByte(0)
|
||||
}
|
||||
}
|
||||
|
||||
// EncodeString 编码字符串
|
||||
func (e *AMF0Encoder) EncodeString(val string) {
|
||||
data := []byte(val)
|
||||
if len(data) > 0xFFFF {
|
||||
// Long string
|
||||
e.buf.WriteByte(AMF0_LONG_STRING)
|
||||
binary.Write(e.buf, binary.BigEndian, uint32(len(data)))
|
||||
} else {
|
||||
e.buf.WriteByte(AMF0_STRING)
|
||||
binary.Write(e.buf, binary.BigEndian, uint16(len(data)))
|
||||
}
|
||||
e.buf.Write(data)
|
||||
}
|
||||
|
||||
// EncodeNull 编码 null
|
||||
func (e *AMF0Encoder) EncodeNull() {
|
||||
e.buf.WriteByte(AMF0_NULL)
|
||||
}
|
||||
|
||||
// EncodeObject 编码对象
|
||||
func (e *AMF0Encoder) EncodeObject(obj AMF0Object) {
|
||||
e.buf.WriteByte(AMF0_OBJECT)
|
||||
for key, val := range obj {
|
||||
// 写入属性名(不带类型标记)
|
||||
binary.Write(e.buf, binary.BigEndian, uint16(len(key)))
|
||||
e.buf.WriteString(key)
|
||||
// 写入值
|
||||
e.EncodeValue(val)
|
||||
}
|
||||
// 写入对象结束标记
|
||||
e.buf.Write([]byte{0, 0, AMF0_OBJECT_END})
|
||||
}
|
||||
|
||||
// EncodeEcmaArray 编码 ECMA 数组
|
||||
func (e *AMF0Encoder) EncodeEcmaArray(obj AMF0Object) {
|
||||
e.buf.WriteByte(AMF0_ECMA_ARRAY)
|
||||
binary.Write(e.buf, binary.BigEndian, uint32(len(obj)))
|
||||
for key, val := range obj {
|
||||
binary.Write(e.buf, binary.BigEndian, uint16(len(key)))
|
||||
e.buf.WriteString(key)
|
||||
e.EncodeValue(val)
|
||||
}
|
||||
e.buf.Write([]byte{0, 0, AMF0_OBJECT_END})
|
||||
}
|
||||
|
||||
// EncodeValue 编码任意值
|
||||
func (e *AMF0Encoder) EncodeValue(val interface{}) {
|
||||
switch v := val.(type) {
|
||||
case float64:
|
||||
e.EncodeNumber(v)
|
||||
case float32:
|
||||
e.EncodeNumber(float64(v))
|
||||
case int:
|
||||
e.EncodeNumber(float64(v))
|
||||
case int64:
|
||||
e.EncodeNumber(float64(v))
|
||||
case int32:
|
||||
e.EncodeNumber(float64(v))
|
||||
case uint32:
|
||||
e.EncodeNumber(float64(v))
|
||||
case bool:
|
||||
e.EncodeBoolean(v)
|
||||
case string:
|
||||
e.EncodeString(v)
|
||||
case nil:
|
||||
e.EncodeNull()
|
||||
case AMF0Object:
|
||||
e.EncodeObject(v)
|
||||
case map[string]interface{}:
|
||||
e.EncodeObject(AMF0Object(v))
|
||||
default:
|
||||
e.EncodeNull()
|
||||
}
|
||||
}
|
||||
|
||||
// Remaining 返回剩余未解析的字节数
|
||||
func (d *AMF0Decoder) Remaining() int {
|
||||
return len(d.data) - d.offset
|
||||
}
|
||||
|
||||
// DecodeAll 解码所有值
|
||||
func (d *AMF0Decoder) DecodeAll() ([]interface{}, error) {
|
||||
var values []interface{}
|
||||
for d.Remaining() > 0 {
|
||||
val, err := d.DecodeValue()
|
||||
if err != nil {
|
||||
break
|
||||
}
|
||||
values = append(values, val)
|
||||
}
|
||||
return values, nil
|
||||
}
|
||||
|
||||
// DecodeValue 解码一个值
|
||||
func (d *AMF0Decoder) DecodeValue() (interface{}, error) {
|
||||
if d.Remaining() < 1 {
|
||||
return nil, fmt.Errorf("数据不足")
|
||||
}
|
||||
|
||||
marker := d.data[d.offset]
|
||||
d.offset++
|
||||
|
||||
switch marker {
|
||||
case AMF0_NUMBER:
|
||||
return d.decodeNumber()
|
||||
case AMF0_BOOLEAN:
|
||||
return d.decodeBoolean()
|
||||
case AMF0_STRING:
|
||||
return d.decodeString()
|
||||
case AMF0_OBJECT:
|
||||
return d.decodeObject()
|
||||
case AMF0_NULL, AMF0_UNDEFINED:
|
||||
return nil, nil
|
||||
case AMF0_ECMA_ARRAY:
|
||||
return d.decodeEcmaArray()
|
||||
case AMF0_STRICT_ARRAY:
|
||||
return d.decodeStrictArray()
|
||||
case AMF0_LONG_STRING:
|
||||
return d.decodeLongString()
|
||||
case AMF0_DATE:
|
||||
return d.decodeDate()
|
||||
default:
|
||||
return nil, fmt.Errorf("未知的 AMF0 类型: 0x%02X", marker)
|
||||
}
|
||||
}
|
||||
|
||||
// decodeNumber 解码数字
|
||||
func (d *AMF0Decoder) decodeNumber() (float64, error) {
|
||||
if d.Remaining() < 8 {
|
||||
return 0, fmt.Errorf("数据不足以解码 number")
|
||||
}
|
||||
bits := binary.BigEndian.Uint64(d.data[d.offset : d.offset+8])
|
||||
d.offset += 8
|
||||
return math.Float64frombits(bits), nil
|
||||
}
|
||||
|
||||
// decodeBoolean 解码布尔值
|
||||
func (d *AMF0Decoder) decodeBoolean() (bool, error) {
|
||||
if d.Remaining() < 1 {
|
||||
return false, fmt.Errorf("数据不足以解码 boolean")
|
||||
}
|
||||
val := d.data[d.offset] != 0
|
||||
d.offset++
|
||||
return val, nil
|
||||
}
|
||||
|
||||
// decodeString 解码字符串
|
||||
func (d *AMF0Decoder) decodeString() (string, error) {
|
||||
if d.Remaining() < 2 {
|
||||
return "", fmt.Errorf("数据不足以解码 string length")
|
||||
}
|
||||
length := int(binary.BigEndian.Uint16(d.data[d.offset : d.offset+2]))
|
||||
d.offset += 2
|
||||
|
||||
if d.Remaining() < length {
|
||||
return "", fmt.Errorf("数据不足以解码 string data")
|
||||
}
|
||||
str := string(d.data[d.offset : d.offset+length])
|
||||
d.offset += length
|
||||
return str, nil
|
||||
}
|
||||
|
||||
// decodeLongString 解码长字符串
|
||||
func (d *AMF0Decoder) decodeLongString() (string, error) {
|
||||
if d.Remaining() < 4 {
|
||||
return "", fmt.Errorf("数据不足以解码 long string length")
|
||||
}
|
||||
length := int(binary.BigEndian.Uint32(d.data[d.offset : d.offset+4]))
|
||||
d.offset += 4
|
||||
|
||||
if d.Remaining() < length {
|
||||
return "", fmt.Errorf("数据不足以解码 long string data")
|
||||
}
|
||||
str := string(d.data[d.offset : d.offset+length])
|
||||
d.offset += length
|
||||
return str, nil
|
||||
}
|
||||
|
||||
// decodeObject 解码对象
|
||||
func (d *AMF0Decoder) decodeObject() (AMF0Object, error) {
|
||||
obj := make(AMF0Object)
|
||||
|
||||
for {
|
||||
// 读取属性名长度
|
||||
if d.Remaining() < 2 {
|
||||
return nil, fmt.Errorf("数据不足以解码对象属性名长度")
|
||||
}
|
||||
nameLen := int(binary.BigEndian.Uint16(d.data[d.offset : d.offset+2]))
|
||||
d.offset += 2
|
||||
|
||||
// 检查是否到达对象结束
|
||||
if nameLen == 0 {
|
||||
if d.Remaining() < 1 || d.data[d.offset] != AMF0_OBJECT_END {
|
||||
return nil, fmt.Errorf("对象结束标记缺失")
|
||||
}
|
||||
d.offset++
|
||||
break
|
||||
}
|
||||
|
||||
// 读取属性名
|
||||
if d.Remaining() < nameLen {
|
||||
return nil, fmt.Errorf("数据不足以解码对象属性名")
|
||||
}
|
||||
name := string(d.data[d.offset : d.offset+nameLen])
|
||||
d.offset += nameLen
|
||||
|
||||
// 读取属性值
|
||||
val, err := d.DecodeValue()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("解码对象属性值失败: %w", err)
|
||||
}
|
||||
obj[name] = val
|
||||
}
|
||||
|
||||
return obj, nil
|
||||
}
|
||||
|
||||
// decodeEcmaArray 解码 ECMA 数组
|
||||
func (d *AMF0Decoder) decodeEcmaArray() (AMF0Object, error) {
|
||||
if d.Remaining() < 4 {
|
||||
return nil, fmt.Errorf("数据不足以解码 ECMA 数组长度")
|
||||
}
|
||||
// 读取数组长度(但实际不使用,因为以 object end 结束)
|
||||
d.offset += 4
|
||||
|
||||
return d.decodeObject()
|
||||
}
|
||||
|
||||
// decodeStrictArray 解码严格数组
|
||||
func (d *AMF0Decoder) decodeStrictArray() ([]interface{}, error) {
|
||||
if d.Remaining() < 4 {
|
||||
return nil, fmt.Errorf("数据不足以解码严格数组长度")
|
||||
}
|
||||
length := int(binary.BigEndian.Uint32(d.data[d.offset : d.offset+4]))
|
||||
d.offset += 4
|
||||
|
||||
arr := make([]interface{}, length)
|
||||
for i := 0; i < length; i++ {
|
||||
val, err := d.DecodeValue()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("解码数组元素失败: %w", err)
|
||||
}
|
||||
arr[i] = val
|
||||
}
|
||||
return arr, nil
|
||||
}
|
||||
|
||||
// decodeDate 解码日期
|
||||
func (d *AMF0Decoder) decodeDate() (float64, error) {
|
||||
if d.Remaining() < 10 {
|
||||
return 0, fmt.Errorf("数据不足以解码 date")
|
||||
}
|
||||
bits := binary.BigEndian.Uint64(d.data[d.offset : d.offset+8])
|
||||
d.offset += 8
|
||||
// 跳过时区信息 (2 bytes)
|
||||
d.offset += 2
|
||||
return math.Float64frombits(bits), nil
|
||||
}
|
||||
|
||||
// EncodeAMF0 编码多个值为 AMF0 格式
|
||||
func EncodeAMF0(values ...interface{}) []byte {
|
||||
encoder := NewAMF0Encoder()
|
||||
for _, val := range values {
|
||||
encoder.EncodeValue(val)
|
||||
}
|
||||
return encoder.Bytes()
|
||||
}
|
||||
|
||||
// DecodeAMF0 从 AMF0 格式解码多个值
|
||||
func DecodeAMF0(data []byte) ([]interface{}, error) {
|
||||
decoder := NewAMF0Decoder(data)
|
||||
return decoder.DecodeAll()
|
||||
}
|
||||
|
||||
// EncodeConnectResult 编码 connect 响应
|
||||
func EncodeConnectResult(transactionID float64) []byte {
|
||||
encoder := NewAMF0Encoder()
|
||||
|
||||
// _result
|
||||
encoder.EncodeString("_result")
|
||||
encoder.EncodeNumber(transactionID)
|
||||
|
||||
// Properties
|
||||
encoder.EncodeObject(AMF0Object{
|
||||
"fmsVer": "FMS/3,5,7,7009",
|
||||
"capabilities": float64(31),
|
||||
"mode": float64(1),
|
||||
})
|
||||
|
||||
// Information
|
||||
encoder.EncodeObject(AMF0Object{
|
||||
"level": "status",
|
||||
"code": "NetConnection.Connect.Success",
|
||||
"description": "Connection succeeded.",
|
||||
"objectEncoding": float64(0),
|
||||
})
|
||||
|
||||
return encoder.Bytes()
|
||||
}
|
||||
|
||||
// EncodeCreateStreamResult 编码 createStream 响应
|
||||
func EncodeCreateStreamResult(transactionID float64, streamID float64) []byte {
|
||||
encoder := NewAMF0Encoder()
|
||||
encoder.EncodeString("_result")
|
||||
encoder.EncodeNumber(transactionID)
|
||||
encoder.EncodeNull()
|
||||
encoder.EncodeNumber(streamID)
|
||||
return encoder.Bytes()
|
||||
}
|
||||
|
||||
// EncodeOnStatus 编码 onStatus 消息
|
||||
func EncodeOnStatus(code, level, description string) []byte {
|
||||
encoder := NewAMF0Encoder()
|
||||
encoder.EncodeString("onStatus")
|
||||
encoder.EncodeNumber(0)
|
||||
encoder.EncodeNull()
|
||||
encoder.EncodeObject(AMF0Object{
|
||||
"level": level,
|
||||
"code": code,
|
||||
"description": description,
|
||||
})
|
||||
return encoder.Bytes()
|
||||
}
|
||||
|
||||
// EncodeOnBWDone 编码 onBWDone 消息
|
||||
func EncodeOnBWDone() []byte {
|
||||
encoder := NewAMF0Encoder()
|
||||
encoder.EncodeString("onBWDone")
|
||||
encoder.EncodeNumber(0)
|
||||
encoder.EncodeNull()
|
||||
return encoder.Bytes()
|
||||
}
|
||||
|
||||
|
||||
161
internal/mediaserver/rtmp_handshake.go
Normal file
161
internal/mediaserver/rtmp_handshake.go
Normal file
@@ -0,0 +1,161 @@
|
||||
/**
|
||||
* package mediaserver
|
||||
*
|
||||
* RTMP 握手实现
|
||||
* 支持 Simple Handshake (Version 3)
|
||||
*/
|
||||
package mediaserver
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"crypto/rand"
|
||||
"encoding/binary"
|
||||
"fmt"
|
||||
"io"
|
||||
"log"
|
||||
"net"
|
||||
"time"
|
||||
)
|
||||
|
||||
// 握手相关常量
|
||||
const (
|
||||
HANDSHAKE_SIZE = 1536
|
||||
RTMP_VERSION = 3
|
||||
)
|
||||
|
||||
// DoHandshake 执行 RTMP 握手(服务器端)
|
||||
// 握手流程:
|
||||
// 1. 接收 C0 + C1 (1 + 1536 bytes)
|
||||
// 2. 发送 S0 + S1 + S2 (1 + 1536 + 1536 bytes)
|
||||
// 3. 接收 C2 (1536 bytes)
|
||||
func DoHandshake(conn net.Conn, timeout time.Duration) error {
|
||||
// 设置超时
|
||||
conn.SetDeadline(time.Now().Add(timeout))
|
||||
defer conn.SetDeadline(time.Time{}) // 清除超时
|
||||
|
||||
// 1. 接收 C0 (1 byte: version)
|
||||
c0 := make([]byte, 1)
|
||||
if _, err := io.ReadFull(conn, c0); err != nil {
|
||||
return fmt.Errorf("读取 C0 失败: %w", err)
|
||||
}
|
||||
|
||||
version := c0[0]
|
||||
if version != RTMP_VERSION {
|
||||
// 尝试兼容其他版本
|
||||
log.Printf("⚠️ [RTMP Handshake] 客户端版本: %d (期望: %d)", version, RTMP_VERSION)
|
||||
}
|
||||
|
||||
// 2. 接收 C1 (1536 bytes)
|
||||
c1 := make([]byte, HANDSHAKE_SIZE)
|
||||
if _, err := io.ReadFull(conn, c1); err != nil {
|
||||
return fmt.Errorf("读取 C1 失败: %w", err)
|
||||
}
|
||||
|
||||
// C1 结构:
|
||||
// - time (4 bytes): 时间戳
|
||||
// - zero (4 bytes): 必须为0(简单握手)或版本信息(复杂握手)
|
||||
// - random (1528 bytes): 随机数据
|
||||
c1Time := binary.BigEndian.Uint32(c1[0:4])
|
||||
c1Zero := binary.BigEndian.Uint32(c1[4:8])
|
||||
|
||||
log.Printf("📡 [RTMP Handshake] C1: time=%d zero=%d", c1Time, c1Zero)
|
||||
|
||||
// 3. 生成 S0 + S1 + S2
|
||||
s0 := []byte{RTMP_VERSION}
|
||||
|
||||
// S1 (1536 bytes): time(4) + zero(4) + random(1528)
|
||||
s1 := make([]byte, HANDSHAKE_SIZE)
|
||||
binary.BigEndian.PutUint32(s1[0:4], uint32(time.Now().Unix())) // time
|
||||
binary.BigEndian.PutUint32(s1[4:8], 0) // zero
|
||||
rand.Read(s1[8:]) // random
|
||||
|
||||
// S2 (1536 bytes): 回显 C1(简单握手)
|
||||
// - time (4 bytes): C1 的时间戳
|
||||
// - time2 (4 bytes): S1 的时间戳
|
||||
// - random echo (1528 bytes): C1 的随机数据
|
||||
s2 := make([]byte, HANDSHAKE_SIZE)
|
||||
copy(s2[0:4], c1[0:4]) // 回显 C1 时间戳
|
||||
binary.BigEndian.PutUint32(s2[4:8], binary.BigEndian.Uint32(s1[0:4])) // S1 时间戳
|
||||
copy(s2[8:], c1[8:]) // 回显 C1 随机数据
|
||||
|
||||
// 4. 发送 S0 + S1 + S2
|
||||
response := make([]byte, 0, 1+HANDSHAKE_SIZE*2)
|
||||
response = append(response, s0...)
|
||||
response = append(response, s1...)
|
||||
response = append(response, s2...)
|
||||
|
||||
if _, err := conn.Write(response); err != nil {
|
||||
return fmt.Errorf("发送 S0+S1+S2 失败: %w", err)
|
||||
}
|
||||
|
||||
// 5. 接收 C2 (1536 bytes)
|
||||
c2 := make([]byte, HANDSHAKE_SIZE)
|
||||
if _, err := io.ReadFull(conn, c2); err != nil {
|
||||
return fmt.Errorf("读取 C2 失败: %w", err)
|
||||
}
|
||||
|
||||
// 验证 C2 (可选,简单握手可以跳过)
|
||||
// C2 应该回显 S1 的数据
|
||||
c2Time := binary.BigEndian.Uint32(c2[0:4])
|
||||
if c2Time != binary.BigEndian.Uint32(s1[0:4]) {
|
||||
log.Printf("⚠️ [RTMP Handshake] C2 时间戳不匹配 (收到: %d, 期望: %d)", c2Time, binary.BigEndian.Uint32(s1[0:4]))
|
||||
// 不返回错误,继续处理
|
||||
}
|
||||
|
||||
// 简单验证:比较随机数据的前几个字节
|
||||
if !bytes.Equal(c2[8:16], s1[8:16]) {
|
||||
log.Printf("⚠️ [RTMP Handshake] C2 随机数据不匹配")
|
||||
// 不返回错误,继续处理
|
||||
}
|
||||
|
||||
log.Printf("✅ [RTMP Handshake] 握手完成 | remote=%s", conn.RemoteAddr())
|
||||
return nil
|
||||
}
|
||||
|
||||
// DoClientHandshake 执行 RTMP 握手(客户端)
|
||||
// 用于测试或代理场景
|
||||
func DoClientHandshake(conn net.Conn, timeout time.Duration) error {
|
||||
// 设置超时
|
||||
conn.SetDeadline(time.Now().Add(timeout))
|
||||
defer conn.SetDeadline(time.Time{})
|
||||
|
||||
// 1. 发送 C0 + C1
|
||||
c0 := []byte{RTMP_VERSION}
|
||||
|
||||
c1 := make([]byte, HANDSHAKE_SIZE)
|
||||
binary.BigEndian.PutUint32(c1[0:4], uint32(time.Now().Unix())) // time
|
||||
binary.BigEndian.PutUint32(c1[4:8], 0) // zero
|
||||
rand.Read(c1[8:]) // random
|
||||
|
||||
if _, err := conn.Write(append(c0, c1...)); err != nil {
|
||||
return fmt.Errorf("发送 C0+C1 失败: %w", err)
|
||||
}
|
||||
|
||||
// 2. 接收 S0 + S1 + S2
|
||||
s0s1s2 := make([]byte, 1+HANDSHAKE_SIZE*2)
|
||||
if _, err := io.ReadFull(conn, s0s1s2); err != nil {
|
||||
return fmt.Errorf("读取 S0+S1+S2 失败: %w", err)
|
||||
}
|
||||
|
||||
// 验证服务器版本
|
||||
serverVersion := s0s1s2[0]
|
||||
if serverVersion != RTMP_VERSION {
|
||||
log.Printf("⚠️ [RTMP Handshake] 服务器版本: %d", serverVersion)
|
||||
}
|
||||
|
||||
s1 := s0s1s2[1 : 1+HANDSHAKE_SIZE]
|
||||
|
||||
// 3. 发送 C2 (回显 S1)
|
||||
c2 := make([]byte, HANDSHAKE_SIZE)
|
||||
copy(c2[0:4], s1[0:4]) // 回显 S1 时间戳
|
||||
binary.BigEndian.PutUint32(c2[4:8], binary.BigEndian.Uint32(c1[0:4])) // C1 时间戳
|
||||
copy(c2[8:], s1[8:]) // 回显 S1 随机数据
|
||||
|
||||
if _, err := conn.Write(c2); err != nil {
|
||||
return fmt.Errorf("发送 C2 失败: %w", err)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
|
||||
509
internal/mediaserver/rtmp_protocol.go
Normal file
509
internal/mediaserver/rtmp_protocol.go
Normal file
@@ -0,0 +1,509 @@
|
||||
/**
|
||||
* 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
|
||||
}
|
||||
|
||||
// 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),
|
||||
}
|
||||
}
|
||||
|
||||
// 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 {
|
||||
// 新消息或上一条消息已完成,创建新状态
|
||||
state = &messageState{
|
||||
header: header,
|
||||
data: make([]byte, 0, header.MessageLength),
|
||||
bytesRead: 0,
|
||||
}
|
||||
r.messageBuffer[header.ChunkStreamID] = state
|
||||
}
|
||||
|
||||
// 计算本次要读取的字节数
|
||||
// 关键修复:使用 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,这是无效的 RTMP 消息
|
||||
// 直接返回错误而不是继续,避免失去同步
|
||||
if toRead == 0 {
|
||||
return nil, fmt.Errorf("invalid RTMP: toRead=0 on csid=%d msgLen=%d bytesRead=%d",
|
||||
header.ChunkStreamID, state.header.MessageLength, state.bytesRead)
|
||||
}
|
||||
|
||||
// 读取 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])
|
||||
|
||||
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]
|
||||
|
||||
case CHUNK_FMT_2:
|
||||
// 验证:fmt 2 必须有有效的 prevHeader(MessageLength 从 prevHeader 继承)
|
||||
if prevHeader.MessageLength == 0 {
|
||||
return nil, fmt.Errorf("invalid RTMP: fmt 2 on csid %d without valid prevHeader (msgLen=0)", csid)
|
||||
}
|
||||
// 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
|
||||
|
||||
case CHUNK_FMT_3:
|
||||
// 验证:fmt 3 必须有有效的 prevHeader(所有字段从 prevHeader 继承)
|
||||
if prevHeader.MessageLength == 0 {
|
||||
return nil, fmt.Errorf("invalid RTMP: fmt 3 on csid %d without valid prevHeader (msgLen=0)", csid)
|
||||
}
|
||||
// 0 bytes: 使用上一个头部的所有字段
|
||||
// 已经复制了 prevHeader 的值
|
||||
}
|
||||
|
||||
// 检查是否有扩展时间戳
|
||||
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,
|
||||
})
|
||||
}
|
||||
|
||||
@@ -18,11 +18,12 @@ import (
|
||||
|
||||
// MediaServer 媒体服务器主结构
|
||||
type MediaServer struct {
|
||||
sfu *SFUServer // WebRTC SFU 服务
|
||||
rtmp *RTMPServer // RTMP 服务
|
||||
rooms map[string]*Room // 通话房间管理
|
||||
mu sync.RWMutex // 房间锁
|
||||
config *MediaServerConfig // 配置
|
||||
sfu *SFUServer // WebRTC SFU 服务
|
||||
rtmp *RTMPServer // RTMP 服务
|
||||
wsProxy *WebMToRTMPProxy // WebSocket to RTMP 代理
|
||||
rooms map[string]*Room // 通话房间管理
|
||||
mu sync.RWMutex // 房间锁
|
||||
config *MediaServerConfig // 配置
|
||||
}
|
||||
|
||||
// MediaServerConfig 媒体服务器配置
|
||||
@@ -65,6 +66,8 @@ func newMediaServer() *MediaServer {
|
||||
if config.Enabled {
|
||||
ms.sfu = NewSFUServer(config)
|
||||
ms.rtmp = NewRTMPServer(config)
|
||||
// 创建 WebSocket to RTMP 代理
|
||||
ms.wsProxy = NewWebMToRTMPProxy(ms.rtmp)
|
||||
}
|
||||
|
||||
return ms
|
||||
@@ -208,6 +211,11 @@ func (ms *MediaServer) GetRTMP() *RTMPServer {
|
||||
return ms.rtmp
|
||||
}
|
||||
|
||||
// GetWSProxy 获取 WebSocket to RTMP 代理
|
||||
func (ms *MediaServer) GetWSProxy() *WebMToRTMPProxy {
|
||||
return ms.wsProxy
|
||||
}
|
||||
|
||||
// GetRoomCount 获取房间数量
|
||||
func (ms *MediaServer) GetRoomCount() int {
|
||||
ms.mu.RLock()
|
||||
|
||||
736
internal/mediaserver/ws_rtmp_proxy.go
Normal file
736
internal/mediaserver/ws_rtmp_proxy.go
Normal file
@@ -0,0 +1,736 @@
|
||||
/**
|
||||
* 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)
|
||||
|
||||
@@ -5,6 +5,7 @@
|
||||
package middleware
|
||||
|
||||
import (
|
||||
"log"
|
||||
"net/http"
|
||||
"xk-websocket-v2/internal/utils"
|
||||
|
||||
@@ -60,6 +61,9 @@ func JWTAuthMiddleware() gin.HandlerFunc {
|
||||
return
|
||||
}
|
||||
|
||||
// 调试日志:打印解析出的用户ID,帮助排查JWT问题
|
||||
log.Printf("🔑 [JWT] 请求: %s %s | 解析出的用户ID: %s", c.Request.Method, c.Request.URL.Path, userID)
|
||||
|
||||
// 步骤6: 将解析出的用户ID注入到Context中
|
||||
// 后续处理器可以通过 c.Get("user_id") 获取当前用户ID
|
||||
c.Set("user_id", userID)
|
||||
|
||||
@@ -40,6 +40,12 @@ var skipResponseBodyRoutes = []string{
|
||||
"/static/",
|
||||
}
|
||||
|
||||
// 需要完全跳过中间件的路由(如 WebSocket)
|
||||
var skipMiddlewareRoutes = []string{
|
||||
"/ws",
|
||||
"/api/call/ws-push",
|
||||
}
|
||||
|
||||
// isBinaryContentType 检测是否为二进制 Content-Type
|
||||
func isBinaryContentType(contentType string) bool {
|
||||
contentType = strings.ToLower(contentType)
|
||||
@@ -113,6 +119,16 @@ func NewRequestLogMiddleware(db *gorm.DB) *RequestLogMiddleware {
|
||||
return &RequestLogMiddleware{DB: db}
|
||||
}
|
||||
|
||||
// shouldSkipMiddleware 检测是否应该完全跳过中间件
|
||||
func shouldSkipMiddleware(path string) bool {
|
||||
for _, route := range skipMiddlewareRoutes {
|
||||
if strings.HasPrefix(path, route) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// Handler 中间件处理函数
|
||||
func (m *RequestLogMiddleware) Handler() gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
@@ -122,6 +138,12 @@ func (m *RequestLogMiddleware) Handler() gin.HandlerFunc {
|
||||
return
|
||||
}
|
||||
|
||||
// 跳过 WebSocket 路由(包装 Writer 会干扰 WebSocket 升级)
|
||||
if shouldSkipMiddleware(c.Request.URL.Path) {
|
||||
c.Next()
|
||||
return
|
||||
}
|
||||
|
||||
// 获取请求IP
|
||||
ip := getClientIP(c)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user