1012 lines
24 KiB
Go
1012 lines
24 KiB
Go
/**
|
||
* package mediaserver
|
||
*
|
||
* RTMP 服务(用于微信小程序 live-pusher/live-player)
|
||
* 使用 github.com/yutopp/go-rtmp 实现完整的 RTMP 协议
|
||
*
|
||
* 功能:
|
||
* 1. 接收小程序推流(publish)
|
||
* 2. 提供拉流服务(play)
|
||
* 3. 支持 HTTP-FLV 协议
|
||
*/
|
||
package mediaserver
|
||
|
||
import (
|
||
"bytes"
|
||
"crypto/hmac"
|
||
"crypto/sha1"
|
||
"encoding/base64"
|
||
"fmt"
|
||
"io"
|
||
"log"
|
||
"net"
|
||
"net/http"
|
||
"strings"
|
||
"sync"
|
||
"time"
|
||
|
||
"github.com/spf13/viper"
|
||
"github.com/yutopp/go-flv"
|
||
flvtag "github.com/yutopp/go-flv/tag"
|
||
"github.com/yutopp/go-rtmp"
|
||
rtmpmsg "github.com/yutopp/go-rtmp/message"
|
||
)
|
||
|
||
// RTMPServer RTMP 服务器
|
||
type RTMPServer struct {
|
||
config *MediaServerConfig
|
||
mu sync.RWMutex
|
||
running bool
|
||
streams map[string]*RTMPStream // streamID -> stream
|
||
streamsMu sync.RWMutex
|
||
|
||
// 服务器监听器
|
||
rtmpListener net.Listener
|
||
httpServer *http.Server
|
||
rtmpServer *rtmp.Server
|
||
|
||
// 订阅者管理
|
||
subscribers map[string]map[string]*Subscriber // streamID -> subscriberID -> subscriber
|
||
subMu sync.RWMutex
|
||
}
|
||
|
||
// RTMPStream RTMP 流信息
|
||
type RTMPStream struct {
|
||
ID string `json:"id"`
|
||
RoomID string `json:"room_id"`
|
||
UserID string `json:"user_id"`
|
||
PushURL string `json:"push_url"`
|
||
PullURL string `json:"pull_url"`
|
||
FLVURL string `json:"flv_url"`
|
||
CreatedAt time.Time `json:"created_at"`
|
||
IsActive bool `json:"is_active"`
|
||
|
||
// 流数据
|
||
mu sync.RWMutex
|
||
flvHeader []byte // FLV header
|
||
metaData []byte // FLV metadata tag
|
||
videoHeader []byte // Video sequence header (AVC)
|
||
audioHeader []byte // Audio sequence header (AAC)
|
||
gopCache [][]byte // GOP cache for new subscribers
|
||
|
||
// 控制
|
||
stopChan chan struct{} `json:"-"`
|
||
}
|
||
|
||
// Subscriber 订阅者
|
||
type Subscriber struct {
|
||
ID string
|
||
StreamID string
|
||
DataChan chan []byte
|
||
Done chan struct{}
|
||
}
|
||
|
||
// NewRTMPServer 创建 RTMP 服务器
|
||
func NewRTMPServer(config *MediaServerConfig) *RTMPServer {
|
||
return &RTMPServer{
|
||
config: config,
|
||
streams: make(map[string]*RTMPStream),
|
||
subscribers: make(map[string]map[string]*Subscriber),
|
||
}
|
||
}
|
||
|
||
// Start 启动 RTMP 服务器
|
||
func (r *RTMPServer) Start() {
|
||
r.mu.Lock()
|
||
if r.running {
|
||
r.mu.Unlock()
|
||
return
|
||
}
|
||
r.running = true
|
||
r.mu.Unlock()
|
||
|
||
// 启动 RTMP 服务器
|
||
go r.startRTMPServer()
|
||
|
||
// 启动 HTTP-FLV 服务器
|
||
go r.startHTTPFLVServer()
|
||
|
||
// 启动流清理任务
|
||
go r.cleanupTask()
|
||
|
||
log.Printf("🚀 [RTMP] RTMP 服务已启动 | RTMP端口: %d | HTTP-FLV端口: %d",
|
||
r.config.RTMPPort, r.config.HTTPFLVPort)
|
||
}
|
||
|
||
// startRTMPServer 启动 RTMP 服务
|
||
func (r *RTMPServer) startRTMPServer() {
|
||
addr := fmt.Sprintf(":%d", r.config.RTMPPort)
|
||
|
||
var err error
|
||
r.rtmpListener, err = net.Listen("tcp", addr)
|
||
if err != nil {
|
||
log.Printf("❌ [RTMP] RTMP 监听失败: %v", err)
|
||
return
|
||
}
|
||
|
||
log.Printf("🎬 [RTMP] RTMP 服务监听: %s", addr)
|
||
|
||
// 创建 RTMP 服务器
|
||
r.rtmpServer = rtmp.NewServer(&rtmp.ServerConfig{
|
||
OnConnect: func(conn net.Conn) (io.ReadWriteCloser, *rtmp.ConnConfig) {
|
||
log.Printf("🎬 [RTMP] 新连接: %s", conn.RemoteAddr().String())
|
||
return conn, &rtmp.ConnConfig{
|
||
Handler: &RTMPHandler{
|
||
server: r,
|
||
conn: conn,
|
||
},
|
||
ControlState: rtmp.StreamControlStateConfig{
|
||
DefaultBandwidthWindowSize: 6 * 1024 * 1024,
|
||
},
|
||
}
|
||
},
|
||
})
|
||
|
||
// 启动服务
|
||
if err := r.rtmpServer.Serve(r.rtmpListener); err != nil {
|
||
if r.running {
|
||
log.Printf("❌ [RTMP] 服务异常: %v", err)
|
||
}
|
||
}
|
||
}
|
||
|
||
// RTMPHandler RTMP 连接处理器
|
||
type RTMPHandler struct {
|
||
rtmp.DefaultHandler
|
||
server *RTMPServer
|
||
conn net.Conn
|
||
streamID string
|
||
isPublish bool
|
||
}
|
||
|
||
// OnServe 连接建立时调用
|
||
func (h *RTMPHandler) OnServe(conn *rtmp.Conn) {
|
||
log.Printf("🎬 [RTMP] OnServe: %s", h.conn.RemoteAddr().String())
|
||
}
|
||
|
||
// OnConnect 处理 connect 命令
|
||
func (h *RTMPHandler) OnConnect(timestamp uint32, cmd *rtmpmsg.NetConnectionConnect) error {
|
||
log.Printf("🎬 [RTMP] OnConnect: app=%v", cmd.Command.App)
|
||
return nil
|
||
}
|
||
|
||
// OnCreateStream 处理 createStream 命令
|
||
func (h *RTMPHandler) OnCreateStream(timestamp uint32, cmd *rtmpmsg.NetConnectionCreateStream) error {
|
||
log.Printf("🎬 [RTMP] OnCreateStream")
|
||
return nil
|
||
}
|
||
|
||
// OnReleaseStream 处理 releaseStream 命令
|
||
func (h *RTMPHandler) OnReleaseStream(timestamp uint32, cmd *rtmpmsg.NetConnectionReleaseStream) error {
|
||
streamName := cmd.StreamName
|
||
if idx := strings.Index(streamName, "?"); idx != -1 {
|
||
streamName = streamName[:idx]
|
||
}
|
||
log.Printf("🎬 [RTMP] OnReleaseStream: %s", streamName)
|
||
return nil
|
||
}
|
||
|
||
// OnDeleteStream 处理 deleteStream 命令
|
||
func (h *RTMPHandler) OnDeleteStream(timestamp uint32, cmd *rtmpmsg.NetStreamDeleteStream) error {
|
||
log.Printf("🎬 [RTMP] OnDeleteStream")
|
||
return nil
|
||
}
|
||
|
||
// OnFCPublish 处理 FCPublish 命令
|
||
func (h *RTMPHandler) OnFCPublish(timestamp uint32, cmd *rtmpmsg.NetStreamFCPublish) error {
|
||
log.Printf("🎬 [RTMP] OnFCPublish: %s", cmd.StreamName)
|
||
return nil
|
||
}
|
||
|
||
// OnFCUnpublish 处理 FCUnpublish 命令
|
||
func (h *RTMPHandler) OnFCUnpublish(timestamp uint32, cmd *rtmpmsg.NetStreamFCUnpublish) error {
|
||
log.Printf("🎬 [RTMP] OnFCUnpublish: %s", cmd.StreamName)
|
||
return nil
|
||
}
|
||
|
||
// OnPublish 处理 publish 命令
|
||
func (h *RTMPHandler) OnPublish(_ *rtmp.StreamContext, timestamp uint32, cmd *rtmpmsg.NetStreamPublish) error {
|
||
log.Printf("🎬 [RTMP] OnPublish: name=%s, type=%s", cmd.PublishingName, cmd.PublishingType)
|
||
|
||
// 解析 stream name,去掉 token 参数
|
||
streamName := cmd.PublishingName
|
||
if idx := strings.Index(streamName, "?"); idx != -1 {
|
||
streamName = streamName[:idx]
|
||
}
|
||
|
||
h.streamID = streamName
|
||
h.isPublish = true
|
||
|
||
log.Printf("🎬 [RTMP] 解析后的流ID: %s", h.streamID)
|
||
|
||
// 检查流是否已注册
|
||
h.server.streamsMu.RLock()
|
||
stream, exists := h.server.streams[h.streamID]
|
||
h.server.streamsMu.RUnlock()
|
||
|
||
if !exists {
|
||
log.Printf("⚠️ [RTMP] 流未注册,自动创建: %s", h.streamID)
|
||
// 自动创建流(用于测试)
|
||
stream = &RTMPStream{
|
||
ID: h.streamID,
|
||
CreatedAt: time.Now(),
|
||
IsActive: true,
|
||
stopChan: make(chan struct{}),
|
||
gopCache: make([][]byte, 0),
|
||
}
|
||
h.server.streamsMu.Lock()
|
||
h.server.streams[h.streamID] = stream
|
||
h.server.streamsMu.Unlock()
|
||
} else {
|
||
log.Printf("✅ [RTMP] 找到已注册的流: %s", h.streamID)
|
||
}
|
||
|
||
stream.mu.Lock()
|
||
stream.IsActive = true
|
||
stream.mu.Unlock()
|
||
|
||
log.Printf("✅ [RTMP] 开始推流: %s", h.streamID)
|
||
|
||
return nil
|
||
}
|
||
|
||
// OnPlay 处理 play 命令
|
||
func (h *RTMPHandler) OnPlay(ctx *rtmp.StreamContext, timestamp uint32, cmd *rtmpmsg.NetStreamPlay) error {
|
||
log.Printf("🎬 [RTMP] OnPlay: name=%s", cmd.StreamName)
|
||
|
||
h.streamID = cmd.StreamName
|
||
h.isPublish = false
|
||
|
||
// 检查流是否存在
|
||
h.server.streamsMu.RLock()
|
||
stream, exists := h.server.streams[h.streamID]
|
||
h.server.streamsMu.RUnlock()
|
||
|
||
if !exists || !stream.IsActive {
|
||
log.Printf("⚠️ [RTMP] 流不存在或未激活: %s", h.streamID)
|
||
return fmt.Errorf("stream not found: %s", h.streamID)
|
||
}
|
||
|
||
log.Printf("✅ [RTMP] 开始播放: %s", h.streamID)
|
||
|
||
return nil
|
||
}
|
||
|
||
// OnSetDataFrame 处理 metadata
|
||
func (h *RTMPHandler) OnSetDataFrame(timestamp uint32, data *rtmpmsg.NetStreamSetDataFrame) error {
|
||
log.Printf("🎬 [RTMP] OnSetDataFrame")
|
||
|
||
if h.streamID == "" {
|
||
return nil
|
||
}
|
||
|
||
h.server.streamsMu.RLock()
|
||
stream, exists := h.server.streams[h.streamID]
|
||
h.server.streamsMu.RUnlock()
|
||
|
||
if exists && len(data.Payload) > 0 {
|
||
stream.mu.Lock()
|
||
stream.metaData = data.Payload
|
||
stream.mu.Unlock()
|
||
}
|
||
|
||
return nil
|
||
}
|
||
|
||
// OnAudio 处理音频数据
|
||
func (h *RTMPHandler) OnAudio(timestamp uint32, payload io.Reader) error {
|
||
if h.streamID == "" || !h.isPublish {
|
||
return nil
|
||
}
|
||
|
||
// 读取音频数据
|
||
data, err := io.ReadAll(payload)
|
||
if err != nil {
|
||
return err
|
||
}
|
||
|
||
if len(data) == 0 {
|
||
return nil
|
||
}
|
||
|
||
h.server.streamsMu.RLock()
|
||
stream, exists := h.server.streams[h.streamID]
|
||
h.server.streamsMu.RUnlock()
|
||
|
||
if !exists {
|
||
return nil
|
||
}
|
||
|
||
// 保存 AAC sequence header
|
||
if len(data) > 1 {
|
||
soundFormat := (data[0] >> 4) & 0x0f
|
||
if soundFormat == 10 { // AAC
|
||
aacPacketType := data[1]
|
||
if aacPacketType == 0 { // Sequence header
|
||
stream.mu.Lock()
|
||
stream.audioHeader = make([]byte, len(data))
|
||
copy(stream.audioHeader, data)
|
||
stream.mu.Unlock()
|
||
}
|
||
}
|
||
}
|
||
|
||
// 创建 FLV audio tag
|
||
flvData := h.createFLVAudioTag(timestamp, data)
|
||
if flvData != nil {
|
||
h.broadcastToSubscribers(h.streamID, flvData)
|
||
}
|
||
|
||
return nil
|
||
}
|
||
|
||
// OnVideo 处理视频数据
|
||
func (h *RTMPHandler) OnVideo(timestamp uint32, payload io.Reader) error {
|
||
if h.streamID == "" || !h.isPublish {
|
||
return nil
|
||
}
|
||
|
||
// 读取视频数据
|
||
data, err := io.ReadAll(payload)
|
||
if err != nil {
|
||
return err
|
||
}
|
||
|
||
if len(data) == 0 {
|
||
return nil
|
||
}
|
||
|
||
h.server.streamsMu.RLock()
|
||
stream, exists := h.server.streams[h.streamID]
|
||
h.server.streamsMu.RUnlock()
|
||
|
||
if !exists {
|
||
return nil
|
||
}
|
||
|
||
// 解析视频帧信息
|
||
frameType := (data[0] >> 4) & 0x0f
|
||
codecID := data[0] & 0x0f
|
||
|
||
// 保存 AVC sequence header
|
||
if codecID == 7 && len(data) > 1 { // AVC
|
||
avcPacketType := data[1]
|
||
if avcPacketType == 0 { // Sequence header
|
||
stream.mu.Lock()
|
||
stream.videoHeader = make([]byte, len(data))
|
||
copy(stream.videoHeader, data)
|
||
stream.mu.Unlock()
|
||
}
|
||
|
||
// 关键帧时清空 GOP 缓存
|
||
if frameType == 1 { // Keyframe
|
||
stream.mu.Lock()
|
||
stream.gopCache = make([][]byte, 0)
|
||
stream.mu.Unlock()
|
||
}
|
||
}
|
||
|
||
// 创建 FLV video tag
|
||
flvData := h.createFLVVideoTag(timestamp, data)
|
||
if flvData != nil {
|
||
// 缓存到 GOP
|
||
stream.mu.Lock()
|
||
stream.gopCache = append(stream.gopCache, flvData)
|
||
// 限制 GOP 缓存大小
|
||
if len(stream.gopCache) > 300 {
|
||
stream.gopCache = stream.gopCache[len(stream.gopCache)-300:]
|
||
}
|
||
stream.mu.Unlock()
|
||
|
||
h.broadcastToSubscribers(h.streamID, flvData)
|
||
}
|
||
|
||
return nil
|
||
}
|
||
|
||
// OnUnknownMessage 处理未知消息
|
||
func (h *RTMPHandler) OnUnknownMessage(timestamp uint32, msg rtmpmsg.Message) error {
|
||
return nil
|
||
}
|
||
|
||
// OnUnknownCommandMessage 处理未知命令消息
|
||
func (h *RTMPHandler) OnUnknownCommandMessage(timestamp uint32, cmd *rtmpmsg.CommandMessage) error {
|
||
return nil
|
||
}
|
||
|
||
// OnUnknownDataMessage 处理未知数据消息
|
||
func (h *RTMPHandler) OnUnknownDataMessage(timestamp uint32, data *rtmpmsg.DataMessage) error {
|
||
return nil
|
||
}
|
||
|
||
// createFLVAudioTag 创建 FLV 音频 tag
|
||
func (h *RTMPHandler) createFLVAudioTag(timestamp uint32, data []byte) []byte {
|
||
// FLV tag 格式:
|
||
// TagType (1 byte): 8 = audio
|
||
// DataSize (3 bytes): big-endian
|
||
// Timestamp (3 bytes): big-endian
|
||
// TimestampExtended (1 byte)
|
||
// StreamID (3 bytes): always 0
|
||
// Data
|
||
// PreviousTagSize (4 bytes): big-endian
|
||
|
||
dataSize := len(data)
|
||
tagSize := 11 + dataSize + 4 // header + data + previous tag size
|
||
|
||
tag := make([]byte, tagSize)
|
||
|
||
// Tag type
|
||
tag[0] = 8 // Audio
|
||
|
||
// 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 (11 + dataSize)
|
||
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
|
||
}
|
||
|
||
// createFLVVideoTag 创建 FLV 视频 tag
|
||
func (h *RTMPHandler) createFLVVideoTag(timestamp uint32, data []byte) []byte {
|
||
dataSize := len(data)
|
||
tagSize := 11 + dataSize + 4
|
||
|
||
tag := make([]byte, tagSize)
|
||
|
||
// Tag type
|
||
tag[0] = 9 // Video
|
||
|
||
// 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
|
||
}
|
||
|
||
// broadcastToSubscribers 向所有订阅者广播数据
|
||
func (h *RTMPHandler) broadcastToSubscribers(streamID string, data []byte) {
|
||
h.server.subMu.RLock()
|
||
subs, exists := h.server.subscribers[streamID]
|
||
if !exists {
|
||
h.server.subMu.RUnlock()
|
||
return
|
||
}
|
||
|
||
for _, sub := range subs {
|
||
select {
|
||
case sub.DataChan <- data:
|
||
default:
|
||
// 缓冲区满,跳过
|
||
}
|
||
}
|
||
h.server.subMu.RUnlock()
|
||
}
|
||
|
||
// OnClose 连接关闭
|
||
func (h *RTMPHandler) OnClose() {
|
||
log.Printf("🎬 [RTMP] OnClose: stream=%s, publish=%v", h.streamID, h.isPublish)
|
||
|
||
if h.isPublish && h.streamID != "" {
|
||
h.server.streamsMu.Lock()
|
||
if stream, exists := h.server.streams[h.streamID]; exists {
|
||
stream.mu.Lock()
|
||
stream.IsActive = false
|
||
stream.mu.Unlock()
|
||
}
|
||
h.server.streamsMu.Unlock()
|
||
log.Printf("⏹️ [RTMP] 停止推流: %s", h.streamID)
|
||
}
|
||
}
|
||
|
||
// startHTTPFLVServer 启动 HTTP-FLV 服务
|
||
func (r *RTMPServer) startHTTPFLVServer() {
|
||
addr := fmt.Sprintf(":%d", r.config.HTTPFLVPort)
|
||
|
||
mux := http.NewServeMux()
|
||
mux.HandleFunc("/live/", r.handleFLVRequest)
|
||
mux.HandleFunc("/api/streams", r.handleStreamsAPI)
|
||
|
||
r.httpServer = &http.Server{
|
||
Addr: addr,
|
||
Handler: mux,
|
||
}
|
||
|
||
log.Printf("🎬 [RTMP] HTTP-FLV 服务监听: %s", addr)
|
||
|
||
if err := r.httpServer.ListenAndServe(); err != nil && err != http.ErrServerClosed {
|
||
log.Printf("❌ [RTMP] HTTP-FLV 启动失败: %v", err)
|
||
}
|
||
}
|
||
|
||
// handleFLVRequest 处理 FLV 请求
|
||
func (r *RTMPServer) handleFLVRequest(w http.ResponseWriter, req *http.Request) {
|
||
// 解析流ID: /live/{streamID}.flv
|
||
path := req.URL.Path
|
||
if len(path) < 7 {
|
||
http.Error(w, "Invalid path", http.StatusBadRequest)
|
||
return
|
||
}
|
||
|
||
// 提取 streamID
|
||
streamPath := path[6:] // 去掉 "/live/"
|
||
if len(streamPath) > 4 && streamPath[len(streamPath)-4:] == ".flv" {
|
||
streamPath = streamPath[:len(streamPath)-4]
|
||
}
|
||
|
||
r.streamsMu.RLock()
|
||
stream, exists := r.streams[streamPath]
|
||
r.streamsMu.RUnlock()
|
||
|
||
if !exists || !stream.IsActive {
|
||
http.Error(w, "Stream not found", http.StatusNotFound)
|
||
return
|
||
}
|
||
|
||
log.Printf("🎬 [RTMP] FLV 请求: %s", streamPath)
|
||
|
||
// 设置响应头
|
||
w.Header().Set("Content-Type", "video/x-flv")
|
||
w.Header().Set("Access-Control-Allow-Origin", "*")
|
||
w.Header().Set("Transfer-Encoding", "chunked")
|
||
w.Header().Set("Connection", "keep-alive")
|
||
w.Header().Set("Cache-Control", "no-cache")
|
||
|
||
// 发送 FLV header
|
||
// FLV signature (3 bytes) + version (1 byte) + flags (1 byte) + header size (4 bytes)
|
||
flvHeader := []byte{0x46, 0x4C, 0x56, 0x01, 0x05, 0x00, 0x00, 0x00, 0x09, 0x00, 0x00, 0x00, 0x00}
|
||
w.Write(flvHeader)
|
||
|
||
if f, ok := w.(http.Flusher); ok {
|
||
f.Flush()
|
||
}
|
||
|
||
// 创建订阅者
|
||
subID := fmt.Sprintf("%d", time.Now().UnixNano())
|
||
sub := &Subscriber{
|
||
ID: subID,
|
||
StreamID: streamPath,
|
||
DataChan: make(chan []byte, 100),
|
||
Done: make(chan struct{}),
|
||
}
|
||
|
||
// 注册订阅者
|
||
r.subMu.Lock()
|
||
if r.subscribers[streamPath] == nil {
|
||
r.subscribers[streamPath] = make(map[string]*Subscriber)
|
||
}
|
||
r.subscribers[streamPath][subID] = sub
|
||
r.subMu.Unlock()
|
||
|
||
// 清理
|
||
defer func() {
|
||
r.subMu.Lock()
|
||
delete(r.subscribers[streamPath], subID)
|
||
r.subMu.Unlock()
|
||
close(sub.Done)
|
||
}()
|
||
|
||
// 发送缓存的头信息
|
||
stream.mu.RLock()
|
||
|
||
// 发送 metadata
|
||
if len(stream.metaData) > 0 {
|
||
w.Write(stream.metaData)
|
||
}
|
||
|
||
// 发送视频 sequence header
|
||
if len(stream.videoHeader) > 0 {
|
||
flvTag := createFLVTagFromData(9, 0, stream.videoHeader)
|
||
w.Write(flvTag)
|
||
}
|
||
|
||
// 发送音频 sequence header
|
||
if len(stream.audioHeader) > 0 {
|
||
flvTag := createFLVTagFromData(8, 0, stream.audioHeader)
|
||
w.Write(flvTag)
|
||
}
|
||
|
||
// 发送 GOP 缓存
|
||
for _, data := range stream.gopCache {
|
||
w.Write(data)
|
||
}
|
||
stream.mu.RUnlock()
|
||
|
||
if f, ok := w.(http.Flusher); ok {
|
||
f.Flush()
|
||
}
|
||
|
||
// 持续发送数据
|
||
for {
|
||
select {
|
||
case <-req.Context().Done():
|
||
return
|
||
case data := <-sub.DataChan:
|
||
if _, err := w.Write(data); err != nil {
|
||
return
|
||
}
|
||
if f, ok := w.(http.Flusher); ok {
|
||
f.Flush()
|
||
}
|
||
}
|
||
}
|
||
}
|
||
|
||
// createFLVTagFromData 从原始数据创建 FLV tag
|
||
func createFLVTagFromData(tagType byte, timestamp uint32, data []byte) []byte {
|
||
dataSize := len(data)
|
||
tagSize := 11 + dataSize + 4
|
||
|
||
tag := make([]byte, tagSize)
|
||
|
||
tag[0] = tagType
|
||
tag[1] = byte((dataSize >> 16) & 0xff)
|
||
tag[2] = byte((dataSize >> 8) & 0xff)
|
||
tag[3] = byte(dataSize & 0xff)
|
||
tag[4] = byte((timestamp >> 16) & 0xff)
|
||
tag[5] = byte((timestamp >> 8) & 0xff)
|
||
tag[6] = byte(timestamp & 0xff)
|
||
tag[7] = byte((timestamp >> 24) & 0xff)
|
||
tag[8] = 0
|
||
tag[9] = 0
|
||
tag[10] = 0
|
||
|
||
copy(tag[11:], data)
|
||
|
||
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
|
||
}
|
||
|
||
// handleStreamsAPI 处理流列表 API
|
||
func (r *RTMPServer) handleStreamsAPI(w http.ResponseWriter, req *http.Request) {
|
||
w.Header().Set("Content-Type", "application/json")
|
||
w.Header().Set("Access-Control-Allow-Origin", "*")
|
||
|
||
r.streamsMu.RLock()
|
||
defer r.streamsMu.RUnlock()
|
||
|
||
streams := make([]map[string]interface{}, 0)
|
||
for _, s := range r.streams {
|
||
if s.IsActive {
|
||
streams = append(streams, map[string]interface{}{
|
||
"id": s.ID,
|
||
"room_id": s.RoomID,
|
||
"user_id": s.UserID,
|
||
"pull_url": s.PullURL,
|
||
"flv_url": s.FLVURL,
|
||
})
|
||
}
|
||
}
|
||
|
||
fmt.Fprintf(w, `{"streams":%d,"data":%v}`, len(streams), streams)
|
||
}
|
||
|
||
// Stop 停止 RTMP 服务器
|
||
func (r *RTMPServer) Stop() {
|
||
r.mu.Lock()
|
||
defer r.mu.Unlock()
|
||
|
||
r.running = false
|
||
|
||
// 关闭 RTMP 服务器
|
||
if r.rtmpServer != nil {
|
||
r.rtmpServer.Close()
|
||
}
|
||
|
||
// 关闭 RTMP 监听器
|
||
if r.rtmpListener != nil {
|
||
r.rtmpListener.Close()
|
||
}
|
||
|
||
// 关闭 HTTP 服务器
|
||
if r.httpServer != nil {
|
||
r.httpServer.Close()
|
||
}
|
||
|
||
// 关闭所有流
|
||
r.streamsMu.Lock()
|
||
for _, stream := range r.streams {
|
||
stream.IsActive = false
|
||
if stream.stopChan != nil {
|
||
select {
|
||
case <-stream.stopChan:
|
||
default:
|
||
close(stream.stopChan)
|
||
}
|
||
}
|
||
}
|
||
r.streams = make(map[string]*RTMPStream)
|
||
r.streamsMu.Unlock()
|
||
|
||
// 关闭所有订阅者
|
||
r.subMu.Lock()
|
||
for _, subs := range r.subscribers {
|
||
for _, sub := range subs {
|
||
close(sub.DataChan)
|
||
}
|
||
}
|
||
r.subscribers = make(map[string]map[string]*Subscriber)
|
||
r.subMu.Unlock()
|
||
|
||
log.Println("🛑 [RTMP] 已停止")
|
||
}
|
||
|
||
// GenerateStreamURLs 为用户生成推拉流地址
|
||
func (r *RTMPServer) GenerateStreamURLs(roomID, userID string) (*RTMPStream, error) {
|
||
r.mu.RLock()
|
||
if !r.running {
|
||
r.mu.RUnlock()
|
||
return nil, ErrRTMPNotReady
|
||
}
|
||
r.mu.RUnlock()
|
||
|
||
// 生成唯一的流ID
|
||
streamID := fmt.Sprintf("%s_%s_%d", roomID, userID, time.Now().UnixNano())
|
||
|
||
// 获取服务器配置
|
||
publicIP := r.config.PublicIP
|
||
if publicIP == "" {
|
||
publicIP = viper.GetString("turn.public_ip")
|
||
}
|
||
|
||
// 生成鉴权 Token
|
||
token := generateStreamToken(streamID)
|
||
|
||
// 生成 URL
|
||
stream := &RTMPStream{
|
||
ID: streamID,
|
||
RoomID: roomID,
|
||
UserID: userID,
|
||
PushURL: fmt.Sprintf("rtmp://%s:%d/live/%s?token=%s", publicIP, r.config.RTMPPort, streamID, token),
|
||
PullURL: fmt.Sprintf("rtmp://%s:%d/live/%s", publicIP, r.config.RTMPPort, streamID),
|
||
FLVURL: fmt.Sprintf("http://%s:%d/live/%s.flv", publicIP, r.config.HTTPFLVPort, streamID),
|
||
CreatedAt: time.Now(),
|
||
IsActive: true,
|
||
stopChan: make(chan struct{}),
|
||
gopCache: make([][]byte, 0),
|
||
}
|
||
|
||
r.streamsMu.Lock()
|
||
r.streams[streamID] = stream
|
||
r.streamsMu.Unlock()
|
||
|
||
log.Printf("🎬 [RTMP] 创建流 | Room:%s User:%s Stream:%s", roomID, userID, streamID)
|
||
|
||
return stream, nil
|
||
}
|
||
|
||
// GetStream 获取流信息
|
||
func (r *RTMPServer) GetStream(streamID string) *RTMPStream {
|
||
r.streamsMu.RLock()
|
||
defer r.streamsMu.RUnlock()
|
||
return r.streams[streamID]
|
||
}
|
||
|
||
// GetStreamsByRoom 获取房间内所有流
|
||
func (r *RTMPServer) GetStreamsByRoom(roomID string) []*RTMPStream {
|
||
r.streamsMu.RLock()
|
||
defer r.streamsMu.RUnlock()
|
||
|
||
streams := make([]*RTMPStream, 0)
|
||
for _, s := range r.streams {
|
||
if s.RoomID == roomID && s.IsActive {
|
||
streams = append(streams, s)
|
||
}
|
||
}
|
||
return streams
|
||
}
|
||
|
||
// GetStreamsByUser 获取用户的所有流
|
||
func (r *RTMPServer) GetStreamsByUser(userID string) []*RTMPStream {
|
||
r.streamsMu.RLock()
|
||
defer r.streamsMu.RUnlock()
|
||
|
||
streams := make([]*RTMPStream, 0)
|
||
for _, s := range r.streams {
|
||
if s.UserID == userID && s.IsActive {
|
||
streams = append(streams, s)
|
||
}
|
||
}
|
||
return streams
|
||
}
|
||
|
||
// RemoveStream 移除流
|
||
func (r *RTMPServer) RemoveStream(streamID string) {
|
||
r.streamsMu.Lock()
|
||
defer r.streamsMu.Unlock()
|
||
|
||
if stream, exists := r.streams[streamID]; exists {
|
||
stream.IsActive = false
|
||
if stream.stopChan != nil {
|
||
select {
|
||
case <-stream.stopChan:
|
||
default:
|
||
close(stream.stopChan)
|
||
}
|
||
}
|
||
delete(r.streams, streamID)
|
||
log.Printf("🎬 [RTMP] 移除流 | Stream:%s", streamID)
|
||
}
|
||
|
||
// 移除订阅者
|
||
r.subMu.Lock()
|
||
delete(r.subscribers, streamID)
|
||
r.subMu.Unlock()
|
||
}
|
||
|
||
// RemoveStreamsByUser 移除用户的所有流
|
||
func (r *RTMPServer) RemoveStreamsByUser(userID string) {
|
||
r.streamsMu.Lock()
|
||
defer r.streamsMu.Unlock()
|
||
|
||
for streamID, stream := range r.streams {
|
||
if stream.UserID == userID {
|
||
stream.IsActive = false
|
||
if stream.stopChan != nil {
|
||
select {
|
||
case <-stream.stopChan:
|
||
default:
|
||
close(stream.stopChan)
|
||
}
|
||
}
|
||
delete(r.streams, streamID)
|
||
log.Printf("🎬 [RTMP] 移除用户流 | User:%s Stream:%s", userID, streamID)
|
||
|
||
// 移除订阅者
|
||
r.subMu.Lock()
|
||
delete(r.subscribers, streamID)
|
||
r.subMu.Unlock()
|
||
}
|
||
}
|
||
}
|
||
|
||
// RemoveStreamsByRoom 移除房间的所有流
|
||
func (r *RTMPServer) RemoveStreamsByRoom(roomID string) {
|
||
r.streamsMu.Lock()
|
||
defer r.streamsMu.Unlock()
|
||
|
||
for streamID, stream := range r.streams {
|
||
if stream.RoomID == roomID {
|
||
stream.IsActive = false
|
||
if stream.stopChan != nil {
|
||
select {
|
||
case <-stream.stopChan:
|
||
default:
|
||
close(stream.stopChan)
|
||
}
|
||
}
|
||
delete(r.streams, streamID)
|
||
log.Printf("🎬 [RTMP] 移除房间流 | Room:%s Stream:%s", roomID, streamID)
|
||
|
||
// 移除订阅者
|
||
r.subMu.Lock()
|
||
delete(r.subscribers, streamID)
|
||
r.subMu.Unlock()
|
||
}
|
||
}
|
||
}
|
||
|
||
// GetStreamCount 获取流数量
|
||
func (r *RTMPServer) GetStreamCount() int {
|
||
r.streamsMu.RLock()
|
||
defer r.streamsMu.RUnlock()
|
||
return len(r.streams)
|
||
}
|
||
|
||
// cleanupTask 清理过期流
|
||
func (r *RTMPServer) cleanupTask() {
|
||
ticker := time.NewTicker(5 * time.Minute)
|
||
defer ticker.Stop()
|
||
|
||
for {
|
||
select {
|
||
case <-ticker.C:
|
||
r.cleanupExpiredStreams()
|
||
}
|
||
|
||
r.mu.RLock()
|
||
if !r.running {
|
||
r.mu.RUnlock()
|
||
return
|
||
}
|
||
r.mu.RUnlock()
|
||
}
|
||
}
|
||
|
||
// cleanupExpiredStreams 清理过期的流
|
||
func (r *RTMPServer) cleanupExpiredStreams() {
|
||
r.streamsMu.Lock()
|
||
defer r.streamsMu.Unlock()
|
||
|
||
expireTime := time.Now().Add(-1 * time.Hour) // 1小时过期
|
||
|
||
for streamID, stream := range r.streams {
|
||
if stream.CreatedAt.Before(expireTime) && !stream.IsActive {
|
||
if stream.stopChan != nil {
|
||
select {
|
||
case <-stream.stopChan:
|
||
default:
|
||
close(stream.stopChan)
|
||
}
|
||
}
|
||
delete(r.streams, streamID)
|
||
log.Printf("🗑️ [RTMP] 清理过期流 | Stream:%s", streamID)
|
||
}
|
||
}
|
||
}
|
||
|
||
// generateStreamToken 生成流鉴权 Token
|
||
func generateStreamToken(streamID string) string {
|
||
secret := viper.GetString("turn.shared_secret")
|
||
timestamp := time.Now().Add(24 * time.Hour).Unix()
|
||
data := fmt.Sprintf("%s:%d", streamID, timestamp)
|
||
|
||
mac := hmac.New(sha1.New, []byte(secret))
|
||
mac.Write([]byte(data))
|
||
return base64.URLEncoding.EncodeToString(mac.Sum(nil))
|
||
}
|
||
|
||
// ValidateStreamToken 验证流鉴权 Token
|
||
func ValidateStreamToken(streamID, token string) bool {
|
||
expectedToken := generateStreamToken(streamID)
|
||
return token == expectedToken
|
||
}
|
||
|
||
// 确保不使用的导入被使用
|
||
var (
|
||
_ = bytes.NewReader
|
||
_ = flv.NewEncoder
|
||
_ = flvtag.TagTypeVideo
|
||
)
|