Files
nl-im-service/internal/mediaserver/rtmp.go
2025-12-15 09:04:14 +08:00

502 lines
12 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters
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 服务(用于微信小程序 live-pusher/live-player
* 功能:
* 1. 接收小程序推流
* 2. 提供拉流服务
* 3. 支持 HTTP-FLV 协议
*
* 注意:当前版本使用简化实现,实际 RTMP 协议处理需要集成完整的 RTMP 服务器
* 生产环境建议使用 nginx-rtmp 或独立部署 livego
*/
package mediaserver
import (
"crypto/hmac"
"crypto/sha1"
"encoding/base64"
"fmt"
"io"
"log"
"net"
"net/http"
"sync"
"time"
"github.com/spf13/viper"
)
// RTMPServer RTMP 服务器
type RTMPServer struct {
config *MediaServerConfig
mu sync.RWMutex
running bool
streams map[string]*RTMPStream // streamID -> stream
// 服务器监听器
rtmpListener net.Listener
httpServer *http.Server
}
// 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"`
// 流数据缓存
dataChan chan []byte `json:"-"`
stopChan chan struct{} `json:"-"`
}
// NewRTMPServer 创建 RTMP 服务器
func NewRTMPServer(config *MediaServerConfig) *RTMPServer {
return &RTMPServer{
config: config,
streams: make(map[string]*RTMPStream),
}
}
// 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)
for {
r.mu.RLock()
if !r.running {
r.mu.RUnlock()
return
}
r.mu.RUnlock()
conn, err := r.rtmpListener.Accept()
if err != nil {
if !r.running {
return
}
log.Printf("⚠️ [RTMP] 接受连接失败: %v", err)
continue
}
// 处理 RTMP 连接
go r.handleRTMPConnection(conn)
}
}
// handleRTMPConnection 处理 RTMP 连接
// 简化版:基本的 RTMP 握手和流处理
func (r *RTMPServer) handleRTMPConnection(conn net.Conn) {
defer conn.Close()
// RTMP 协议处理
// 步骤1: RTMP 握手
if err := r.doHandshake(conn); err != nil {
log.Printf("⚠️ [RTMP] 握手失败: %v", err)
return
}
log.Printf("🎬 [RTMP] 新连接: %s", conn.RemoteAddr().String())
// 步骤2: 读取 RTMP 消息
// 简化实现:等待数据并转发
buf := make([]byte, 4096)
for {
r.mu.RLock()
if !r.running {
r.mu.RUnlock()
return
}
r.mu.RUnlock()
conn.SetReadDeadline(time.Now().Add(30 * time.Second))
n, err := conn.Read(buf)
if err != nil {
if err != io.EOF {
log.Printf("⚠️ [RTMP] 读取失败: %v", err)
}
return
}
// TODO: 解析 RTMP 消息并路由到对应的流
_ = n
}
}
// doHandshake 执行 RTMP 握手
func (r *RTMPServer) doHandshake(conn net.Conn) error {
// RTMP 握手协议
// C0 + C1 -> S0 + S1 + S2 -> C2
c0c1 := make([]byte, 1537) // C0 (1) + C1 (1536)
conn.SetReadDeadline(time.Now().Add(10 * time.Second))
if _, err := io.ReadFull(conn, c0c1); err != nil {
return fmt.Errorf("read C0C1 failed: %w", err)
}
// 验证版本号 (C0 应该是 0x03)
if c0c1[0] != 0x03 {
return fmt.Errorf("invalid RTMP version: %d", c0c1[0])
}
// 发送 S0 + S1 + S2
s0s1s2 := make([]byte, 3073) // S0 (1) + S1 (1536) + S2 (1536)
s0s1s2[0] = 0x03 // RTMP version 3
// S1: 时间戳(4) + 零(4) + 随机数据(1528)
// S2: 客户端时间戳(4) + 服务器时间戳(4) + 随机数据(1528)
copy(s0s1s2[1:1537], c0c1[1:]) // S1 = C1
copy(s0s1s2[1537:], c0c1[1:]) // S2 = C1
conn.SetWriteDeadline(time.Now().Add(10 * time.Second))
if _, err := conn.Write(s0s1s2); err != nil {
return fmt.Errorf("write S0S1S2 failed: %w", err)
}
// 读取 C2
c2 := make([]byte, 1536)
conn.SetReadDeadline(time.Now().Add(10 * time.Second))
if _, err := io.ReadFull(conn, c2); err != nil {
return fmt.Errorf("read C2 failed: %w", err)
}
return nil
}
// 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 { // "/live/" = 6, 至少还需要1个字符
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.mu.RLock()
stream, exists := r.streams[streamPath]
r.mu.RUnlock()
if !exists || !stream.IsActive {
http.Error(w, "Stream not found", http.StatusNotFound)
return
}
// 设置响应头
w.Header().Set("Content-Type", "video/x-flv")
w.Header().Set("Access-Control-Allow-Origin", "*")
w.Header().Set("Transfer-Encoding", "chunked")
// TODO: 实际的 FLV 数据推送
// 当前为占位实现
log.Printf("🎬 [RTMP] FLV 请求: %s", streamPath)
// 发送 FLV header
flvHeader := []byte{0x46, 0x4C, 0x56, 0x01, 0x05, 0x00, 0x00, 0x00, 0x09, 0x00, 0x00, 0x00, 0x00}
w.Write(flvHeader)
// 保持连接直到客户端断开
<-req.Context().Done()
}
// 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.mu.RLock()
defer r.mu.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.rtmpListener != nil {
r.rtmpListener.Close()
}
// 关闭 HTTP 服务器
if r.httpServer != nil {
r.httpServer.Close()
}
// 关闭所有流
for _, stream := range r.streams {
if stream.stopChan != nil {
close(stream.stopChan)
}
}
r.streams = make(map[string]*RTMPStream)
log.Println("🛑 [RTMP] 已停止")
}
// GenerateStreamURLs 为用户生成推拉流地址
func (r *RTMPServer) GenerateStreamURLs(roomID, userID string) (*RTMPStream, error) {
r.mu.Lock()
defer r.mu.Unlock()
if !r.running {
return nil, ErrRTMPNotReady
}
// 生成唯一的流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,
dataChan: make(chan []byte, 100),
stopChan: make(chan struct{}),
}
r.streams[streamID] = stream
log.Printf("🎬 [RTMP] 创建流 | Room:%s User:%s Stream:%s", roomID, userID, streamID)
return stream, nil
}
// GetStream 获取流信息
func (r *RTMPServer) GetStream(streamID string) *RTMPStream {
r.mu.RLock()
defer r.mu.RUnlock()
return r.streams[streamID]
}
// GetStreamsByRoom 获取房间内所有流
func (r *RTMPServer) GetStreamsByRoom(roomID string) []*RTMPStream {
r.mu.RLock()
defer r.mu.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.mu.RLock()
defer r.mu.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.mu.Lock()
defer r.mu.Unlock()
if stream, exists := r.streams[streamID]; exists {
stream.IsActive = false
if stream.stopChan != nil {
close(stream.stopChan)
}
delete(r.streams, streamID)
log.Printf("🎬 [RTMP] 移除流 | Stream:%s", streamID)
}
}
// RemoveStreamsByUser 移除用户的所有流
func (r *RTMPServer) RemoveStreamsByUser(userID string) {
r.mu.Lock()
defer r.mu.Unlock()
for streamID, stream := range r.streams {
if stream.UserID == userID {
stream.IsActive = false
if stream.stopChan != nil {
close(stream.stopChan)
}
delete(r.streams, streamID)
log.Printf("🎬 [RTMP] 移除用户流 | User:%s Stream:%s", userID, streamID)
}
}
}
// RemoveStreamsByRoom 移除房间的所有流
func (r *RTMPServer) RemoveStreamsByRoom(roomID string) {
r.mu.Lock()
defer r.mu.Unlock()
for streamID, stream := range r.streams {
if stream.RoomID == roomID {
stream.IsActive = false
if stream.stopChan != nil {
close(stream.stopChan)
}
delete(r.streams, streamID)
log.Printf("🎬 [RTMP] 移除房间流 | Room:%s Stream:%s", roomID, streamID)
}
}
}
// GetStreamCount 获取流数量
func (r *RTMPServer) GetStreamCount() int {
r.mu.RLock()
defer r.mu.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.mu.Lock()
defer r.mu.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 {
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
}