diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml new file mode 100644 index 0000000..e748c71 --- /dev/null +++ b/.github/workflows/ci.yml @@ -0,0 +1,105 @@ +name: CI + +on: + push: + branches: [main] + pull_request: + branches: [main] + +permissions: + contents: read + +jobs: + lint: + name: Lint + runs-on: ubuntu-latest + steps: + - uses: actions/checkout@v4 + + - uses: actions/setup-go@v5 + with: + go-version-file: go.mod + + - name: golangci-lint + uses: golangci/golangci-lint-action@v6 + with: + version: latest + + test: + name: Test + runs-on: ubuntu-latest + steps: + - uses: actions/checkout@v4 + + - uses: actions/setup-go@v5 + with: + go-version-file: go.mod + + - name: Install FFmpeg dev libraries + run: | + sudo apt-get update + sudo apt-get install -y --no-install-recommends \ + pkg-config libavcodec-dev libswresample-dev libavutil-dev + + - name: Build + run: CGO_ENABLED=1 go build -tags audiocodec ./cmd/liveforge + + - name: Test + run: CGO_ENABLED=1 go test -race -coverprofile=coverage.out -covermode=atomic ./... + + - name: Upload coverage + if: github.event_name == 'pull_request' + uses: actions/upload-artifact@v4 + with: + name: coverage + path: coverage.out + + security: + name: Security Scan + runs-on: ubuntu-latest + steps: + - uses: actions/checkout@v4 + + - uses: actions/setup-go@v5 + with: + go-version-file: go.mod + + - name: Install gosec + run: go install github.com/securego/gosec/v2/cmd/gosec@latest + + - name: Run gosec + run: gosec -exclude-generated ./... + + docker: + name: Docker Build + runs-on: ubuntu-latest + needs: [lint, test, security] + if: github.ref == 'refs/heads/main' + steps: + - uses: actions/checkout@v4 + + - name: Set up Docker Buildx + uses: docker/setup-buildx-action@v3 + + - name: Login to Docker Hub + uses: docker/login-action@v3 + with: + username: ${{ secrets.DOCKERHUB_USERNAME }} + password: ${{ secrets.DOCKERHUB_TOKEN }} + + - name: Extract version + id: version + run: echo "sha_short=$(git rev-parse --short HEAD)" >> "$GITHUB_OUTPUT" + + - name: Build and push + uses: docker/build-push-action@v6 + with: + context: . + push: true + platforms: linux/amd64,linux/arm64 + build-args: VERSION=${{ steps.version.outputs.sha_short }} + tags: | + impingo/liveforge:latest + impingo/liveforge:${{ steps.version.outputs.sha_short }} + cache-from: type=gha + cache-to: type=gha,mode=max diff --git a/.github/workflows/release.yml b/.github/workflows/release.yml new file mode 100644 index 0000000..f190fc7 --- /dev/null +++ b/.github/workflows/release.yml @@ -0,0 +1,68 @@ +name: Release + +on: + push: + tags: ["v*"] + +permissions: + contents: write + packages: write + +jobs: + release: + name: Release + runs-on: ubuntu-latest + steps: + - uses: actions/checkout@v4 + with: + fetch-depth: 0 + + - uses: actions/setup-go@v5 + with: + go-version-file: go.mod + + - name: Set up Docker Buildx + uses: docker/setup-buildx-action@v3 + + - name: Login to Docker Hub + uses: docker/login-action@v3 + with: + username: ${{ secrets.DOCKERHUB_USERNAME }} + password: ${{ secrets.DOCKERHUB_TOKEN }} + + - name: Extract tag + id: tag + run: echo "version=${GITHUB_REF#refs/tags/}" >> "$GITHUB_OUTPUT" + + - name: Build and push Docker image + uses: docker/build-push-action@v6 + with: + context: . + push: true + platforms: linux/amd64,linux/arm64 + build-args: VERSION=${{ steps.tag.outputs.version }} + tags: | + impingo/liveforge:${{ steps.tag.outputs.version }} + impingo/liveforge:latest + cache-from: type=gha + cache-to: type=gha,mode=max + + - name: Build release binaries + run: | + mkdir -p dist + for os_arch in linux/amd64 linux/arm64 darwin/amd64 darwin/arm64; do + os="${os_arch%/*}" + arch="${os_arch#*/}" + output="dist/liveforge-${os}-${arch}" + [ "$os" = "linux" ] && cgo=1 || cgo=0 + CGO_ENABLED=$cgo GOOS=$os GOARCH=$arch \ + go build -trimpath \ + -ldflags "-s -w -X main.version=${{ steps.tag.outputs.version }}" \ + -o "$output" ./cmd/liveforge + done + + - name: Create GitHub Release + uses: softprops/action-gh-release@v2 + with: + generate_release_notes: true + files: dist/* diff --git a/.golangci.yml b/.golangci.yml new file mode 100644 index 0000000..fe8a1a5 --- /dev/null +++ b/.golangci.yml @@ -0,0 +1,48 @@ +run: + timeout: 5m + go: "1.26" + +linters: + enable: + - errcheck + - govet + - staticcheck + - unused + - ineffassign + - gosimple + - typecheck + - bodyclose + - durationcheck + - errname + - errorlint + - exportloopref + - gosec + - makezero + - nilerr + - prealloc + - unconvert + - unparam + - wastedassign + +linters-settings: + govet: + enable-all: true + disable: + - fieldalignment + gosec: + excludes: + - G104 # unhandled errors (too noisy for streaming server) + - G304 # file path from variable (expected for config loading) + errcheck: + exclude-functions: + - (net.Conn).Close + - (io.Closer).Close + - (*os.File).Close + +issues: + exclude-dirs: + - vendor + - third_party + - tools + max-issues-per-linter: 50 + max-same-issues: 5 diff --git a/cmd/liveforge/main.go b/cmd/liveforge/main.go index de4b7c8..7210877 100644 --- a/cmd/liveforge/main.go +++ b/cmd/liveforge/main.go @@ -23,6 +23,8 @@ import ( gb28181mod "github.com/im-pingo/liveforge/module/gb28181" metricsmod "github.com/im-pingo/liveforge/module/metrics" sipmod "github.com/im-pingo/liveforge/module/sip" + sipgwmod "github.com/im-pingo/liveforge/module/sipgateway" + dvrmod "github.com/im-pingo/liveforge/module/dvr" srtmod "github.com/im-pingo/liveforge/module/srt" webrtcmod "github.com/im-pingo/liveforge/module/webrtc" ) @@ -92,6 +94,13 @@ func main() { s.RegisterModule(gb28181mod.NewModule(sipModule.Service())) } + if cfg.SIP.Gateway.Enabled { + if sipModule == nil { + log.Fatal("sip gateway requires sip to be enabled") + } + s.RegisterModule(sipgwmod.NewModule(sipModule.Service())) + } + // Notify must be registered before API so its WebSocket handler // is available when the API module registers routes. if cfg.Notify.HTTP.Enabled || cfg.Notify.WebSocket.Enabled { @@ -112,6 +121,10 @@ func main() { s.RegisterModule(record.NewModule()) } + if cfg.DVR.Enabled { + s.RegisterModule(dvrmod.NewModule()) + } + if cfg.Metrics.Enabled { s.RegisterModule(metricsmod.NewModule()) } @@ -124,9 +137,24 @@ func main() { // Block until signal sigCh := make(chan os.Signal, 1) - signal.Notify(sigCh, syscall.SIGINT, syscall.SIGTERM) - sig := <-sigCh - slog.Info("shutting down", "signal", sig.String()) + signal.Notify(sigCh, syscall.SIGINT, syscall.SIGTERM, syscall.SIGHUP) + for { + sig := <-sigCh + if sig == syscall.SIGHUP { + slog.Info("received SIGHUP, reloading config", "path", *configPath) + newCfg, err := config.Load(*configPath) + if err != nil { + slog.Error("config reload failed", "error", err) + continue + } + logger.Init(newCfg.Server.LogLevel) + s.UpdateConfig(newCfg) + slog.Info("config reloaded successfully") + continue + } + slog.Info("shutting down", "signal", sig.String()) + break + } s.Shutdown() slog.Info("server stopped") diff --git a/config/config.go b/config/config.go index c6a855a..8c97a82 100644 --- a/config/config.go +++ b/config/config.go @@ -20,6 +20,7 @@ type Config struct { Notify NotifyConfig `yaml:"notify"` Cluster ClusterConfig `yaml:"cluster"` Record RecordConfig `yaml:"record"` + DVR DVRConfig `yaml:"dvr"` API APIConfig `yaml:"api"` Metrics MetricsConfig `yaml:"metrics"` AudioCodec AudioCodecConfig `yaml:"audio_codec"` @@ -79,10 +80,20 @@ type RTSPConfig struct { Enabled bool `yaml:"enabled"` Listen string `yaml:"listen"` RTPPortRange []int `yaml:"rtp_port_range"` + Multicast MulticastConfig `yaml:"multicast"` TLS *bool `yaml:"tls,omitempty"` // nil=follow global, true=force on, false=force off SkipTracker *SkipTrackerConfig `yaml:"skip_tracker,omitempty"` } +// MulticastConfig holds RTSP multicast delivery settings. +type MulticastConfig struct { + Enabled bool `yaml:"enabled"` + Address string `yaml:"address"` // multicast group IP (e.g., "239.0.0.1") + BasePort int `yaml:"base_port"` // starting port for multicast RTP (even number) + TTL int `yaml:"ttl"` // multicast TTL (default 16) + Interface string `yaml:"interface"` // network interface name (empty = default route) +} + // HTTPConfig holds HTTP-FLV/TS/FMP4/HLS/DASH module settings. type HTTPConfig struct { Enabled bool `yaml:"enabled"` @@ -160,12 +171,13 @@ type SRTConfig struct { // SIPConfig holds SIP module settings. type SIPConfig struct { - Enabled bool `yaml:"enabled"` - Listen string `yaml:"listen"` - Transport []string `yaml:"transport"` - ServerID string `yaml:"server_id"` - Domain string `yaml:"domain"` - Auth SIPAuth `yaml:"auth"` + Enabled bool `yaml:"enabled"` + Listen string `yaml:"listen"` + Transport []string `yaml:"transport"` + ServerID string `yaml:"server_id"` + Domain string `yaml:"domain"` + Auth SIPAuth `yaml:"auth"` + Gateway SIPGatewayConfig `yaml:"gateway"` } // SIPAuth holds SIP digest authentication settings. @@ -174,6 +186,15 @@ type SIPAuth struct { Password string `yaml:"password"` } +// SIPGatewayConfig holds SIP-to-stream gateway settings. +type SIPGatewayConfig struct { + Enabled bool `yaml:"enabled"` + StreamPrefix string `yaml:"stream_prefix"` // stream key prefix (default "sip") + RTPPortRange []int `yaml:"rtp_port_range"` // [min, max] for RTP port allocation + Codecs []string `yaml:"codecs"` // preferred codecs (default: opus, PCMA, PCMU) + MaxCalls int `yaml:"max_calls"` // max concurrent calls (default 100) +} + // GB28181Config holds GB28181 module settings. type GB28181Config struct { Enabled bool `yaml:"enabled"` @@ -317,12 +338,27 @@ type NotifyWSConfig struct { // ClusterConfig holds cluster settings. type ClusterConfig struct { - Forward ForwardConfig `yaml:"forward"` - Origin OriginConfig `yaml:"origin"` - SRT ClusterSRTConfig `yaml:"srt"` - RTSP ClusterRTSPConfig `yaml:"rtsp"` - RTP ClusterRTPConfig `yaml:"rtp"` - GB28181 ClusterGBConfig `yaml:"gb28181"` + Forward ForwardConfig `yaml:"forward"` + Origin OriginConfig `yaml:"origin"` + HealthCheck HealthCheckConfig `yaml:"health_check"` + RelayPool RelayPoolConfig `yaml:"relay_pool"` + SRT ClusterSRTConfig `yaml:"srt"` + RTSP ClusterRTSPConfig `yaml:"rtsp"` + RTP ClusterRTPConfig `yaml:"rtp"` + GB28181 ClusterGBConfig `yaml:"gb28181"` +} + +// HealthCheckConfig holds cluster node health monitoring settings. +type HealthCheckConfig struct { + Enabled bool `yaml:"enabled"` + Interval time.Duration `yaml:"interval"` // probe interval for evicted nodes + Timeout time.Duration `yaml:"timeout"` // TCP dial timeout per probe + EvictThreshold int `yaml:"evict_threshold"` // consecutive failures before eviction +} + +// RelayPoolConfig holds cluster relay connection pool settings. +type RelayPoolConfig struct { + MaxPerHost int `yaml:"max_per_host"` // max concurrent relay connections per peer host } // ClusterGBConfig holds GB28181 cluster relay settings. @@ -398,6 +434,17 @@ type FileCompleteConfig struct { URL string `yaml:"url"` } +// DVRConfig holds DVR/time-shift playback settings. +type DVRConfig struct { + Enabled bool `yaml:"enabled"` + Listen string `yaml:"listen"` + StreamPattern string `yaml:"stream_pattern"` + Path string `yaml:"path"` + Window time.Duration `yaml:"window"` + SegmentDuration time.Duration `yaml:"segment_duration"` + CleanupInterval time.Duration `yaml:"cleanup_interval"` +} + // MetricsConfig holds Prometheus metrics settings. type MetricsConfig struct { Enabled bool `yaml:"enabled"` diff --git a/config/loader.go b/config/loader.go index 4b72b6d..527bcf1 100644 --- a/config/loader.go +++ b/config/loader.go @@ -110,6 +110,13 @@ func defaults() *Config { Listen: ":9090", Path: "/metrics", }, + DVR: DVRConfig{ + Listen: ":8070", + Path: "./dvr/{stream_key}", + Window: 2 * time.Hour, + SegmentDuration: 6 * time.Second, + CleanupInterval: 30 * time.Second, + }, Cluster: ClusterConfig{ SRT: ClusterSRTConfig{ Latency: 120 * time.Millisecond, diff --git a/configs/liveforge.yaml b/configs/liveforge.yaml index b4d9448..5b2f115 100644 --- a/configs/liveforge.yaml +++ b/configs/liveforge.yaml @@ -30,6 +30,12 @@ rtsp: enabled: true listen: ":8554" rtp_port_range: [10000, 20000] + multicast: + enabled: false + address: "239.0.0.1" # multicast group IP + base_port: 40000 # starting RTP port for multicast streams + ttl: 16 # multicast TTL + interface: "" # network interface (empty = default) # skip_tracker: # max_count: 3 # window: 60s @@ -61,6 +67,13 @@ webrtc: ice_lite: true ice_servers: - urls: ["stun:stun.l.google.com:19302"] + # TURN server example (uncomment for NAT traversal): + # - urls: ["turn:turn.example.com:3478"] + # username: "user" + # credential: "pass" + # - urls: ["turns:turn.example.com:5349"] + # username: "user" + # credential: "pass" udp_port_range: [20000, 30000] candidates: [] gcc: @@ -211,6 +224,15 @@ record: on_file_complete: url: "" +dvr: + enabled: false + listen: ":8070" + stream_pattern: "live/*" + path: "./dvr/{stream_key}" + window: 2h + segment_duration: 6s + cleanup_interval: 30s + metrics: enabled: false listen: ":9090" diff --git a/core/module.go b/core/module.go index 9c46dc5..349b297 100644 --- a/core/module.go +++ b/core/module.go @@ -58,3 +58,10 @@ type Module interface { Hooks() []HookRegistration Close() error } + +// Reloadable is an optional interface modules can implement to support +// config hot-reload via SIGHUP. Only modules whose config has actually +// changed need to do anything in OnReload. +type Reloadable interface { + OnReload(s *Server) error +} diff --git a/core/server.go b/core/server.go index f3c4ee6..a8faac2 100644 --- a/core/server.go +++ b/core/server.go @@ -25,7 +25,7 @@ var Version = "dev" // Server is the main application server that manages modules and lifecycle. type Server struct { - config *config.Config + configPtr atomic.Pointer[config.Config] eventBus *EventBus hub *StreamHub modules []Module @@ -43,19 +43,34 @@ type Server struct { // NewServer creates a new Server instance. func NewServer(cfg *config.Config) *Server { bus := NewEventBus() - return &Server{ - config: cfg, + s := &Server{ eventBus: bus, hub: NewStreamHub(cfg.Stream, cfg.Limits, bus), startTime: time.Now(), done: make(chan struct{}), apiHandlers: make(map[string]http.Handler), } + s.configPtr.Store(cfg) + return s } // Config returns the server configuration. func (s *Server) Config() *config.Config { - return s.config + return s.configPtr.Load() +} + +// UpdateConfig atomically swaps the server configuration and notifies all +// Reloadable modules. Errors from individual modules are logged but do not +// stop the reload process. +func (s *Server) UpdateConfig(cfg *config.Config) { + s.configPtr.Store(cfg) + for _, m := range s.modules { + if r, ok := m.(Reloadable); ok { + if err := r.OnReload(s); err != nil { + slog.Error("module reload failed", "module", m.Name(), "error", err) + } + } + } } // GetEventBus returns the server's event bus. @@ -116,6 +131,16 @@ func (s *Server) ModuleNames() []string { return names } +// ModuleByName returns the module with the given name, or nil if not found. +func (s *Server) ModuleByName(name string) Module { + for _, m := range s.modules { + if m.Name() == name { + return m + } + } + return nil +} + // RegisterAPIHandler registers an HTTP handler for the given pattern on the API mux. // Modules call this during Init to expose HTTP/WebSocket endpoints on the API server. func (s *Server) RegisterAPIHandler(pattern string, h http.Handler) { @@ -137,7 +162,7 @@ func (s *Server) APIHandlers() map[string]http.Handler { // AcquireConn increments the connection counter. Returns false if max_connections is exceeded. func (s *Server) AcquireConn() bool { - max := s.config.Limits.MaxConnections + max := s.Config().Limits.MaxConnections if max > 0 { if s.connCount.Load() >= int64(max) { return false @@ -164,16 +189,16 @@ func (s *Server) ConnectionCount() int64 { // - true → force TLS on (error if global cert/key not configured) // - false → force TLS off (plain TCP even if global cert/key are configured) func (s *Server) MakeListener(addr string, moduleTLS *bool) (net.Listener, error) { - useTLS := s.config.TLS.Configured() // default: follow global + useTLS := s.Config().TLS.Configured() // default: follow global if moduleTLS != nil { useTLS = *moduleTLS } if useTLS { - if !s.config.TLS.Configured() { + if !s.Config().TLS.Configured() { return nil, fmt.Errorf("TLS enabled but tls.cert_file and tls.key_file are not configured") } - cert, err := tls.LoadX509KeyPair(s.config.TLS.CertFile, s.config.TLS.KeyFile) + cert, err := tls.LoadX509KeyPair(s.Config().TLS.CertFile, s.Config().TLS.KeyFile) if err != nil { return nil, fmt.Errorf("failed to load TLS certificate: %w", err) } @@ -202,8 +227,8 @@ func (s *Server) MakeListenerAutoTLS(addr string, moduleTLS *bool) (net.Listener } // If file-based TLS is configured, use it. - if s.config.TLS.Configured() { - cert, err := tls.LoadX509KeyPair(s.config.TLS.CertFile, s.config.TLS.KeyFile) + if s.Config().TLS.Configured() { + cert, err := tls.LoadX509KeyPair(s.Config().TLS.CertFile, s.Config().TLS.KeyFile) if err != nil { return nil, fmt.Errorf("failed to load TLS certificate: %w", err) } @@ -215,7 +240,7 @@ func (s *Server) MakeListenerAutoTLS(addr string, moduleTLS *bool) (net.Listener } // Auto-generate self-signed cert only when tls.auto is enabled. - if s.config.TLS.Auto { + if s.Config().TLS.Auto { autoCert := s.getOrCreateAutoCert() if autoCert == nil { return nil, fmt.Errorf("failed to generate self-signed TLS certificate") @@ -233,10 +258,10 @@ func (s *Server) MakeListenerAutoTLS(addr string, moduleTLS *bool) (net.Listener // HasTLS returns true if TLS is available (either file-based or auto-generated). func (s *Server) HasTLS() bool { - if s.config.TLS.Configured() { + if s.Config().TLS.Configured() { return true } - return s.config.TLS.Auto && s.getOrCreateAutoCert() != nil + return s.Config().TLS.Auto && s.getOrCreateAutoCert() != nil } // AutoCertPEM returns the auto-generated certificate in PEM format, or nil @@ -316,7 +341,7 @@ func generateSelfSignedCert() (*tls.Certificate, error) { // aliveLoop periodically emits alive events for all active streams. func (s *Server) aliveLoop() { - interval := s.config.Notify.AliveInterval + interval := s.Config().Notify.AliveInterval if interval <= 0 { interval = 10 * time.Second } diff --git a/core/server_test.go b/core/server_test.go index 0c2d186..0c457c0 100644 --- a/core/server_test.go +++ b/core/server_test.go @@ -369,3 +369,38 @@ func TestTLSConfigConfigured(t *testing.T) { }) } } + +type reloadableMockModule struct { + mockModule + reloadCfgName string + reloadErr error +} + +func (m *reloadableMockModule) OnReload(s *Server) error { + m.reloadCfgName = s.Config().Server.Name + return m.reloadErr +} + +func TestServerUpdateConfig(t *testing.T) { + cfg := &config.Config{} + cfg.Server.Name = "original" + s := NewServer(cfg) + + rm := &reloadableMockModule{mockModule: mockModule{name: "reloadable"}} + plain := &mockModule{name: "plain"} + s.RegisterModule(rm) + s.RegisterModule(plain) + _ = s.Init() + defer s.Shutdown() + + newCfg := &config.Config{} + newCfg.Server.Name = "reloaded" + s.UpdateConfig(newCfg) + + if s.Config().Server.Name != "reloaded" { + t.Errorf("expected config name 'reloaded', got %q", s.Config().Server.Name) + } + if rm.reloadCfgName != "reloaded" { + t.Errorf("expected reloadable module to see 'reloaded', got %q", rm.reloadCfgName) + } +} diff --git a/module/api/console.html b/module/api/console.html index bcf566d..23e35c5 100644 --- a/module/api/console.html +++ b/module/api/console.html @@ -670,6 +670,15 @@

LiveForge Console

Uptime
-
+ + + +
+
@@ -1614,7 +1623,10 @@

Playback Playback Playback = ht.evictThreshold { + ns.Evicted = true + slog.Warn("cluster node evicted", "module", "cluster", + "host", host, "failures", ns.ConsecutiveFailures) + return true + } + return false +} + +// IsEvicted returns true if the host extracted from rawURL is currently evicted. +func (ht *HealthTracker) IsEvicted(rawURL string) bool { + host := extractHost(rawURL) + if host == "" { + return false + } + ht.mu.RLock() + defer ht.mu.RUnlock() + if ns, ok := ht.nodes[host]; ok { + return ns.Evicted + } + return false +} + +// FilterHealthy returns only the URLs whose hosts are not evicted. +func (ht *HealthTracker) FilterHealthy(urls []string) []string { + ht.mu.RLock() + defer ht.mu.RUnlock() + + healthy := make([]string, 0, len(urls)) + for _, u := range urls { + host := extractHost(u) + if host == "" { + healthy = append(healthy, u) + continue + } + if ns, ok := ht.nodes[host]; ok && ns.Evicted { + continue + } + healthy = append(healthy, u) + } + return healthy +} + +// Snapshot returns a copy of all tracked nodes (for API/diagnostics). +func (ht *HealthTracker) Snapshot() map[string]NodeStatus { + ht.mu.RLock() + defer ht.mu.RUnlock() + out := make(map[string]NodeStatus, len(ht.nodes)) + for k, v := range ht.nodes { + out[k] = *v + } + return out +} + +// Close stops the background recovery loop. +func (ht *HealthTracker) Close() { + select { + case <-ht.closed: + default: + close(ht.closed) + } +} + +// recoveryLoop periodically probes evicted nodes via TCP dial. +func (ht *HealthTracker) recoveryLoop() { + ticker := time.NewTicker(ht.probeInterval) + defer ticker.Stop() + + for { + select { + case <-ht.closed: + return + case <-ticker.C: + ht.probeEvicted() + } + } +} + +func (ht *HealthTracker) probeEvicted() { + ht.mu.RLock() + var evicted []string + for host, ns := range ht.nodes { + if ns.Evicted { + evicted = append(evicted, host) + } + } + ht.mu.RUnlock() + + for _, host := range evicted { + conn, err := net.DialTimeout("tcp", host, ht.probeTimeout) + if err != nil { + continue + } + conn.Close() + + ht.mu.Lock() + if ns, ok := ht.nodes[host]; ok && ns.Evicted { + ns.Evicted = false + ns.ConsecutiveFailures = 0 + ns.LastSuccess = time.Now() + slog.Info("cluster node recovered via probe", "module", "cluster", "host", host) + } + ht.mu.Unlock() + } +} + +// extractHost returns "host:port" from a URL string. +func extractHost(rawURL string) string { + u, err := url.Parse(rawURL) + if err != nil || u.Host == "" { + return "" + } + return u.Host +} diff --git a/module/cluster/health_test.go b/module/cluster/health_test.go new file mode 100644 index 0000000..815015b --- /dev/null +++ b/module/cluster/health_test.go @@ -0,0 +1,169 @@ +package cluster + +import ( + "testing" + "time" + + "github.com/im-pingo/liveforge/config" +) + +func TestHealthTrackerRecordSuccessClearsFailures(t *testing.T) { + ht := NewHealthTracker(config.HealthCheckConfig{ + Enabled: true, + EvictThreshold: 3, + Interval: time.Hour, + Timeout: time.Second, + }) + defer ht.Close() + + ht.RecordFailure("rtmp://node1:1935/live") + ht.RecordFailure("rtmp://node1:1935/live") + ht.RecordSuccess("rtmp://node1:1935/live") + + snap := ht.Snapshot() + ns := snap["node1:1935"] + if ns.ConsecutiveFailures != 0 { + t.Errorf("ConsecutiveFailures = %d, want 0", ns.ConsecutiveFailures) + } + if ns.Evicted { + t.Error("should not be evicted after success") + } +} + +func TestHealthTrackerEviction(t *testing.T) { + ht := NewHealthTracker(config.HealthCheckConfig{ + Enabled: true, + EvictThreshold: 2, + Interval: time.Hour, + Timeout: time.Second, + }) + defer ht.Close() + + evicted := ht.RecordFailure("rtmp://node1:1935/live") + if evicted { + t.Error("should not evict after 1 failure") + } + + evicted = ht.RecordFailure("rtmp://node1:1935/live") + if !evicted { + t.Error("should evict after 2 failures (threshold=2)") + } + + if !ht.IsEvicted("rtmp://node1:1935/live") { + t.Error("IsEvicted should return true") + } +} + +func TestHealthTrackerFilterHealthy(t *testing.T) { + ht := NewHealthTracker(config.HealthCheckConfig{ + Enabled: true, + EvictThreshold: 1, + Interval: time.Hour, + Timeout: time.Second, + }) + defer ht.Close() + + ht.RecordFailure("rtmp://bad:1935/live") + + urls := []string{ + "rtmp://good:1935/live", + "rtmp://bad:1935/live", + "rtmp://also-good:1935/live", + } + + healthy := ht.FilterHealthy(urls) + if len(healthy) != 2 { + t.Fatalf("FilterHealthy = %d URLs, want 2", len(healthy)) + } + if healthy[0] != "rtmp://good:1935/live" || healthy[1] != "rtmp://also-good:1935/live" { + t.Errorf("unexpected healthy list: %v", healthy) + } +} + +func TestHealthTrackerSuccessRecovery(t *testing.T) { + ht := NewHealthTracker(config.HealthCheckConfig{ + Enabled: true, + EvictThreshold: 1, + Interval: time.Hour, + Timeout: time.Second, + }) + defer ht.Close() + + ht.RecordFailure("rtmp://node1:1935/live") + if !ht.IsEvicted("rtmp://node1:1935/live") { + t.Fatal("should be evicted") + } + + ht.RecordSuccess("rtmp://node1:1935/live") + if ht.IsEvicted("rtmp://node1:1935/live") { + t.Error("should recover after success") + } +} + +func TestHealthTrackerExtractHost(t *testing.T) { + tests := []struct { + url string + want string + }{ + {"rtmp://host:1935/live/stream", "host:1935"}, + {"srt://192.168.1.1:6000", "192.168.1.1:6000"}, + {"rtsp://node:554/live/test", "node:554"}, + {"invalid-url", ""}, + } + + for _, tt := range tests { + got := extractHost(tt.url) + if got != tt.want { + t.Errorf("extractHost(%q) = %q, want %q", tt.url, got, tt.want) + } + } +} + +func TestHealthTrackerUnknownURLNotEvicted(t *testing.T) { + ht := NewHealthTracker(config.HealthCheckConfig{ + Enabled: true, + EvictThreshold: 3, + Interval: time.Hour, + Timeout: time.Second, + }) + defer ht.Close() + + if ht.IsEvicted("rtmp://never-seen:1935/live") { + t.Error("unknown host should not be evicted") + } +} + +func TestHealthTrackerSnapshot(t *testing.T) { + ht := NewHealthTracker(config.HealthCheckConfig{ + Enabled: true, + EvictThreshold: 3, + Interval: time.Hour, + Timeout: time.Second, + }) + defer ht.Close() + + ht.RecordSuccess("rtmp://node1:1935/live") + ht.RecordFailure("rtmp://node2:1935/live") + + snap := ht.Snapshot() + if len(snap) != 2 { + t.Fatalf("Snapshot has %d entries, want 2", len(snap)) + } + if snap["node1:1935"].ConsecutiveFailures != 0 { + t.Error("node1 should have 0 failures") + } + if snap["node2:1935"].ConsecutiveFailures != 1 { + t.Error("node2 should have 1 failure") + } +} + +func TestHealthTrackerCloseIdempotent(t *testing.T) { + ht := NewHealthTracker(config.HealthCheckConfig{ + Enabled: true, + EvictThreshold: 3, + Interval: time.Hour, + Timeout: time.Second, + }) + ht.Close() + ht.Close() // should not panic +} diff --git a/module/cluster/integration_test.go b/module/cluster/integration_test.go index 38c558d..163956a 100644 --- a/module/cluster/integration_test.go +++ b/module/cluster/integration_test.go @@ -157,7 +157,7 @@ func TestForwardToMockRTMPServer(t *testing.T) { Payload: []byte{0x12, 0x10}, }) - ft := NewForwardTarget("live/fwdtest", "rtmp://"+addr+"/live/fwdtest", stream, NewRTMPTransport(), 1, time.Second) + ft := NewForwardTarget("live/fwdtest", "rtmp://"+addr+"/live/fwdtest", stream, NewRTMPTransport(), nil, nil, 1, time.Second) done := make(chan struct{}) go func() { @@ -211,7 +211,7 @@ func TestOriginPullFromMockServer(t *testing.T) { hub, _ := newTestHub() stream, _ := hub.GetOrCreate("live/pulltest") - op := NewOriginPull("live/pulltest", []string{"rtmp://" + addr + "/live"}, stream, newTestRegistry(), 1, 2*time.Second, 5*time.Second) + op := NewOriginPull("live/pulltest", []string{"rtmp://" + addr + "/live"}, stream, newTestRegistry(), nil, nil, 1, 2*time.Second, 5*time.Second) done := make(chan struct{}) go func() { @@ -251,7 +251,7 @@ func TestForwardManagerMultiProtocol(t *testing.T) { "srt://127.0.0.1:19997/live/stream", } scheduler := NewScheduler("", targets, "", 0) - fm := NewForwardManager(hub, bus, scheduler, registry, 1, time.Millisecond) + fm := NewForwardManager(hub, bus, scheduler, registry, nil, nil, 1, time.Millisecond) stream, _ := hub.GetOrCreate("live/multitest") pub := &originPublisher{id: "test-pub", info: &avframe.MediaInfo{ diff --git a/module/cluster/module.go b/module/cluster/module.go index f99ec19..728d508 100644 --- a/module/cluster/module.go +++ b/module/cluster/module.go @@ -8,9 +8,11 @@ import ( // Module implements core.Module for cluster forwarding and origin pull. type Module struct { - forward *ForwardManager - origin *OriginManager - registry *TransportRegistry + forward *ForwardManager + origin *OriginManager + health *HealthTracker + relayPool *RelayPool + registry *TransportRegistry } // NewModule creates a new cluster module. @@ -34,6 +36,19 @@ func (m *Module) Init(s *core.Server) error { m.registry.Register(NewRTPTransport(cfg.RTP, s)) m.registry.Register(NewGBTransport(cfg.GB28181, s)) + if cfg.HealthCheck.Enabled { + m.health = NewHealthTracker(cfg.HealthCheck) + slog.Info("cluster health check enabled", "module", "cluster", + "evict_threshold", cfg.HealthCheck.EvictThreshold, + "interval", cfg.HealthCheck.Interval) + } + + if cfg.RelayPool.MaxPerHost > 0 { + m.relayPool = NewRelayPool(cfg.RelayPool.MaxPerHost) + slog.Info("cluster relay pool enabled", "module", "cluster", + "max_per_host", cfg.RelayPool.MaxPerHost) + } + if cfg.Forward.Enabled && (len(cfg.Forward.Targets) > 0 || cfg.Forward.ScheduleURL != "") { fwdScheduler := NewScheduler( cfg.Forward.ScheduleURL, @@ -45,6 +60,8 @@ func (m *Module) Init(s *core.Server) error { hub, bus, fwdScheduler, m.registry, + m.health, + m.relayPool, cfg.Forward.RetryMax, cfg.Forward.RetryInterval, ) @@ -64,6 +81,8 @@ func (m *Module) Init(s *core.Server) error { hub, bus, origScheduler, m.registry, + m.health, + m.relayPool, cfg.Origin.RetryMax, cfg.Origin.RetryDelay, cfg.Origin.IdleTimeout, @@ -96,6 +115,9 @@ func (m *Module) Close() error { if m.origin != nil { m.origin.Close() } + if m.health != nil { + m.health.Close() + } if m.registry != nil { m.registry.Close() } diff --git a/module/cluster/module_test.go b/module/cluster/module_test.go index 0cac137..d5ed570 100644 --- a/module/cluster/module_test.go +++ b/module/cluster/module_test.go @@ -148,7 +148,7 @@ func newTestRegistry() *TransportRegistry { func TestForwardManagerDefaults(t *testing.T) { hub, bus := newTestHub() - fm := NewForwardManager(hub, bus, NewScheduler("", []string{"rtmp://target/live/stream"}, "", 0), newTestRegistry(), 0, 0) + fm := NewForwardManager(hub, bus, NewScheduler("", []string{"rtmp://target/live/stream"}, "", 0), newTestRegistry(), nil, nil, 0, 0) if fm.retryMax != 3 { t.Errorf("retryMax = %d, want 3", fm.retryMax) @@ -163,7 +163,7 @@ func TestForwardManagerDefaults(t *testing.T) { func TestForwardManagerOnPublishNoStream(t *testing.T) { hub, bus := newTestHub() - fm := NewForwardManager(hub, bus, NewScheduler("", []string{"rtmp://target/live/stream"}, "", 0), newTestRegistry(), 1, time.Millisecond) + fm := NewForwardManager(hub, bus, NewScheduler("", []string{"rtmp://target/live/stream"}, "", 0), newTestRegistry(), nil, nil, 1, time.Millisecond) defer fm.Close() // Publish event for non-existent stream should not create targets @@ -178,7 +178,7 @@ func TestForwardManagerOnPublishNoStream(t *testing.T) { func TestForwardManagerOnPublishStop(t *testing.T) { hub, bus := newTestHub() - fm := NewForwardManager(hub, bus, NewScheduler("", []string{"rtmp://127.0.0.1:19999/live/stream"}, "", 0), newTestRegistry(), 1, time.Millisecond) + fm := NewForwardManager(hub, bus, NewScheduler("", []string{"rtmp://127.0.0.1:19999/live/stream"}, "", 0), newTestRegistry(), nil, nil, 1, time.Millisecond) stream, _ := hub.GetOrCreate("live/test") // Set a dummy publisher so the stream is in publishing state @@ -205,7 +205,7 @@ func TestForwardManagerOnPublishStop(t *testing.T) { func TestForwardManagerDuplicatePublish(t *testing.T) { hub, bus := newTestHub() - fm := NewForwardManager(hub, bus, NewScheduler("", []string{"rtmp://127.0.0.1:19999/live/stream"}, "", 0), newTestRegistry(), 1, time.Millisecond) + fm := NewForwardManager(hub, bus, NewScheduler("", []string{"rtmp://127.0.0.1:19999/live/stream"}, "", 0), newTestRegistry(), nil, nil, 1, time.Millisecond) defer fm.Close() stream, _ := hub.GetOrCreate("live/test") @@ -222,7 +222,7 @@ func TestForwardManagerDuplicatePublish(t *testing.T) { func TestForwardManagerClose(t *testing.T) { hub, bus := newTestHub() - fm := NewForwardManager(hub, bus, NewScheduler("", []string{"rtmp://127.0.0.1:19999/live/stream"}, "", 0), newTestRegistry(), 1, time.Millisecond) + fm := NewForwardManager(hub, bus, NewScheduler("", []string{"rtmp://127.0.0.1:19999/live/stream"}, "", 0), newTestRegistry(), nil, nil, 1, time.Millisecond) stream, _ := hub.GetOrCreate("live/test") pub := &originPublisher{id: "test", info: &avframe.MediaInfo{}} @@ -239,7 +239,7 @@ func TestForwardManagerClose(t *testing.T) { func TestOriginManagerDefaults(t *testing.T) { hub, bus := newTestHub() - om := NewOriginManager(hub, bus, NewScheduler("", []string{"rtmp://origin/live"}, "", 0), newTestRegistry(), 0, 0, 0) + om := NewOriginManager(hub, bus, NewScheduler("", []string{"rtmp://origin/live"}, "", 0), newTestRegistry(), nil, nil, 0, 0, 0) if om.retryMax != 3 { t.Errorf("retryMax = %d, want 3", om.retryMax) @@ -254,7 +254,7 @@ func TestOriginManagerDefaults(t *testing.T) { func TestOriginManagerOnSubscribeNoStream(t *testing.T) { hub, bus := newTestHub() - om := NewOriginManager(hub, bus, NewScheduler("", []string{"rtmp://127.0.0.1:19999/live"}, "", 0), newTestRegistry(), 1, time.Second, time.Second) + om := NewOriginManager(hub, bus, NewScheduler("", []string{"rtmp://127.0.0.1:19999/live"}, "", 0), newTestRegistry(), nil, nil, 1, time.Second, time.Second) defer om.Close() err := om.onSubscribe(&core.EventContext{StreamKey: "nonexistent/stream"}) @@ -268,7 +268,7 @@ func TestOriginManagerOnSubscribeNoStream(t *testing.T) { func TestOriginManagerOnSubscribeWithPublisher(t *testing.T) { hub, bus := newTestHub() - om := NewOriginManager(hub, bus, NewScheduler("", []string{"rtmp://127.0.0.1:19999/live"}, "", 0), newTestRegistry(), 1, time.Second, time.Second) + om := NewOriginManager(hub, bus, NewScheduler("", []string{"rtmp://127.0.0.1:19999/live"}, "", 0), newTestRegistry(), nil, nil, 1, time.Second, time.Second) defer om.Close() stream, _ := hub.GetOrCreate("live/test") @@ -284,7 +284,7 @@ func TestOriginManagerOnSubscribeWithPublisher(t *testing.T) { func TestOriginManagerOnSubscribeTriggersPull(t *testing.T) { hub, bus := newTestHub() - om := NewOriginManager(hub, bus, NewScheduler("", []string{"rtmp://127.0.0.1:19999/live"}, "", 0), newTestRegistry(), 1, time.Second, time.Second) + om := NewOriginManager(hub, bus, NewScheduler("", []string{"rtmp://127.0.0.1:19999/live"}, "", 0), newTestRegistry(), nil, nil, 1, time.Second, time.Second) // Create stream without publisher hub.GetOrCreate("live/test") @@ -303,7 +303,7 @@ func TestOriginManagerOnSubscribeTriggersPull(t *testing.T) { func TestOriginManagerDuplicateSubscribe(t *testing.T) { hub, bus := newTestHub() - om := NewOriginManager(hub, bus, NewScheduler("", []string{"rtmp://127.0.0.1:19999/live"}, "", 0), newTestRegistry(), 1, time.Second, time.Second) + om := NewOriginManager(hub, bus, NewScheduler("", []string{"rtmp://127.0.0.1:19999/live"}, "", 0), newTestRegistry(), nil, nil, 1, time.Second, time.Second) defer om.Close() hub.GetOrCreate("live/test") @@ -318,7 +318,7 @@ func TestOriginManagerDuplicateSubscribe(t *testing.T) { func TestOriginManagerClose(t *testing.T) { hub, bus := newTestHub() - om := NewOriginManager(hub, bus, NewScheduler("", []string{"rtmp://127.0.0.1:19999/live"}, "", 0), newTestRegistry(), 1, time.Second, time.Second) + om := NewOriginManager(hub, bus, NewScheduler("", []string{"rtmp://127.0.0.1:19999/live"}, "", 0), newTestRegistry(), nil, nil, 1, time.Second, time.Second) hub.GetOrCreate("live/test") om.onSubscribe(&core.EventContext{StreamKey: "live/test"}) @@ -353,7 +353,7 @@ func TestForwardTargetClose(t *testing.T) { stream, _ := hub.GetOrCreate("live/test") _ = bus - ft := NewForwardTarget("live/test", "rtmp://127.0.0.1:19999/live/test", stream, NewRTMPTransport(), 1, time.Millisecond) + ft := NewForwardTarget("live/test", "rtmp://127.0.0.1:19999/live/test", stream, NewRTMPTransport(), nil, nil, 1, time.Millisecond) // Close before Run ft.Close() @@ -366,7 +366,7 @@ func TestOriginPullClose(t *testing.T) { hub, _ := newTestHub() stream, _ := hub.GetOrCreate("live/test") - op := NewOriginPull("live/test", []string{"rtmp://127.0.0.1:19999/live"}, stream, newTestRegistry(), 1, time.Second, time.Second) + op := NewOriginPull("live/test", []string{"rtmp://127.0.0.1:19999/live"}, stream, newTestRegistry(), nil, nil, 1, time.Second, time.Second) // Close before Run op.Close() @@ -379,7 +379,7 @@ func TestForwardTargetRunWithClosedTarget(t *testing.T) { hub, _ := newTestHub() stream, _ := hub.GetOrCreate("live/test") - ft := NewForwardTarget("live/test", "rtmp://127.0.0.1:19999/live/test", stream, NewRTMPTransport(), 1, time.Millisecond) + ft := NewForwardTarget("live/test", "rtmp://127.0.0.1:19999/live/test", stream, NewRTMPTransport(), nil, nil, 1, time.Millisecond) ft.Close() // Run should return immediately when already closed @@ -400,7 +400,7 @@ func TestOriginPullRunWithClosedPull(t *testing.T) { hub, _ := newTestHub() stream, _ := hub.GetOrCreate("live/test") - op := NewOriginPull("live/test", []string{"rtmp://127.0.0.1:19999/live"}, stream, newTestRegistry(), 1, time.Second, time.Second) + op := NewOriginPull("live/test", []string{"rtmp://127.0.0.1:19999/live"}, stream, newTestRegistry(), nil, nil, 1, time.Second, time.Second) op.Close() done := make(chan struct{}) diff --git a/module/cluster/origin.go b/module/cluster/origin.go index d0b2677..b20039b 100644 --- a/module/cluster/origin.go +++ b/module/cluster/origin.go @@ -17,6 +17,8 @@ type OriginPull struct { servers []string stream *core.Stream registry *TransportRegistry + health *HealthTracker + pool *RelayPool retryMax int retryDelay time.Duration idleTimeout time.Duration @@ -26,12 +28,14 @@ type OriginPull struct { } // NewOriginPull creates a new origin pull instance. -func NewOriginPull(streamKey string, servers []string, stream *core.Stream, registry *TransportRegistry, retryMax int, retryDelay, idleTimeout time.Duration) *OriginPull { +func NewOriginPull(streamKey string, servers []string, stream *core.Stream, registry *TransportRegistry, health *HealthTracker, pool *RelayPool, retryMax int, retryDelay, idleTimeout time.Duration) *OriginPull { return &OriginPull{ streamKey: streamKey, servers: servers, stream: stream, registry: registry, + health: health, + pool: pool, retryMax: retryMax, retryDelay: retryDelay, idleTimeout: idleTimeout, @@ -114,7 +118,25 @@ func (op *OriginPull) pullOnce(sourceURL string) error { } }() - return transport.Pull(ctx, sourceURL, op.stream) + if op.pool != nil { + host := extractHost(sourceURL) + if err := op.pool.Acquire(ctx, host); err != nil { + return err + } + defer op.pool.Release(host) + } + + err = transport.Pull(ctx, sourceURL, op.stream) + if err != nil { + if op.health != nil { + op.health.RecordFailure(sourceURL) + } + return err + } + if op.health != nil { + op.health.RecordSuccess(sourceURL) + } + return nil } // Close stops the origin pull. @@ -144,6 +166,8 @@ type OriginManager struct { eventBus *core.EventBus scheduler *Scheduler registry *TransportRegistry + health *HealthTracker + pool *RelayPool retryMax int retryDelay time.Duration idleTimeout time.Duration @@ -154,7 +178,7 @@ type OriginManager struct { } // NewOriginManager creates a new origin manager. -func NewOriginManager(hub *core.StreamHub, bus *core.EventBus, scheduler *Scheduler, registry *TransportRegistry, retryMax int, retryDelay, idleTimeout time.Duration) *OriginManager { +func NewOriginManager(hub *core.StreamHub, bus *core.EventBus, scheduler *Scheduler, registry *TransportRegistry, health *HealthTracker, pool *RelayPool, retryMax int, retryDelay, idleTimeout time.Duration) *OriginManager { if retryMax <= 0 { retryMax = 3 } @@ -169,6 +193,8 @@ func NewOriginManager(hub *core.StreamHub, bus *core.EventBus, scheduler *Schedu eventBus: bus, scheduler: scheduler, registry: registry, + health: health, + pool: pool, retryMax: retryMax, retryDelay: retryDelay, idleTimeout: idleTimeout, @@ -215,7 +241,11 @@ func (om *OriginManager) onSubscribe(ctx *core.EventContext) error { return nil } - op := NewOriginPull(ctx.StreamKey, servers, stream, om.registry, om.retryMax, om.retryDelay, om.idleTimeout) + if om.health != nil { + servers = om.health.FilterHealthy(servers) + } + + op := NewOriginPull(ctx.StreamKey, servers, stream, om.registry, om.health, om.pool, om.retryMax, om.retryDelay, om.idleTimeout) om.active[ctx.StreamKey] = op om.eventBus.Emit(core.EventOriginPullStart, &core.EventContext{ //nolint:errcheck diff --git a/module/cluster/relay_pool.go b/module/cluster/relay_pool.go new file mode 100644 index 0000000..703fd06 --- /dev/null +++ b/module/cluster/relay_pool.go @@ -0,0 +1,70 @@ +package cluster + +import ( + "context" + "fmt" + "sync" +) + +// RelayPool limits concurrent relay connections per peer host. +// This prevents overwhelming a single cluster node when many streams +// need to be forwarded or pulled simultaneously. +type RelayPool struct { + maxPerHost int + mu sync.Mutex + hosts map[string]chan struct{} +} + +// NewRelayPool creates a relay pool with the given per-host concurrency limit. +func NewRelayPool(maxPerHost int) *RelayPool { + if maxPerHost <= 0 { + maxPerHost = 10 + } + return &RelayPool{ + maxPerHost: maxPerHost, + hosts: make(map[string]chan struct{}), + } +} + +// Acquire blocks until a relay slot is available for the given host, +// or the context is cancelled. +func (p *RelayPool) Acquire(ctx context.Context, host string) error { + sem := p.getSemaphore(host) + select { + case sem <- struct{}{}: + return nil + case <-ctx.Done(): + return fmt.Errorf("relay pool acquire %s: %w", host, ctx.Err()) + } +} + +// Release returns a relay slot for the given host. +func (p *RelayPool) Release(host string) { + sem := p.getSemaphore(host) + select { + case <-sem: + default: + } +} + +// ActiveCount returns the number of active relay connections for a host. +func (p *RelayPool) ActiveCount(host string) int { + p.mu.Lock() + sem, ok := p.hosts[host] + p.mu.Unlock() + if !ok { + return 0 + } + return len(sem) +} + +func (p *RelayPool) getSemaphore(host string) chan struct{} { + p.mu.Lock() + defer p.mu.Unlock() + sem, ok := p.hosts[host] + if !ok { + sem = make(chan struct{}, p.maxPerHost) + p.hosts[host] = sem + } + return sem +} diff --git a/module/cluster/relay_pool_test.go b/module/cluster/relay_pool_test.go new file mode 100644 index 0000000..f6cf7bf --- /dev/null +++ b/module/cluster/relay_pool_test.go @@ -0,0 +1,135 @@ +package cluster + +import ( + "context" + "sync" + "testing" + "time" +) + +func TestRelayPoolAcquireRelease(t *testing.T) { + pool := NewRelayPool(2) + + ctx := context.Background() + if err := pool.Acquire(ctx, "host1:1935"); err != nil { + t.Fatalf("Acquire 1: %v", err) + } + if err := pool.Acquire(ctx, "host1:1935"); err != nil { + t.Fatalf("Acquire 2: %v", err) + } + + if pool.ActiveCount("host1:1935") != 2 { + t.Errorf("ActiveCount = %d, want 2", pool.ActiveCount("host1:1935")) + } + + pool.Release("host1:1935") + if pool.ActiveCount("host1:1935") != 1 { + t.Errorf("ActiveCount after release = %d, want 1", pool.ActiveCount("host1:1935")) + } + + pool.Release("host1:1935") + if pool.ActiveCount("host1:1935") != 0 { + t.Errorf("ActiveCount after full release = %d, want 0", pool.ActiveCount("host1:1935")) + } +} + +func TestRelayPoolBlocksAtLimit(t *testing.T) { + pool := NewRelayPool(1) + + ctx := context.Background() + if err := pool.Acquire(ctx, "host1:1935"); err != nil { + t.Fatalf("Acquire: %v", err) + } + + ctx2, cancel := context.WithTimeout(context.Background(), 50*time.Millisecond) + defer cancel() + + err := pool.Acquire(ctx2, "host1:1935") + if err == nil { + t.Error("expected error when pool is full and context expires") + } +} + +func TestRelayPoolUnblocksOnRelease(t *testing.T) { + pool := NewRelayPool(1) + + ctx := context.Background() + pool.Acquire(ctx, "host1:1935") + + acquired := make(chan struct{}) + go func() { + pool.Acquire(context.Background(), "host1:1935") + close(acquired) + }() + + time.Sleep(20 * time.Millisecond) + pool.Release("host1:1935") + + select { + case <-acquired: + case <-time.After(time.Second): + t.Error("second Acquire did not unblock after Release") + } +} + +func TestRelayPoolIndependentHosts(t *testing.T) { + pool := NewRelayPool(1) + ctx := context.Background() + + if err := pool.Acquire(ctx, "host1:1935"); err != nil { + t.Fatalf("Acquire host1: %v", err) + } + if err := pool.Acquire(ctx, "host2:1935"); err != nil { + t.Fatalf("Acquire host2: %v", err) + } + + if pool.ActiveCount("host1:1935") != 1 { + t.Errorf("host1 ActiveCount = %d, want 1", pool.ActiveCount("host1:1935")) + } + if pool.ActiveCount("host2:1935") != 1 { + t.Errorf("host2 ActiveCount = %d, want 1", pool.ActiveCount("host2:1935")) + } +} + +func TestRelayPoolConcurrentAccess(t *testing.T) { + pool := NewRelayPool(10) + ctx := context.Background() + + var wg sync.WaitGroup + for i := 0; i < 20; i++ { + wg.Add(1) + go func() { + defer wg.Done() + pool.Acquire(ctx, "host:1935") + time.Sleep(time.Millisecond) + pool.Release("host:1935") + }() + } + wg.Wait() + + if pool.ActiveCount("host:1935") != 0 { + t.Errorf("ActiveCount = %d, want 0 after all released", pool.ActiveCount("host:1935")) + } +} + +func TestRelayPoolDefaultMaxPerHost(t *testing.T) { + pool := NewRelayPool(0) + if pool.maxPerHost != 10 { + t.Errorf("maxPerHost = %d, want 10 (default)", pool.maxPerHost) + } +} + +func TestRelayPoolUnknownHostActiveCount(t *testing.T) { + pool := NewRelayPool(5) + if pool.ActiveCount("unknown:1935") != 0 { + t.Errorf("ActiveCount for unknown host = %d, want 0", pool.ActiveCount("unknown:1935")) + } +} + +func TestRelayPoolReleaseWithoutAcquire(t *testing.T) { + pool := NewRelayPool(5) + pool.Release("host:1935") + if pool.ActiveCount("host:1935") != 0 { + t.Errorf("ActiveCount = %d, want 0", pool.ActiveCount("host:1935")) + } +} diff --git a/module/dvr/cleanup.go b/module/dvr/cleanup.go new file mode 100644 index 0000000..68fdc4a --- /dev/null +++ b/module/dvr/cleanup.go @@ -0,0 +1,61 @@ +package dvr + +import ( + "context" + "log/slog" + "time" +) + +func (m *Module) runCleanup(ctx context.Context) { + interval := m.cfg.CleanupInterval + if interval <= 0 { + interval = 30 * time.Second + } + + ticker := time.NewTicker(interval) + defer ticker.Stop() + + for { + select { + case <-ctx.Done(): + return + case <-ticker.C: + m.cleanExpiredSegments() + } + } +} + +func (m *Module) cleanExpiredSegments() { + cutoff := time.Now().Add(-m.cfg.Window) + + m.mu.Lock() + keys := make([]string, 0, len(m.sessions)) + for k := range m.sessions { + keys = append(keys, k) + } + m.mu.Unlock() + + for _, key := range keys { + m.mu.Lock() + session := m.sessions[key] + m.mu.Unlock() + + if session == nil { + continue + } + + removed := session.Index().CleanBefore(cutoff) + if len(removed) > 0 { + slog.Debug("dvr cleanup", "stream", key, "removed", len(removed)) + } + + if !session.IsLive() && session.Index().Len() == 0 { + m.mu.Lock() + if m.sessions[key] == session { + delete(m.sessions, key) + } + m.mu.Unlock() + slog.Info("dvr session expired", "module", "dvr", "stream", key) + } + } +} diff --git a/module/dvr/dvr_test.go b/module/dvr/dvr_test.go new file mode 100644 index 0000000..76e8326 --- /dev/null +++ b/module/dvr/dvr_test.go @@ -0,0 +1,374 @@ +package dvr + +import ( + "os" + "path/filepath" + "strings" + "testing" + "time" + + "github.com/im-pingo/liveforge/config" + "github.com/im-pingo/liveforge/core" + "github.com/im-pingo/liveforge/pkg/avframe" +) + +func TestSegmentIndex_AddAndLookup(t *testing.T) { + idx := NewSegmentIndex() + + now := time.Now() + for i := 0; i < 5; i++ { + idx.Add(Segment{ + SeqNum: i, + StartTime: now.Add(time.Duration(i) * 6 * time.Second), + Duration: 6.0, + Filename: segFilename(i), + DiskPath: "/tmp/test/" + segFilename(i), + }) + } + + if idx.Len() != 5 { + t.Fatalf("Len = %d, want 5", idx.Len()) + } + + seg, ok := idx.SegmentBySeqNum(3) + if !ok { + t.Fatal("SegmentBySeqNum(3) not found") + } + if seg.SeqNum != 3 { + t.Errorf("SeqNum = %d, want 3", seg.SeqNum) + } + + _, ok = idx.SegmentBySeqNum(99) + if ok { + t.Error("SegmentBySeqNum(99) should not be found") + } + + first, _ := idx.First() + if first.SeqNum != 0 { + t.Errorf("First.SeqNum = %d, want 0", first.SeqNum) + } + + last, _ := idx.Last() + if last.SeqNum != 4 { + t.Errorf("Last.SeqNum = %d, want 4", last.SeqNum) + } + + if idx.MaxDuration() != 6.0 { + t.Errorf("MaxDuration = %f, want 6.0", idx.MaxDuration()) + } +} + +func TestSegmentIndex_CleanBefore(t *testing.T) { + dir := t.TempDir() + idx := NewSegmentIndex() + + now := time.Now() + for i := 0; i < 5; i++ { + filename := segFilename(i) + path := filepath.Join(dir, filename) + os.WriteFile(path, []byte("test data"), 0644) + + idx.Add(Segment{ + SeqNum: i, + StartTime: now.Add(time.Duration(i-3) * time.Hour), + Duration: 6.0, + Filename: filename, + DiskPath: path, + }) + } + + // Clean segments older than 1 hour + cutoff := now.Add(-1 * time.Hour) + removed := idx.CleanBefore(cutoff) + + if len(removed) != 2 { + t.Fatalf("removed %d segments, want 2", len(removed)) + } + + if idx.Len() != 3 { + t.Errorf("remaining = %d, want 3", idx.Len()) + } + + // Verify files were deleted + for _, seg := range removed { + if _, err := os.Stat(seg.DiskPath); !os.IsNotExist(err) { + t.Errorf("file %s should have been deleted", seg.DiskPath) + } + } + + // Verify remaining files still exist + segs := idx.Segments() + for _, seg := range segs { + if _, err := os.Stat(seg.DiskPath); err != nil { + t.Errorf("file %s should still exist", seg.DiskPath) + } + } +} + +func TestSegmentIndex_EmptyIndex(t *testing.T) { + idx := NewSegmentIndex() + + if idx.Len() != 0 { + t.Error("expected empty index") + } + + _, ok := idx.First() + if ok { + t.Error("First on empty index should return false") + } + + _, ok = idx.Last() + if ok { + t.Error("Last on empty index should return false") + } + + if idx.MaxDuration() != 0 { + t.Error("MaxDuration on empty index should be 0") + } +} + +func TestGeneratePlaylist_Live(t *testing.T) { + idx := NewSegmentIndex() + now := time.Date(2026, 5, 25, 12, 0, 0, 0, time.UTC) + + for i := 0; i < 3; i++ { + idx.Add(Segment{ + SeqNum: i, + StartTime: now.Add(time.Duration(i) * 6 * time.Second), + Duration: 6.0, + Filename: segFilename(i), + }) + } + + playlist := GeneratePlaylist(idx, "live/mystream", true) + + if !strings.Contains(playlist, "#EXTM3U") { + t.Error("missing #EXTM3U") + } + if !strings.Contains(playlist, "#EXT-X-VERSION:6") { + t.Error("missing VERSION:6") + } + if !strings.Contains(playlist, "#EXT-X-MEDIA-SEQUENCE:0") { + t.Error("missing MEDIA-SEQUENCE") + } + if !strings.Contains(playlist, "#EXT-X-PROGRAM-DATE-TIME:2026-05-25T12:00:00.000Z") { + t.Error("missing PROGRAM-DATE-TIME for first segment") + } + if !strings.Contains(playlist, "#EXTINF:6.000,") { + t.Error("missing EXTINF") + } + if !strings.Contains(playlist, "mystream/seg_000000.ts") { + t.Error("missing segment URL") + } + if strings.Contains(playlist, "#EXT-X-ENDLIST") { + t.Error("live playlist should not have ENDLIST") + } +} + +func TestGeneratePlaylist_Stopped(t *testing.T) { + idx := NewSegmentIndex() + now := time.Date(2026, 5, 25, 12, 0, 0, 0, time.UTC) + + idx.Add(Segment{ + SeqNum: 5, + StartTime: now, + Duration: 4.5, + Filename: segFilename(5), + }) + + playlist := GeneratePlaylist(idx, "live/test", false) + + if !strings.Contains(playlist, "#EXT-X-ENDLIST") { + t.Error("stopped playlist should have ENDLIST") + } + if !strings.Contains(playlist, "#EXT-X-MEDIA-SEQUENCE:5") { + t.Error("media sequence should be 5") + } +} + +func TestGeneratePlaylist_Empty(t *testing.T) { + idx := NewSegmentIndex() + playlist := GeneratePlaylist(idx, "live/test", true) + if playlist != "" { + t.Errorf("expected empty playlist, got %q", playlist) + } +} + +func TestSession_WritesSegments(t *testing.T) { + dir := t.TempDir() + bus := core.NewEventBus() + hub := core.NewStreamHub(config.StreamConfig{RingBufferSize: 256}, config.LimitsConfig{}, bus) + + stream, _ := hub.GetOrCreate("live/dvr-test") + + cfg := config.DVRConfig{ + Enabled: true, + Path: filepath.Join(dir, "{stream_key}"), + Window: 2 * time.Hour, + SegmentDuration: 100 * time.Millisecond, + CleanupInterval: 30 * time.Second, + } + + session, err := NewSession("live/dvr-test", stream, cfg, nil, 0) + if err != nil { + t.Fatalf("NewSession: %v", err) + } + + // Feed frames in a goroutine + go func() { + // Sequence headers + stream.WriteFrame(&avframe.AVFrame{ + MediaType: avframe.MediaTypeVideo, + Codec: avframe.CodecH264, + FrameType: avframe.FrameTypeSequenceHeader, + DTS: 0, + PTS: 0, + Payload: []byte{0x67, 0x42, 0x00, 0x1e, 0x67, 0x42, 0x00, 0x1e, 0x68, 0xce, 0x38, 0x80}, + }) + + // Write 3 keyframes to trigger segment splits + for i := 0; i < 3; i++ { + dts := int64(i * 200) + stream.WriteFrame(&avframe.AVFrame{ + MediaType: avframe.MediaTypeVideo, + Codec: avframe.CodecH264, + FrameType: avframe.FrameTypeKeyframe, + DTS: dts, + PTS: dts, + Payload: []byte{0x65, 0x88, 0x84, 0x00, 0x33}, + }) + // Inter frame + stream.WriteFrame(&avframe.AVFrame{ + MediaType: avframe.MediaTypeVideo, + Codec: avframe.CodecH264, + FrameType: avframe.FrameTypeInterframe, + DTS: dts + 33, + PTS: dts + 33, + Payload: []byte{0x41, 0x9a, 0x24}, + }) + } + + // Give session time to process, then stop + time.Sleep(50 * time.Millisecond) + session.Stop() + }() + + session.Run() + + // Verify segments were created + segCount := session.Index().Len() + if segCount < 2 { + t.Errorf("expected at least 2 segments, got %d", segCount) + } + + // Verify segment files exist on disk + segs := session.Index().Segments() + for _, seg := range segs { + if _, err := os.Stat(seg.DiskPath); err != nil { + t.Errorf("segment file missing: %s", seg.DiskPath) + } + } +} + +func TestModuleHooks(t *testing.T) { + m := NewModule() + hooks := m.Hooks() + + if len(hooks) != 2 { + t.Fatalf("expected 2 hooks, got %d", len(hooks)) + } + + if hooks[0].Event != core.EventPublish { + t.Errorf("hook[0] event = %v, want EventPublish", hooks[0].Event) + } + if hooks[1].Event != core.EventPublishStop { + t.Errorf("hook[1] event = %v, want EventPublishStop", hooks[1].Event) + } +} + +func TestModuleName(t *testing.T) { + m := NewModule() + if m.Name() != "dvr" { + t.Errorf("Name = %q, want dvr", m.Name()) + } +} + +func TestResolvePath(t *testing.T) { + tests := []struct { + template string + streamKey string + want string + }{ + {"./dvr/{stream_key}", "live/test", "dvr/live/test"}, + {"/data/dvr/{stream_key}", "app/cam1", "/data/dvr/app/cam1"}, + } + + for _, tt := range tests { + got := resolvePath(tt.template, tt.streamKey) + if got != tt.want { + t.Errorf("resolvePath(%q, %q) = %q, want %q", tt.template, tt.streamKey, got, tt.want) + } + } +} + +func TestParseSeqNum(t *testing.T) { + tests := []struct { + filename string + want int + }{ + {"seg_000000.ts", 0}, + {"seg_000042.ts", 42}, + {"seg_001234.ts", 1234}, + {"bad_name.ts", -1}, + {"seg_.ts", 0}, + {"seg_abc.ts", -1}, + } + + for _, tt := range tests { + got := parseSeqNum(tt.filename) + if got != tt.want { + t.Errorf("parseSeqNum(%q) = %d, want %d", tt.filename, got, tt.want) + } + } +} + +func TestMatchPattern(t *testing.T) { + tests := []struct { + pattern string + key string + want bool + }{ + {"", "anything", true}, + {"*", "anything", true}, + {"live/*", "live/test", true}, + {"live/*", "vod/test", false}, + {"live/cam*", "live/cam1", true}, + } + + for _, tt := range tests { + got := matchPattern(tt.pattern, tt.key) + if got != tt.want { + t.Errorf("matchPattern(%q, %q) = %v, want %v", tt.pattern, tt.key, got, tt.want) + } + } +} + +func segFilename(seq int) string { + return "seg_" + padInt(seq) + ".ts" +} + +func padInt(n int) string { + s := "000000" + ns := "" + for n > 0 { + ns = string(rune('0'+n%10)) + ns + n /= 10 + } + if ns == "" { + ns = "0" + } + if len(ns) < 6 { + return s[:6-len(ns)] + ns + } + return ns +} diff --git a/module/dvr/handler.go b/module/dvr/handler.go new file mode 100644 index 0000000..7b86b30 --- /dev/null +++ b/module/dvr/handler.go @@ -0,0 +1,89 @@ +package dvr + +import ( + "net/http" + "path" + "strings" +) + +func (m *Module) handlePlaylist(w http.ResponseWriter, r *http.Request) { + app := r.PathValue("app") + key := r.PathValue("key") + streamKey := app + "/" + key + + m.mu.Lock() + session := m.sessions[streamKey] + m.mu.Unlock() + + if session == nil { + http.NotFound(w, r) + return + } + + playlist := GeneratePlaylist(session.Index(), streamKey, session.IsLive()) + if playlist == "" { + http.NotFound(w, r) + return + } + + w.Header().Set("Content-Type", "application/vnd.apple.mpegurl") + w.Header().Set("Cache-Control", "no-cache, no-store") + w.Write([]byte(playlist)) +} + +func (m *Module) handleSegment(w http.ResponseWriter, r *http.Request) { + app := r.PathValue("app") + key := r.PathValue("key") + filename := r.PathValue("filename") + streamKey := app + "/" + key + + m.mu.Lock() + session := m.sessions[streamKey] + m.mu.Unlock() + + if session == nil { + http.NotFound(w, r) + return + } + + seqNum := parseSeqNum(filename) + if seqNum < 0 { + http.NotFound(w, r) + return + } + + seg, ok := session.Index().SegmentBySeqNum(seqNum) + if !ok { + http.NotFound(w, r) + return + } + + w.Header().Set("Cache-Control", "public, max-age=3600") + http.ServeFile(w, r, seg.DiskPath) +} + +func parseSeqNum(filename string) int { + // Expected format: seg_000042.ts + name := strings.TrimSuffix(filename, ".ts") + if !strings.HasPrefix(name, "seg_") { + return -1 + } + numStr := name[4:] + n := 0 + for _, c := range numStr { + if c < '0' || c > '9' { + return -1 + } + n = n*10 + int(c-'0') + } + return n +} + +// matchPattern checks if a stream key matches a glob pattern. +func matchPattern(pattern, key string) bool { + if pattern == "" || pattern == "*" { + return true + } + matched, _ := path.Match(pattern, key) + return matched +} diff --git a/module/dvr/module.go b/module/dvr/module.go new file mode 100644 index 0000000..fc0d9b2 --- /dev/null +++ b/module/dvr/module.go @@ -0,0 +1,195 @@ +package dvr + +import ( + "context" + "log/slog" + "net" + "net/http" + "sync" + + "github.com/im-pingo/liveforge/config" + "github.com/im-pingo/liveforge/core" +) + +// Module implements core.Module for DVR/time-shift playback. +type Module struct { + server *core.Server + cfg config.DVRConfig + listener net.Listener + httpSrv *http.Server + cancel context.CancelFunc + wg sync.WaitGroup + mu sync.Mutex + sessions map[string]*Session +} + +// NewModule creates a new DVR module. +func NewModule() *Module { + return &Module{ + sessions: make(map[string]*Session), + } +} + +// Name returns the module name. +func (m *Module) Name() string { return "dvr" } + +// Init starts the DVR HTTP server and cleanup goroutine. +func (m *Module) Init(s *core.Server) error { + m.server = s + m.cfg = s.Config().DVR + + ln, err := s.MakeListenerAutoTLS(m.cfg.Listen, nil) + if err != nil { + return err + } + m.listener = ln + + mux := http.NewServeMux() + mux.HandleFunc("GET /dvr/{app}/{key}.m3u8", m.handlePlaylist) + mux.HandleFunc("GET /dvr/{app}/{key}/{filename}", m.handleSegment) + + m.httpSrv = &http.Server{Handler: mux} + + ctx, cancel := context.WithCancel(context.Background()) + m.cancel = cancel + + m.wg.Add(2) + go func() { + defer m.wg.Done() + if err := m.httpSrv.Serve(ln); err != nil && err != http.ErrServerClosed { + slog.Error("serve error", "module", "dvr", "error", err) + } + }() + go func() { + defer m.wg.Done() + m.runCleanup(ctx) + }() + + slog.Info("listening", "module", "dvr", "addr", ln.Addr(), + "window", m.cfg.Window, "segment", m.cfg.SegmentDuration) + + return nil +} + +// Hooks returns async hooks for publish start/stop events. +func (m *Module) Hooks() []core.HookRegistration { + return []core.HookRegistration{ + { + Event: core.EventPublish, + Mode: core.HookAsync, + Priority: 60, + Handler: m.onPublish, + }, + { + Event: core.EventPublishStop, + Mode: core.HookAsync, + Priority: 60, + Handler: m.onPublishStop, + }, + } +} + +// Close shuts down the DVR server and stops all sessions. +func (m *Module) Close() error { + if m.cancel != nil { + m.cancel() + } + if m.httpSrv != nil { + m.httpSrv.Close() + } + + m.mu.Lock() + sessions := make([]*Session, 0, len(m.sessions)) + for _, s := range m.sessions { + sessions = append(sessions, s) + } + m.mu.Unlock() + + for _, s := range sessions { + s.Stop() + } + + m.wg.Wait() + slog.Info("stopped", "module", "dvr") + return nil +} + +func (m *Module) onPublish(ctx *core.EventContext) error { + if !matchPattern(m.cfg.StreamPattern, ctx.StreamKey) { + return nil + } + + stream, ok := m.server.StreamHub().Find(ctx.StreamKey) + if !ok { + return nil + } + + m.mu.Lock() + defer m.mu.Unlock() + + existing := m.sessions[ctx.StreamKey] + if existing != nil && existing.IsLive() { + return nil + } + + var existingIndex *SegmentIndex + startSeq := 0 + if existing != nil { + existingIndex = existing.Index() + if last, ok := existingIndex.Last(); ok { + startSeq = last.SeqNum + 1 + } + } + + session, err := NewSession(ctx.StreamKey, stream, m.cfg, existingIndex, startSeq) + if err != nil { + slog.Error("failed to start dvr session", "module", "dvr", "stream", ctx.StreamKey, "error", err) + return nil + } + + m.sessions[ctx.StreamKey] = session + go session.Run() + slog.Info("started", "module", "dvr", "stream", ctx.StreamKey) + return nil +} + +func (m *Module) onPublishStop(ctx *core.EventContext) error { + m.mu.Lock() + session := m.sessions[ctx.StreamKey] + m.mu.Unlock() + + if session != nil { + session.Stop() + slog.Info("stream stopped, segments retained", "module", "dvr", "stream", ctx.StreamKey) + } + return nil +} + +// SessionStatus returns the status of all DVR sessions (implements DVRStatusProvider for the API module). +func (m *Module) SessionStatus() any { + m.mu.Lock() + result := make([]DVRSessionStatus, 0, len(m.sessions)) + for key, session := range m.sessions { + segs := session.Index().Segments() + var totalDur float64 + for _, seg := range segs { + totalDur += seg.Duration + } + result = append(result, DVRSessionStatus{ + StreamKey: key, + Live: session.IsLive(), + Segments: len(segs), + Duration: totalDur, + }) + } + m.mu.Unlock() + return result +} + +// DVRSessionStatus represents a single DVR session's status. +type DVRSessionStatus struct { + StreamKey string `json:"stream_key"` + Live bool `json:"live"` + Segments int `json:"segments"` + Duration float64 `json:"duration_sec"` +} diff --git a/module/dvr/playlist.go b/module/dvr/playlist.go new file mode 100644 index 0000000..0b12398 --- /dev/null +++ b/module/dvr/playlist.go @@ -0,0 +1,45 @@ +package dvr + +import ( + "fmt" + "math" + "strings" +) + +// GeneratePlaylist builds an HLS m3u8 playlist from the segment index. +func GeneratePlaylist(index *SegmentIndex, streamKey string, live bool) string { + segments := index.Segments() + if len(segments) == 0 { + return "" + } + + // Extract the key portion (after app/) for relative segment URLs + key := streamKey + if i := strings.IndexByte(streamKey, '/'); i >= 0 { + key = streamKey[i+1:] + } + + maxDur := index.MaxDuration() + targetDur := int(math.Ceil(maxDur)) + if targetDur < 1 { + targetDur = 1 + } + + var b strings.Builder + b.WriteString("#EXTM3U\n") + b.WriteString("#EXT-X-VERSION:6\n") + b.WriteString(fmt.Sprintf("#EXT-X-TARGETDURATION:%d\n", targetDur)) + b.WriteString(fmt.Sprintf("#EXT-X-MEDIA-SEQUENCE:%d\n", segments[0].SeqNum)) + + for _, seg := range segments { + b.WriteString(fmt.Sprintf("#EXT-X-PROGRAM-DATE-TIME:%s\n", seg.StartTime.UTC().Format("2006-01-02T15:04:05.000Z"))) + b.WriteString(fmt.Sprintf("#EXTINF:%.3f,\n", seg.Duration)) + b.WriteString(fmt.Sprintf("%s/%s\n", key, seg.Filename)) + } + + if !live { + b.WriteString("#EXT-X-ENDLIST\n") + } + + return b.String() +} diff --git a/module/dvr/segment_index.go b/module/dvr/segment_index.go new file mode 100644 index 0000000..b1da933 --- /dev/null +++ b/module/dvr/segment_index.go @@ -0,0 +1,124 @@ +package dvr + +import ( + "os" + "sort" + "sync" + "time" +) + +// Segment represents a single TS segment on disk. +type Segment struct { + SeqNum int + StartTime time.Time + StartDTS int64 + Duration float64 + Filename string + Size int64 + DiskPath string +} + +// SegmentIndex is a thread-safe ordered list of segments. +type SegmentIndex struct { + mu sync.RWMutex + segments []Segment +} + +// NewSegmentIndex creates a new segment index. +func NewSegmentIndex() *SegmentIndex { + return &SegmentIndex{} +} + +// Add appends a segment to the index. +func (idx *SegmentIndex) Add(seg Segment) { + idx.mu.Lock() + idx.segments = append(idx.segments, seg) + idx.mu.Unlock() +} + +// Segments returns a copy of all segments. +func (idx *SegmentIndex) Segments() []Segment { + idx.mu.RLock() + out := make([]Segment, len(idx.segments)) + copy(out, idx.segments) + idx.mu.RUnlock() + return out +} + +// SegmentBySeqNum finds a segment by sequence number using binary search. +func (idx *SegmentIndex) SegmentBySeqNum(seqNum int) (Segment, bool) { + idx.mu.RLock() + defer idx.mu.RUnlock() + + i := sort.Search(len(idx.segments), func(i int) bool { + return idx.segments[i].SeqNum >= seqNum + }) + if i < len(idx.segments) && idx.segments[i].SeqNum == seqNum { + return idx.segments[i], true + } + return Segment{}, false +} + +// CleanBefore removes segments with StartTime before cutoff and deletes their disk files. +func (idx *SegmentIndex) CleanBefore(cutoff time.Time) []Segment { + idx.mu.Lock() + defer idx.mu.Unlock() + + splitAt := 0 + for splitAt < len(idx.segments) && idx.segments[splitAt].StartTime.Before(cutoff) { + splitAt++ + } + if splitAt == 0 { + return nil + } + + removed := make([]Segment, splitAt) + copy(removed, idx.segments[:splitAt]) + idx.segments = idx.segments[splitAt:] + + for _, seg := range removed { + os.Remove(seg.DiskPath) + } + + return removed +} + +// Len returns the number of segments. +func (idx *SegmentIndex) Len() int { + idx.mu.RLock() + defer idx.mu.RUnlock() + return len(idx.segments) +} + +// First returns the first (oldest) segment. +func (idx *SegmentIndex) First() (Segment, bool) { + idx.mu.RLock() + defer idx.mu.RUnlock() + if len(idx.segments) == 0 { + return Segment{}, false + } + return idx.segments[0], true +} + +// Last returns the last (newest) segment. +func (idx *SegmentIndex) Last() (Segment, bool) { + idx.mu.RLock() + defer idx.mu.RUnlock() + if len(idx.segments) == 0 { + return Segment{}, false + } + return idx.segments[len(idx.segments)-1], true +} + +// MaxDuration returns the maximum segment duration. +func (idx *SegmentIndex) MaxDuration() float64 { + idx.mu.RLock() + defer idx.mu.RUnlock() + var max float64 + for _, seg := range idx.segments { + if seg.Duration > max { + max = seg.Duration + } + } + return max +} diff --git a/module/dvr/session.go b/module/dvr/session.go new file mode 100644 index 0000000..c81cf97 --- /dev/null +++ b/module/dvr/session.go @@ -0,0 +1,217 @@ +package dvr + +import ( + "fmt" + "log/slog" + "os" + "path/filepath" + "strings" + "sync/atomic" + "time" + + "github.com/im-pingo/liveforge/config" + "github.com/im-pingo/liveforge/core" + "github.com/im-pingo/liveforge/pkg/avframe" + "github.com/im-pingo/liveforge/pkg/muxer/ts" + "github.com/im-pingo/liveforge/pkg/util" +) + +// Session manages DVR segment writing for a single stream. +type Session struct { + streamKey string + stream *core.Stream + cfg config.DVRConfig + index *SegmentIndex + segDir string + + muxer *ts.Muxer + segFile *os.File + segStartDTS int64 + lastDTS int64 + segSeqNum int + wallStart time.Time + + videoSeq []byte + audioSeq []byte + videoCodec avframe.CodecType + audioCodec avframe.CodecType + + reader *util.RingReader[*avframe.AVFrame] + done chan struct{} + stopped atomic.Bool +} + +// NewSession creates a DVR session for the given stream. +func NewSession(streamKey string, stream *core.Stream, cfg config.DVRConfig, existingIndex *SegmentIndex, startSeq int) (*Session, error) { + segDir := resolvePath(cfg.Path, streamKey) + if err := os.MkdirAll(segDir, 0755); err != nil { + return nil, fmt.Errorf("dvr: create dir %s: %w", segDir, err) + } + + idx := existingIndex + if idx == nil { + idx = NewSegmentIndex() + } + + return &Session{ + streamKey: streamKey, + stream: stream, + cfg: cfg, + index: idx, + segDir: segDir, + segStartDTS: -1, + lastDTS: -1, + segSeqNum: startSeq, + reader: stream.RingBuffer().NewReader(), + done: make(chan struct{}), + }, nil +} + +// Run starts the DVR segment writing loop. Blocks until Stop or stream closes. +func (s *Session) Run() { + defer s.closeSegment() + + if vsh := s.stream.VideoSeqHeader(); vsh != nil { + s.videoSeq = append([]byte(nil), vsh.Payload...) + s.videoCodec = vsh.Codec + } + if ash := s.stream.AudioSeqHeader(); ash != nil { + s.audioSeq = append([]byte(nil), ash.Payload...) + s.audioCodec = ash.Codec + } + + for { + frame, ok := s.reader.TryRead() + if ok { + s.processFrame(frame) + continue + } + + if s.stream.RingBuffer().IsClosed() { + return + } + + select { + case <-s.done: + return + case <-s.stream.RingBuffer().Signal(): + } + } +} + +func (s *Session) processFrame(frame *avframe.AVFrame) { + if frame.FrameType == avframe.FrameTypeSequenceHeader { + if frame.MediaType.IsVideo() { + s.videoSeq = append([]byte(nil), frame.Payload...) + s.videoCodec = frame.Codec + } else if frame.MediaType.IsAudio() { + s.audioSeq = append([]byte(nil), frame.Payload...) + s.audioCodec = frame.Codec + } + return + } + + if s.muxer == nil { + s.muxer = ts.NewMuxer(s.videoCodec, s.audioCodec, s.videoSeq, s.audioSeq) + } + + if s.segFile == nil || (frame.MediaType.IsVideo() && frame.FrameType.IsKeyframe() && s.shouldSplit()) { + s.closeSegment() + s.openSegment() + } + + data := s.muxer.WriteFrame(frame) + if data != nil && s.segFile != nil { + s.segFile.Write(data) + } + + if s.segStartDTS < 0 { + s.segStartDTS = frame.DTS + } + s.lastDTS = frame.DTS +} + +func (s *Session) shouldSplit() bool { + if s.segStartDTS < 0 { + return false + } + dur := s.cfg.SegmentDuration + if dur <= 0 { + dur = 6 * time.Second + } + elapsed := time.Duration(s.lastDTS-s.segStartDTS) * time.Millisecond + return elapsed >= dur +} + +func (s *Session) openSegment() { + filename := fmt.Sprintf("seg_%06d.ts", s.segSeqNum) + path := filepath.Join(s.segDir, filename) + + f, err := os.Create(path) + if err != nil { + slog.Error("dvr: create segment", "stream", s.streamKey, "error", err) + return + } + + s.segFile = f + s.wallStart = time.Now() + s.segStartDTS = -1 +} + +func (s *Session) closeSegment() { + if s.segFile == nil { + return + } + + s.segFile.Close() + + dur := float64(0) + if s.segStartDTS >= 0 && s.lastDTS > s.segStartDTS { + dur = float64(s.lastDTS-s.segStartDTS) / 1000.0 + } + + info, _ := os.Stat(s.segFile.Name()) + size := int64(0) + if info != nil { + size = info.Size() + } + + filename := filepath.Base(s.segFile.Name()) + s.index.Add(Segment{ + SeqNum: s.segSeqNum, + StartTime: s.wallStart, + StartDTS: s.segStartDTS, + Duration: dur, + Filename: filename, + Size: size, + DiskPath: s.segFile.Name(), + }) + + s.segSeqNum++ + s.segFile = nil +} + +// Stop signals the session to exit its write loop. +func (s *Session) Stop() { + select { + case <-s.done: + default: + close(s.done) + } + s.stopped.Store(true) +} + +// IsLive returns true if the stream is still publishing. +func (s *Session) IsLive() bool { + return !s.stopped.Load() +} + +// Index returns the session's segment index. +func (s *Session) Index() *SegmentIndex { + return s.index +} + +func resolvePath(pathTemplate, streamKey string) string { + p := strings.ReplaceAll(pathTemplate, "{stream_key}", streamKey) + return filepath.Clean(p) +} diff --git a/module/metrics/metrics_test.go b/module/metrics/metrics_test.go index 3658a8b..b86101e 100644 --- a/module/metrics/metrics_test.go +++ b/module/metrics/metrics_test.go @@ -1,6 +1,7 @@ package metrics import ( + "fmt" "io" "net/http" "strings" @@ -190,3 +191,206 @@ func TestMetricsDefaultPath(t *testing.T) { t.Fatalf("expected 200, got %d", resp.StatusCode) } } + +// TestPerStreamGaugesCreated verifies that per-stream gauges (bitrate, fps, uptime, +// gop_cache_frames) are emitted for each active stream. +func TestPerStreamGaugesCreated(t *testing.T) { + cfg := testConfig() + s := core.NewServer(cfg) + m := NewModule() + s.RegisterModule(m) + + if err := s.Init(); err != nil { + t.Fatalf("init failed: %v", err) + } + defer s.Shutdown() + + stream, err := s.StreamHub().GetOrCreate("live/gauges") + if err != nil { + t.Fatalf("create stream failed: %v", err) + } + + pub := &stubPublisher{ + id: "pub-gauges", + mediaInfo: &avframe.MediaInfo{VideoCodec: avframe.CodecH264, AudioCodec: avframe.CodecAAC}, + } + if err := stream.SetPublisher(pub); err != nil { + t.Fatalf("set publisher failed: %v", err) + } + + // Write a keyframe to populate GOP cache and stats + stream.WriteFrame(&avframe.AVFrame{ + MediaType: avframe.MediaTypeVideo, + FrameType: avframe.FrameTypeKeyframe, + Codec: avframe.CodecH264, + DTS: 0, PTS: 0, + Payload: make([]byte, 2000), + }) + + time.Sleep(10 * time.Millisecond) + + addr := m.Addr() + resp, err := http.Get("http://" + addr.String() + "/metrics") + if err != nil { + t.Fatalf("GET /metrics failed: %v", err) + } + defer resp.Body.Close() + + body, _ := io.ReadAll(resp.Body) + content := string(body) + + // Per-stream gauges that should be present + expectedMetrics := []string{ + `liveforge_stream_bitrate_kbps{stream_key="live/gauges"}`, + `liveforge_stream_fps{stream_key="live/gauges"}`, + `liveforge_stream_uptime_seconds{stream_key="live/gauges"}`, + `liveforge_stream_gop_cache_frames{stream_key="live/gauges"}`, + } + for _, metric := range expectedMetrics { + if !strings.Contains(content, metric) { + t.Errorf("missing per-stream gauge: %s", metric) + } + } +} + +// TestCollectMultipleStreams verifies that metrics are emitted for multiple +// concurrently active streams. +func TestCollectMultipleStreams(t *testing.T) { + cfg := testConfig() + s := core.NewServer(cfg) + m := NewModule() + s.RegisterModule(m) + + if err := s.Init(); err != nil { + t.Fatalf("init failed: %v", err) + } + defer s.Shutdown() + + streamKeys := []string{"live/stream_a", "live/stream_b", "live/stream_c"} + for i, key := range streamKeys { + stream, err := s.StreamHub().GetOrCreate(key) + if err != nil { + t.Fatalf("create stream %s failed: %v", key, err) + } + pub := &stubPublisher{ + id: fmt.Sprintf("pub-%d", i), + mediaInfo: &avframe.MediaInfo{VideoCodec: avframe.CodecH264, AudioCodec: avframe.CodecAAC}, + } + if err := stream.SetPublisher(pub); err != nil { + t.Fatalf("set publisher for %s failed: %v", key, err) + } + stream.WriteFrame(&avframe.AVFrame{ + MediaType: avframe.MediaTypeVideo, + FrameType: avframe.FrameTypeKeyframe, + Codec: avframe.CodecH264, + DTS: int64(i * 1000), PTS: int64(i * 1000), + Payload: make([]byte, 500*(i+1)), + }) + } + + time.Sleep(10 * time.Millisecond) + + addr := m.Addr() + resp, err := http.Get("http://" + addr.String() + "/metrics") + if err != nil { + t.Fatalf("GET /metrics failed: %v", err) + } + defer resp.Body.Close() + + body, _ := io.ReadAll(resp.Body) + content := string(body) + + // Verify active stream count + if !strings.Contains(content, "liveforge_server_streams_active 3") { + t.Error("expected liveforge_server_streams_active to be 3") + } + + // Verify per-stream metrics for each stream + for _, key := range streamKeys { + label := fmt.Sprintf(`stream_key="%s"`, key) + if !strings.Contains(content, label) { + t.Errorf("missing metrics for stream %s", key) + } + } +} + +// TestCollectCleanupRemovedStreams verifies that after a stream is removed from the +// hub, its metrics are no longer emitted. +func TestCollectCleanupRemovedStreams(t *testing.T) { + cfg := testConfig() + s := core.NewServer(cfg) + m := NewModule() + s.RegisterModule(m) + + if err := s.Init(); err != nil { + t.Fatalf("init failed: %v", err) + } + defer s.Shutdown() + + // Create two streams + for _, key := range []string{"live/keep", "live/remove"} { + stream, err := s.StreamHub().GetOrCreate(key) + if err != nil { + t.Fatalf("create stream %s failed: %v", key, err) + } + pub := &stubPublisher{ + id: "pub-" + key, + mediaInfo: &avframe.MediaInfo{VideoCodec: avframe.CodecH264}, + } + if err := stream.SetPublisher(pub); err != nil { + t.Fatalf("set publisher for %s failed: %v", key, err) + } + stream.WriteFrame(&avframe.AVFrame{ + MediaType: avframe.MediaTypeVideo, + FrameType: avframe.FrameTypeKeyframe, + Codec: avframe.CodecH264, + DTS: 0, PTS: 0, + Payload: make([]byte, 500), + }) + } + + time.Sleep(10 * time.Millisecond) + + // Verify both streams appear before removal + addr := m.Addr() + resp1, err := http.Get("http://" + addr.String() + "/metrics") + if err != nil { + t.Fatalf("GET /metrics failed: %v", err) + } + body1, _ := io.ReadAll(resp1.Body) + resp1.Body.Close() + content1 := string(body1) + + if !strings.Contains(content1, `stream_key="live/keep"`) { + t.Error("live/keep should be present before removal") + } + if !strings.Contains(content1, `stream_key="live/remove"`) { + t.Error("live/remove should be present before removal") + } + + // Remove the stream from the hub + s.StreamHub().Remove("live/remove") + + time.Sleep(10 * time.Millisecond) + + // Verify only the kept stream remains + resp2, err := http.Get("http://" + addr.String() + "/metrics") + if err != nil { + t.Fatalf("GET /metrics failed after removal: %v", err) + } + body2, _ := io.ReadAll(resp2.Body) + resp2.Body.Close() + content2 := string(body2) + + if !strings.Contains(content2, `stream_key="live/keep"`) { + t.Error("live/keep should still be present after removing live/remove") + } + if strings.Contains(content2, `stream_key="live/remove"`) { + t.Error("live/remove should not appear after being removed from the hub") + } + + // Stream count should now be 1 + if !strings.Contains(content2, "liveforge_server_streams_active 1") { + t.Error("expected liveforge_server_streams_active to be 1 after removal") + } +} diff --git a/module/notify/notify_test.go b/module/notify/notify_test.go index 5946dec..e649ad9 100644 --- a/module/notify/notify_test.go +++ b/module/notify/notify_test.go @@ -219,3 +219,329 @@ func TestModuleHooks(t *testing.T) { } } } + +// TestHTTPSenderHMACSignatureVerification verifies that the X-Signature header +// contains a valid HMAC-SHA256 of the JSON body using the configured secret. +// This is a more thorough test than TestHTTPSenderHMACSignature, verifying the +// raw signature computation independently. +func TestHTTPSenderHMACSignatureVerification(t *testing.T) { + secret := "webhook-secret-key-123" + var mu sync.Mutex + var receivedSig string + var receivedBody []byte + + ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + body, _ := io.ReadAll(r.Body) + mu.Lock() + receivedSig = r.Header.Get("X-Signature") + receivedBody = body + mu.Unlock() + w.WriteHeader(200) + })) + defer ts.Close() + + endpoints := []config.NotifyEndpointConfig{ + {URL: ts.URL, Secret: secret, Retry: 1, Timeout: 2 * time.Second}, + } + sender := NewHTTPSender(endpoints) + sender.Start() + + payload := &NotifyPayload{ + Event: "on_subscribe", + StreamKey: "live/hmac-verify", + Protocol: "http-flv", + Timestamp: 1700000000, + Extra: map[string]any{"client_id": "abc123"}, + } + sender.Send(payload) + + time.Sleep(200 * time.Millisecond) + sender.Stop() + + mu.Lock() + sig := receivedSig + body := receivedBody + mu.Unlock() + + if sig == "" { + t.Fatal("expected X-Signature header to be set") + } + if len(body) == 0 { + t.Fatal("expected non-empty request body") + } + + // Independently compute HMAC-SHA256 + mac := hmac.New(sha256.New, []byte(secret)) + mac.Write(body) + expectedSig := hex.EncodeToString(mac.Sum(nil)) + + if sig != expectedSig { + t.Errorf("HMAC signature mismatch:\n got: %s\n want: %s", sig, expectedSig) + } + + // Verify the body can be decoded back to the payload + var decoded NotifyPayload + if err := json.Unmarshal(body, &decoded); err != nil { + t.Fatalf("failed to decode body: %v", err) + } + if decoded.Event != "on_subscribe" { + t.Errorf("expected event on_subscribe, got %s", decoded.Event) + } + if decoded.StreamKey != "live/hmac-verify" { + t.Errorf("expected stream_key live/hmac-verify, got %s", decoded.StreamKey) + } +} + +// TestHTTPSenderNoSignatureWithoutSecret verifies that when no secret is +// configured, the X-Signature header is absent. +func TestHTTPSenderNoSignatureWithoutSecret(t *testing.T) { + var mu sync.Mutex + var receivedSig string + + ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + mu.Lock() + receivedSig = r.Header.Get("X-Signature") + mu.Unlock() + w.WriteHeader(200) + })) + defer ts.Close() + + endpoints := []config.NotifyEndpointConfig{ + {URL: ts.URL, Secret: "", Retry: 1, Timeout: 2 * time.Second}, + } + sender := NewHTTPSender(endpoints) + sender.Start() + + sender.Send(&NotifyPayload{Event: "on_publish", StreamKey: "live/nosig", Timestamp: 1}) + + time.Sleep(200 * time.Millisecond) + sender.Stop() + + mu.Lock() + sig := receivedSig + mu.Unlock() + + if sig != "" { + t.Errorf("expected no X-Signature header when secret is empty, got %s", sig) + } +} + +// TestHTTPSenderRetryBehavior verifies that the sender retries delivery when +// the endpoint returns a server error, and eventually succeeds. +func TestHTTPSenderRetryBehavior(t *testing.T) { + var mu sync.Mutex + var attempts int + + ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + mu.Lock() + attempts++ + current := attempts + mu.Unlock() + + if current < 3 { + // Fail the first two attempts + w.WriteHeader(http.StatusInternalServerError) + return + } + // Succeed on the third attempt + w.WriteHeader(http.StatusOK) + })) + defer ts.Close() + + endpoints := []config.NotifyEndpointConfig{ + {URL: ts.URL, Retry: 3, Timeout: 2 * time.Second}, + } + sender := NewHTTPSender(endpoints) + sender.Start() + + sender.Send(&NotifyPayload{Event: "on_publish", StreamKey: "live/retry", Timestamp: 1}) + + // Wait long enough for retries (backoff: 1s, 2s) + time.Sleep(5 * time.Second) + sender.Stop() + + mu.Lock() + finalAttempts := attempts + mu.Unlock() + + if finalAttempts != 3 { + t.Errorf("expected 3 attempts (2 failures + 1 success), got %d", finalAttempts) + } +} + +// TestHTTPSenderRetryExhausted verifies that after all retries are exhausted, +// no further attempts are made. +func TestHTTPSenderRetryExhausted(t *testing.T) { + var mu sync.Mutex + var attempts int + + ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + mu.Lock() + attempts++ + mu.Unlock() + w.WriteHeader(http.StatusBadGateway) + })) + defer ts.Close() + + endpoints := []config.NotifyEndpointConfig{ + {URL: ts.URL, Retry: 2, Timeout: 2 * time.Second}, + } + sender := NewHTTPSender(endpoints) + sender.Start() + + sender.Send(&NotifyPayload{Event: "on_publish", StreamKey: "live/exhaust", Timestamp: 1}) + + // Wait for retries (backoff: 1s) + time.Sleep(3 * time.Second) + sender.Stop() + + mu.Lock() + finalAttempts := attempts + mu.Unlock() + + if finalAttempts != 2 { + t.Errorf("expected exactly 2 attempts (all exhausted), got %d", finalAttempts) + } +} + +// TestHTTPSenderQueueOverflow verifies that when the queue is full, new events +// are dropped without blocking the caller. +func TestHTTPSenderQueueOverflow(t *testing.T) { + // Create a sender with a tiny queue but do NOT start the worker, + // so the queue cannot be drained. + endpoints := []config.NotifyEndpointConfig{ + {URL: "http://localhost:1/noop", Retry: 1, Timeout: 1 * time.Second}, + } + sender := &HTTPSender{ + endpoints: endpoints, + queue: make(chan *NotifyPayload, 2), // capacity of 2 + done: make(chan struct{}), + } + // Do not call sender.Start() — worker is not running, so queue stays full. + + // Fill the queue + sender.Send(&NotifyPayload{Event: "on_publish", StreamKey: "live/a", Timestamp: 1}) + sender.Send(&NotifyPayload{Event: "on_publish", StreamKey: "live/b", Timestamp: 2}) + + // This should be dropped (queue full) without blocking + done := make(chan struct{}) + go func() { + sender.Send(&NotifyPayload{Event: "on_publish", StreamKey: "live/c", Timestamp: 3}) + close(done) + }() + + select { + case <-done: + // Send returned without blocking — correct behavior + case <-time.After(1 * time.Second): + t.Fatal("Send() blocked when queue was full; expected non-blocking drop") + } + + // Queue should have exactly 2 items (capacity) + if len(sender.queue) != 2 { + t.Errorf("expected queue length 2, got %d", len(sender.queue)) + } +} + +// TestHTTPSenderMultipleEndpoints verifies that a single event is delivered +// to all matching endpoints. +func TestHTTPSenderMultipleEndpoints(t *testing.T) { + var mu sync.Mutex + received := make(map[string]int) + + makeServer := func(name string) *httptest.Server { + return httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + mu.Lock() + received[name]++ + mu.Unlock() + w.WriteHeader(200) + })) + } + + ts1 := makeServer("endpoint1") + defer ts1.Close() + ts2 := makeServer("endpoint2") + defer ts2.Close() + ts3 := makeServer("endpoint3") + defer ts3.Close() + + endpoints := []config.NotifyEndpointConfig{ + {URL: ts1.URL, Retry: 1, Timeout: 2 * time.Second}, + {URL: ts2.URL, Retry: 1, Timeout: 2 * time.Second}, + {URL: ts3.URL, Retry: 1, Timeout: 2 * time.Second}, + } + sender := NewHTTPSender(endpoints) + sender.Start() + + sender.Send(&NotifyPayload{Event: "on_publish", StreamKey: "live/multi", Timestamp: 1}) + + time.Sleep(300 * time.Millisecond) + sender.Stop() + + mu.Lock() + defer mu.Unlock() + + for _, name := range []string{"endpoint1", "endpoint2", "endpoint3"} { + if received[name] != 1 { + t.Errorf("expected %s to receive 1 event, got %d", name, received[name]) + } + } +} + +// TestHTTPSenderMultipleEndpointsWithFilters verifies that events are delivered +// only to endpoints whose event filter matches. +func TestHTTPSenderMultipleEndpointsWithFilters(t *testing.T) { + var mu sync.Mutex + received := make(map[string][]string) + + makeServer := func(name string) *httptest.Server { + return httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + body, _ := io.ReadAll(r.Body) + var p NotifyPayload + json.Unmarshal(body, &p) + mu.Lock() + received[name] = append(received[name], p.Event) + mu.Unlock() + w.WriteHeader(200) + })) + } + + ts1 := makeServer("pub_only") + defer ts1.Close() + ts2 := makeServer("sub_only") + defer ts2.Close() + ts3 := makeServer("all_events") + defer ts3.Close() + + endpoints := []config.NotifyEndpointConfig{ + {URL: ts1.URL, Events: []string{"on_publish"}, Retry: 1, Timeout: 2 * time.Second}, + {URL: ts2.URL, Events: []string{"on_subscribe"}, Retry: 1, Timeout: 2 * time.Second}, + {URL: ts3.URL, Retry: 1, Timeout: 2 * time.Second}, // empty filter = all events + } + sender := NewHTTPSender(endpoints) + sender.Start() + + sender.Send(&NotifyPayload{Event: "on_publish", StreamKey: "live/a", Timestamp: 1}) + sender.Send(&NotifyPayload{Event: "on_subscribe", StreamKey: "live/a", Timestamp: 2}) + + time.Sleep(300 * time.Millisecond) + sender.Stop() + + mu.Lock() + defer mu.Unlock() + + // pub_only should receive only on_publish + if len(received["pub_only"]) != 1 || received["pub_only"][0] != "on_publish" { + t.Errorf("pub_only expected [on_publish], got %v", received["pub_only"]) + } + + // sub_only should receive only on_subscribe + if len(received["sub_only"]) != 1 || received["sub_only"][0] != "on_subscribe" { + t.Errorf("sub_only expected [on_subscribe], got %v", received["sub_only"]) + } + + // all_events should receive both + if len(received["all_events"]) != 2 { + t.Errorf("all_events expected 2 events, got %d: %v", len(received["all_events"]), received["all_events"]) + } +} diff --git a/module/record/file_writer.go b/module/record/file_writer.go index a73cb69..5adf530 100644 --- a/module/record/file_writer.go +++ b/module/record/file_writer.go @@ -15,6 +15,7 @@ import ( "github.com/im-pingo/liveforge/pkg/avframe" "github.com/im-pingo/liveforge/pkg/muxer/flv" "github.com/im-pingo/liveforge/pkg/muxer/fmp4" + "github.com/im-pingo/liveforge/pkg/muxer/mp4" ) // frameWriter is the interface for format-specific frame writing. @@ -117,6 +118,56 @@ func (w *fmp4FrameWriter) flush(f *os.File) error { return nil } +// mp4FrameWriter writes AVFrames as a classic MP4 with moov atom at end. +type mp4FrameWriter struct { + muxer *mp4.Muxer + started bool + prevDTS int64 +} + +func (w *mp4FrameWriter) writeHeader(f *os.File, frame *avframe.AVFrame) error { + return nil // defer until we have codec info +} + +func (w *mp4FrameWriter) writeFrame(f *os.File, frame *avframe.AVFrame) error { + if frame.FrameType == avframe.FrameTypeSequenceHeader { + if w.muxer == nil { + videoCodec := avframe.CodecType(0) + audioCodec := avframe.CodecType(0) + if frame.MediaType.IsVideo() { + videoCodec = frame.Codec + } else { + audioCodec = frame.Codec + } + w.muxer = mp4.NewMuxer(videoCodec, audioCodec) + } + w.muxer.WriteFrame(f, frame, -1) + + if !w.started { + w.muxer.WriteFtyp(f) + w.muxer.WriteMdatHeader(f) + w.started = true + w.prevDTS = -1 + } + return nil + } + + if !w.started || w.muxer == nil { + return nil + } + + _, err := w.muxer.WriteFrame(f, frame, w.prevDTS) + w.prevDTS = frame.DTS + return err +} + +func (w *mp4FrameWriter) finalize(f *os.File) error { + if w.muxer == nil || !w.started { + return nil + } + return w.muxer.Finalize(f) +} + // FileWriter manages writing AVFrames to files with optional segmentation. type FileWriter struct { cfg config.RecordConfig @@ -135,7 +186,7 @@ func NewFileWriter(streamKey string, cfg config.RecordConfig) (*FileWriter, erro w := &FileWriter{ cfg: cfg, streamKey: streamKey, - format: newFrameWriter(cfg.Format), + format: newFrameWriterWithContext(cfg.Format, streamKey, cfg), startTime: time.Now(), } @@ -147,19 +198,31 @@ func NewFileWriter(streamKey string, cfg config.RecordConfig) (*FileWriter, erro // newFrameWriter creates a format-specific writer based on the config format string. func newFrameWriter(format string) frameWriter { + return newFrameWriterWithContext(format, "", config.RecordConfig{}) +} + +func newFrameWriterWithContext(format, streamKey string, cfg config.RecordConfig) frameWriter { switch strings.ToLower(format) { - case "fmp4", "mp4": + case "mp4": + return &mp4FrameWriter{} + case "fmp4": return &fmp4FrameWriter{} + case "ts", "hls": + return newTSFrameWriter(streamKey, cfg) default: return &flvFrameWriter{muxer: flv.NewMuxer()} } } -// Format returns the recording format string ("flv" or "fmp4"). +// Format returns the recording format string ("flv", "fmp4", "mp4", or "ts"). func (w *FileWriter) Format() string { switch w.format.(type) { + case *mp4FrameWriter: + return "mp4" case *fmp4FrameWriter: return "fmp4" + case *tsFrameWriter: + return "ts" default: return "flv" } @@ -200,9 +263,13 @@ func (w *FileWriter) WriteFrame(frame *avframe.AVFrame) error { // Close flushes and closes the current file. func (w *FileWriter) Close() { if w.file != nil { - // Flush any buffered data for fMP4 - if fmp4w, ok := w.format.(*fmp4FrameWriter); ok { - fmp4w.flush(w.file) //nolint:errcheck + switch fw := w.format.(type) { + case *fmp4FrameWriter: + fw.flush(w.file) //nolint:errcheck + case *mp4FrameWriter: + fw.finalize(w.file) //nolint:errcheck + case *tsFrameWriter: + fw.flush(w.file) //nolint:errcheck } w.file.Close() w.notifyFileComplete() @@ -238,9 +305,13 @@ func (w *FileWriter) openFile() error { } func (w *FileWriter) rotate() error { - // Flush any buffered data for fMP4 - if fmp4w, ok := w.format.(*fmp4FrameWriter); ok { - fmp4w.flush(w.file) //nolint:errcheck + switch fw := w.format.(type) { + case *fmp4FrameWriter: + fw.flush(w.file) //nolint:errcheck + case *mp4FrameWriter: + fw.finalize(w.file) //nolint:errcheck + case *tsFrameWriter: + fw.flush(w.file) //nolint:errcheck } // Close current file diff --git a/module/record/record_test.go b/module/record/record_test.go index a2f9f81..f293a8a 100644 --- a/module/record/record_test.go +++ b/module/record/record_test.go @@ -3,6 +3,7 @@ package record import ( "os" "path/filepath" + "strings" "testing" "time" @@ -246,8 +247,8 @@ func TestNewFrameWriterFMP4(t *testing.T) { } w = newFrameWriter("mp4") - if _, ok := w.(*fmp4FrameWriter); !ok { - t.Errorf("mp4 should map to fmp4FrameWriter, got %T", w) + if _, ok := w.(*mp4FrameWriter); !ok { + t.Errorf("mp4 should map to mp4FrameWriter, got %T", w) } } @@ -593,6 +594,141 @@ func TestNotifyFileComplete(t *testing.T) { w.notifyFileComplete() } +func TestNewFrameWriterTS(t *testing.T) { + w := newFrameWriter("ts") + if _, ok := w.(*tsFrameWriter); !ok { + t.Errorf("expected tsFrameWriter, got %T", w) + } + w = newFrameWriter("hls") + if _, ok := w.(*tsFrameWriter); !ok { + t.Errorf("hls should map to tsFrameWriter, got %T", w) + } +} + +func TestTSRecordingCreatesSegmentsAndPlaylist(t *testing.T) { + dir := t.TempDir() + cfg := config.RecordConfig{ + Format: "ts", + Path: filepath.Join(dir, "{stream_key}", "{date}_{time}.ts"), + Segment: config.SegmentConfig{ + Duration: 100 * time.Millisecond, + }, + } + + w, err := NewFileWriter("live/tstest", cfg) + if err != nil { + t.Fatal(err) + } + + // Write H.264 sequence header + seqHeader := avframe.NewAVFrame( + avframe.MediaTypeVideo, avframe.CodecH264, avframe.FrameTypeSequenceHeader, + 0, 0, []byte{0x01, 0x64, 0x00, 0x28, 0xFF, 0xE1, 0x00, 0x04, 0x67, 0x64, 0x00, 0x28, 0x01, 0x00, 0x04, 0x68, 0xEE, 0x3C, 0x80}, + ) + if err := w.WriteFrame(seqHeader); err != nil { + t.Fatalf("write seq header: %v", err) + } + + // Write first GOP: keyframe + interframes + keyframe1 := avframe.NewAVFrame( + avframe.MediaTypeVideo, avframe.CodecH264, avframe.FrameTypeKeyframe, + 0, 0, []byte{0x00, 0x00, 0x00, 0x02, 0x65, 0x88}, + ) + if err := w.WriteFrame(keyframe1); err != nil { + t.Fatalf("write keyframe1: %v", err) + } + + for i := 1; i <= 5; i++ { + inter := avframe.NewAVFrame( + avframe.MediaTypeVideo, avframe.CodecH264, avframe.FrameTypeInterframe, + int64(i*33), int64(i*33), []byte{0x00, 0x00, 0x00, 0x02, 0x41, 0x9A}, + ) + if err := w.WriteFrame(inter); err != nil { + t.Fatalf("write interframe %d: %v", i, err) + } + } + + // Write second GOP to trigger segment split + keyframe2 := avframe.NewAVFrame( + avframe.MediaTypeVideo, avframe.CodecH264, avframe.FrameTypeKeyframe, + 200, 200, []byte{0x00, 0x00, 0x00, 0x02, 0x65, 0x89}, + ) + if err := w.WriteFrame(keyframe2); err != nil { + t.Fatalf("write keyframe2: %v", err) + } + + for i := 1; i <= 3; i++ { + inter := avframe.NewAVFrame( + avframe.MediaTypeVideo, avframe.CodecH264, avframe.FrameTypeInterframe, + int64(200+i*33), int64(200+i*33), []byte{0x00, 0x00, 0x00, 0x02, 0x41, 0x9B}, + ) + if err := w.WriteFrame(inter); err != nil { + t.Fatalf("write interframe2 %d: %v", i, err) + } + } + + w.Close() + + // Verify: the TS writer should have created segment files and a playlist + parentDir := filepath.Dir(w.FilePath()) + entries, err := os.ReadDir(parentDir) + if err != nil { + t.Fatalf("read dir: %v", err) + } + + var tsFiles, m3u8Files []string + for _, e := range entries { + if filepath.Ext(e.Name()) == ".ts" { + tsFiles = append(tsFiles, e.Name()) + } + if filepath.Ext(e.Name()) == ".m3u8" { + m3u8Files = append(m3u8Files, e.Name()) + } + } + + if len(tsFiles) == 0 { + t.Error("expected at least one .ts segment file") + } + if len(m3u8Files) == 0 { + t.Error("expected index.m3u8 playlist file") + } + + // Read playlist and verify VOD markers + if len(m3u8Files) > 0 { + playlist, err := os.ReadFile(filepath.Join(parentDir, m3u8Files[0])) + if err != nil { + t.Fatalf("read playlist: %v", err) + } + content := string(playlist) + if !strings.Contains(content, "#EXTM3U") { + t.Error("playlist missing #EXTM3U header") + } + if !strings.Contains(content, "#EXT-X-PLAYLIST-TYPE:VOD") { + t.Error("playlist missing VOD type") + } + if !strings.Contains(content, "#EXT-X-ENDLIST") { + t.Error("playlist missing #EXT-X-ENDLIST") + } + t.Logf("TS segments: %v, playlist:\n%s", tsFiles, content) + } +} + +func TestFileWriterFormatTS(t *testing.T) { + dir := t.TempDir() + cfg := config.RecordConfig{ + Format: "ts", + Path: filepath.Join(dir, "{stream_key}", "{date}_{time}.ts"), + } + w, err := NewFileWriter("live/test", cfg) + if err != nil { + t.Fatal(err) + } + if w.Format() != "ts" { + t.Errorf("Format = %q, want ts", w.Format()) + } + w.Close() +} + func TestModuleHooks(t *testing.T) { dir := t.TempDir() cfg := newTestConfig(dir) diff --git a/module/record/ts_writer.go b/module/record/ts_writer.go new file mode 100644 index 0000000..0cafdaa --- /dev/null +++ b/module/record/ts_writer.go @@ -0,0 +1,184 @@ +package record + +import ( + "fmt" + "os" + "path/filepath" + "strings" + "time" + + "github.com/im-pingo/liveforge/config" + "github.com/im-pingo/liveforge/pkg/avframe" + "github.com/im-pingo/liveforge/pkg/muxer/ts" +) + +// tsFrameWriter writes AVFrames as MPEG-TS segments with an HLS playlist. +type tsFrameWriter struct { + cfg config.RecordConfig + streamKey string + + muxer *ts.Muxer + videoSeq []byte + audioSeq []byte + videoCodec avframe.CodecType + audioCodec avframe.CodecType + + segmentDir string + segmentFile *os.File + segmentPath string + segmentIdx int + segmentDur time.Duration + segStartTS int64 // DTS of first frame in segment (ms) + lastDTS int64 + + segments []segmentInfo +} + +type segmentInfo struct { + filename string + duration float64 // seconds +} + +func newTSFrameWriter(streamKey string, cfg config.RecordConfig) *tsFrameWriter { + return &tsFrameWriter{ + cfg: cfg, + streamKey: streamKey, + lastDTS: -1, + segStartTS: -1, + } +} + +func (w *tsFrameWriter) writeHeader(f *os.File, frame *avframe.AVFrame) error { + return nil +} + +func (w *tsFrameWriter) writeFrame(f *os.File, frame *avframe.AVFrame) error { + if frame.FrameType == avframe.FrameTypeSequenceHeader { + if frame.MediaType.IsVideo() { + w.videoSeq = append([]byte(nil), frame.Payload...) + w.videoCodec = frame.Codec + } else if frame.MediaType.IsAudio() { + w.audioSeq = append([]byte(nil), frame.Payload...) + w.audioCodec = frame.Codec + } + return nil + } + + if w.muxer == nil { + w.muxer = ts.NewMuxer(w.videoCodec, w.audioCodec, w.videoSeq, w.audioSeq) + w.segmentDir = filepath.Dir(f.Name()) + } + + // Start new segment on keyframe or if no segment is open + if w.segmentFile == nil || (frame.MediaType.IsVideo() && frame.FrameType.IsKeyframe() && w.shouldSplit()) { + if err := w.closeSegment(); err != nil { + return err + } + if err := w.openSegment(); err != nil { + return err + } + } + + data := w.muxer.WriteFrame(frame) + if data != nil { + if _, err := w.segmentFile.Write(data); err != nil { + return fmt.Errorf("write TS data: %w", err) + } + } + + if w.segStartTS < 0 { + w.segStartTS = frame.DTS + } + w.lastDTS = frame.DTS + + return nil +} + +func (w *tsFrameWriter) shouldSplit() bool { + if w.segStartTS < 0 { + return false + } + dur := w.cfg.Segment.Duration + if dur <= 0 { + dur = 6 * time.Second + } + elapsed := time.Duration(w.lastDTS-w.segStartTS) * time.Millisecond + return elapsed >= dur +} + +func (w *tsFrameWriter) openSegment() error { + filename := fmt.Sprintf("segment_%05d.ts", w.segmentIdx) + path := filepath.Join(w.segmentDir, filename) + + sf, err := os.Create(path) + if err != nil { + return fmt.Errorf("create segment %s: %w", path, err) + } + + w.segmentFile = sf + w.segmentPath = path + w.segStartTS = -1 + return nil +} + +func (w *tsFrameWriter) closeSegment() error { + if w.segmentFile == nil { + return nil + } + + w.segmentFile.Close() + + dur := float64(0) + if w.segStartTS >= 0 && w.lastDTS > w.segStartTS { + dur = float64(w.lastDTS-w.segStartTS) / 1000.0 + } + + w.segments = append(w.segments, segmentInfo{ + filename: filepath.Base(w.segmentPath), + duration: dur, + }) + w.segmentIdx++ + w.segmentFile = nil + + return nil +} + +// flush writes the final segment and playlist. +func (w *tsFrameWriter) flush(f *os.File) error { + if err := w.closeSegment(); err != nil { + return err + } + return w.writePlaylist() +} + +func (w *tsFrameWriter) writePlaylist() error { + if len(w.segments) == 0 { + return nil + } + + dir := w.segmentDir + playlistPath := filepath.Join(dir, "index.m3u8") + + var maxDur float64 + for _, seg := range w.segments { + if seg.duration > maxDur { + maxDur = seg.duration + } + } + + var b strings.Builder + b.WriteString("#EXTM3U\n") + b.WriteString("#EXT-X-VERSION:3\n") + b.WriteString(fmt.Sprintf("#EXT-X-TARGETDURATION:%d\n", int(maxDur)+1)) + b.WriteString("#EXT-X-MEDIA-SEQUENCE:0\n") + b.WriteString("#EXT-X-PLAYLIST-TYPE:VOD\n") + + for _, seg := range w.segments { + b.WriteString(fmt.Sprintf("#EXTINF:%.3f,\n", seg.duration)) + b.WriteString(seg.filename + "\n") + } + + b.WriteString("#EXT-X-ENDLIST\n") + + return os.WriteFile(playlistPath, []byte(b.String()), 0644) +} diff --git a/module/rtmp/handler.go b/module/rtmp/handler.go index a8fe967..96ccea5 100644 --- a/module/rtmp/handler.go +++ b/module/rtmp/handler.go @@ -11,6 +11,7 @@ import ( "github.com/im-pingo/liveforge/config" "github.com/im-pingo/liveforge/core" "github.com/im-pingo/liveforge/pkg/avframe" + "github.com/im-pingo/liveforge/pkg/codec/aac" ) // isConnClosed checks if an error indicates a closed connection. @@ -345,8 +346,16 @@ func (h *Handler) handleMediaMessage(msg *Message) error { mi := rp.MediaInfo() if frame.MediaType.IsVideo() { mi.VideoCodec = frame.Codec + mi.VideoSequenceHeader = append([]byte(nil), frame.Payload...) } else if frame.MediaType.IsAudio() { mi.AudioCodec = frame.Codec + mi.AudioSequenceHeader = append([]byte(nil), frame.Payload...) + if frame.Codec == avframe.CodecAAC && len(frame.Payload) >= 2 { + if ascInfo, err := aac.ParseAudioSpecificConfig(frame.Payload); err == nil { + mi.SampleRate = ascInfo.SampleRate + mi.Channels = ascInfo.Channels + } + } } } } diff --git a/module/rtsp/handler.go b/module/rtsp/handler.go index 228a0b1..0828cdd 100644 --- a/module/rtsp/handler.go +++ b/module/rtsp/handler.go @@ -7,6 +7,7 @@ import ( "net/http" "strings" + "github.com/im-pingo/liveforge/config" "github.com/im-pingo/liveforge/core" "github.com/im-pingo/liveforge/pkg/avframe" "github.com/im-pingo/liveforge/pkg/portalloc" @@ -15,13 +16,14 @@ import ( // Handler processes RTSP requests. type Handler struct { - server *core.Server - ports *portalloc.PortAllocator + server *core.Server + ports *portalloc.PortAllocator + multicast *config.MulticastConfig // nil if multicast disabled } // NewHandler creates a new RTSP handler. -func NewHandler(server *core.Server, ports *portalloc.PortAllocator) *Handler { - return &Handler{server: server, ports: ports} +func NewHandler(server *core.Server, ports *portalloc.PortAllocator, multicast *config.MulticastConfig) *Handler { + return &Handler{server: server, ports: ports, multicast: multicast} } // newResponse creates a base response with CSeq from request. @@ -82,6 +84,7 @@ func (h *Handler) HandleDescribe(req *Request, session *RTSPSession) *Response { // TransportConfig holds parsed Transport header data. type TransportConfig struct { IsTCP bool + IsMulticast bool Interleaved [2]int // channel pair for TCP ClientPorts [2]int // client ports for UDP ServerPorts [2]int // allocated server ports for UDP @@ -93,7 +96,19 @@ func (h *Handler) HandleSetup(req *Request, session *RTSPSession, remoteAddr str tc := parseTransportHeader(transport) var udpTransport *UDPTransport - if !tc.IsTCP && h.ports != nil { + var mcastTransport *MulticastTransport + + if tc.IsMulticast && h.multicast != nil { + mt, err := NewMulticastTransport(*h.multicast) + if err != nil { + return newResponse(500, "Internal Server Error", req) + } + rtpPort, rtcpPort := mt.ServerPorts() + tc.ServerPorts = [2]int{rtpPort, rtcpPort} + mcastTransport = mt + } else if tc.IsMulticast { + return newResponse(461, "Unsupported Transport", req) + } else if !tc.IsTCP && h.ports != nil { ut, err := NewUDPTransport(h.ports) if err != nil { return newResponse(500, "Internal Server Error", req) @@ -102,7 +117,6 @@ func (h *Handler) HandleSetup(req *Request, session *RTSPSession, remoteAddr str tc.ServerPorts = [2]int{rtpPort, rtcpPort} udpTransport = ut - // Set client address from client_port and remote IP. host, _, _ := net.SplitHostPort(remoteAddr) clientIP := net.ParseIP(host) if clientIP != nil { @@ -110,15 +124,14 @@ func (h *Handler) HandleSetup(req *Request, session *RTSPSession, remoteAddr str } } - // Store transport config per track on the session. if session != nil { trackID, _ := extractTrackID(req.URL) ts := TrackSetup{ TrackID: trackID, Transport: tc, UDP: udpTransport, + Multicast: mcastTransport, } - // Assign codec from MediaInfo based on track order. if session.MediaInfo != nil { idx := len(session.Tracks) if idx == 0 && session.MediaInfo.HasVideo() { @@ -133,6 +146,9 @@ func (h *Handler) HandleSetup(req *Request, session *RTSPSession, remoteAddr str resp := newResponse(200, "OK", req) if tc.IsTCP { resp.Headers.Set("Transport", fmt.Sprintf("RTP/AVP/TCP;unicast;interleaved=%d-%d", tc.Interleaved[0], tc.Interleaved[1])) + } else if tc.IsMulticast && mcastTransport != nil { + resp.Headers.Set("Transport", fmt.Sprintf("RTP/AVP;multicast;destination=%s;port=%d-%d;ttl=%d", + mcastTransport.MulticastAddr(), tc.ServerPorts[0], tc.ServerPorts[1], h.multicast.TTL)) } else { resp.Headers.Set("Transport", fmt.Sprintf("RTP/AVP;unicast;client_port=%d-%d;server_port=%d-%d", tc.ClientPorts[0], tc.ClientPorts[1], tc.ServerPorts[0], tc.ServerPorts[1])) @@ -282,6 +298,9 @@ func parseTransportHeader(transport string) TransportConfig { if part == "RTP/AVP/TCP" { tc.IsTCP = true } + if part == "multicast" { + tc.IsMulticast = true + } if strings.HasPrefix(part, "interleaved=") { fmt.Sscanf(part, "interleaved=%d-%d", &tc.Interleaved[0], &tc.Interleaved[1]) } diff --git a/module/rtsp/handler_test.go b/module/rtsp/handler_test.go index da829a3..9626b69 100644 --- a/module/rtsp/handler_test.go +++ b/module/rtsp/handler_test.go @@ -9,7 +9,7 @@ import ( ) func TestHandleOptions(t *testing.T) { - h := NewHandler(nil, nil) + h := NewHandler(nil, nil, nil) req := &Request{Method: "OPTIONS", URL: "*", Headers: make(http.Header)} req.Headers.Set("CSeq", "1") resp := h.HandleOptions(req) @@ -28,7 +28,7 @@ func TestHandleOptions(t *testing.T) { } func TestHandleGetParameter(t *testing.T) { - h := NewHandler(nil, nil) + h := NewHandler(nil, nil, nil) req := &Request{Method: "GET_PARAMETER", URL: "rtsp://host/live/test", Headers: make(http.Header)} req.Headers.Set("CSeq", "2") resp := h.HandleGetParameter(req) @@ -38,7 +38,7 @@ func TestHandleGetParameter(t *testing.T) { } func TestHandleSetupTCPInterleaved(t *testing.T) { - h := NewHandler(nil, nil) + h := NewHandler(nil, nil, nil) session := NewRTSPSession("test-id", "live/room1") session.Transition(StateDescribed) req := &Request{Method: "SETUP", URL: "rtsp://host/live/test/trackID=0", Headers: make(http.Header)} @@ -59,7 +59,7 @@ func TestHandleSetupTCPInterleaved(t *testing.T) { func TestHandleSetupUDP(t *testing.T) { pa, _ := portalloc.New(10000, 10010) - h := NewHandler(nil, pa) + h := NewHandler(nil, pa, nil) session := NewRTSPSession("test-id", "live/room1") session.Transition(StateDescribed) req := &Request{Method: "SETUP", URL: "rtsp://host/live/test/trackID=0", Headers: make(http.Header)} @@ -76,7 +76,7 @@ func TestHandleSetupUDP(t *testing.T) { } func TestHandleAnnounce(t *testing.T) { - h := NewHandler(nil, nil) + h := NewHandler(nil, nil, nil) session := NewRTSPSession("test-id", "live/room1") sdpBody := "v=0\r\no=- 0 0 IN IP4 0.0.0.0\r\ns=test\r\nt=0 0\r\nm=video 0 RTP/AVP 96\r\na=rtpmap:96 H264/90000\r\n" req := &Request{Method: "ANNOUNCE", URL: "rtsp://host/live/test", Headers: make(http.Header), Body: []byte(sdpBody)} @@ -91,7 +91,7 @@ func TestHandleAnnounce(t *testing.T) { } func TestHandleAnnounceNoBody(t *testing.T) { - h := NewHandler(nil, nil) + h := NewHandler(nil, nil, nil) session := NewRTSPSession("test-id", "live/room1") req := &Request{Method: "ANNOUNCE", URL: "rtsp://host/live/test", Headers: make(http.Header)} req.Headers.Set("CSeq", "1") @@ -102,7 +102,7 @@ func TestHandleAnnounceNoBody(t *testing.T) { } func TestHandleRecord(t *testing.T) { - h := NewHandler(nil, nil) + h := NewHandler(nil, nil, nil) session := NewRTSPSession("test-id", "live/room1") session.Transition(StateAnnounced) session.Transition(StateReady) @@ -118,7 +118,7 @@ func TestHandleRecord(t *testing.T) { } func TestHandleRecordInvalidState(t *testing.T) { - h := NewHandler(nil, nil) + h := NewHandler(nil, nil, nil) session := NewRTSPSession("test-id", "live/room1") // Init state -> cannot RECORD directly req := &Request{Method: "RECORD", URL: "rtsp://host/live/test", Headers: make(http.Header)} @@ -130,7 +130,7 @@ func TestHandleRecordInvalidState(t *testing.T) { } func TestHandlePlay(t *testing.T) { - h := NewHandler(nil, nil) + h := NewHandler(nil, nil, nil) session := NewRTSPSession("test-id", "live/room1") session.Transition(StateDescribed) session.Transition(StateReady) @@ -149,7 +149,7 @@ func TestHandlePlay(t *testing.T) { } func TestHandlePause(t *testing.T) { - h := NewHandler(nil, nil) + h := NewHandler(nil, nil, nil) session := NewRTSPSession("test-id", "live/room1") session.Transition(StateDescribed) session.Transition(StateReady) @@ -166,7 +166,7 @@ func TestHandlePause(t *testing.T) { } func TestHandleTeardown(t *testing.T) { - h := NewHandler(nil, nil) + h := NewHandler(nil, nil, nil) session := NewRTSPSession("test-id", "live/room1") session.Transition(StateDescribed) session.Transition(StateReady) diff --git a/module/rtsp/helpers_test.go b/module/rtsp/helpers_test.go index aa098b0..d9957e8 100644 --- a/module/rtsp/helpers_test.go +++ b/module/rtsp/helpers_test.go @@ -156,7 +156,7 @@ func TestNewRTSPSubscriberVideoOnly(t *testing.T) { } func TestHandleDescribeNoServer(t *testing.T) { - h := NewHandler(nil, nil) + h := NewHandler(nil, nil, nil) session := NewRTSPSession("test-id", "live/test") req := &Request{Method: "DESCRIBE", URL: "rtsp://host/live/test", Headers: make(map[string][]string)} req.Headers.Set("CSeq", "1") @@ -213,7 +213,7 @@ func TestNewResponseNilRequest(t *testing.T) { } func TestHandlePlayInvalidState(t *testing.T) { - h := NewHandler(nil, nil) + h := NewHandler(nil, nil, nil) session := NewRTSPSession("test-id", "live/test") // Init state -> cannot PLAY directly req := &Request{Method: "PLAY", URL: "rtsp://host/live/test", Headers: make(map[string][]string)} @@ -226,7 +226,7 @@ func TestHandlePlayInvalidState(t *testing.T) { func TestHandleSetupNoPortManager(t *testing.T) { // TCP transport should work without port manager - h := NewHandler(nil, nil) + h := NewHandler(nil, nil, nil) session := NewRTSPSession("test-id", "live/room1") session.Transition(StateDescribed) req := &Request{Method: "SETUP", URL: "rtsp://host/live/test/trackID=0", Headers: make(map[string][]string)} @@ -239,7 +239,7 @@ func TestHandleSetupNoPortManager(t *testing.T) { } func TestHandleSetupWithMediaInfo(t *testing.T) { - h := NewHandler(nil, nil) + h := NewHandler(nil, nil, nil) session := NewRTSPSession("test-id", "live/room1") session.Transition(StateDescribed) session.MediaInfo = &avframe.MediaInfo{ diff --git a/module/rtsp/multicast.go b/module/rtsp/multicast.go new file mode 100644 index 0000000..de95c00 --- /dev/null +++ b/module/rtsp/multicast.go @@ -0,0 +1,143 @@ +package rtsp + +import ( + "fmt" + "net" + "sync" + "sync/atomic" + + "github.com/im-pingo/liveforge/config" + "golang.org/x/net/ipv4" +) + +// MulticastTransport sends RTP/RTCP to a multicast group address. +// Multiple subscribers share the same multicast stream; each subscriber +// joins the group on their end and receives via IGMP. +type MulticastTransport struct { + rtpConn net.PacketConn + rtcpConn net.PacketConn + rtpAddr *net.UDPAddr + rtcpAddr *net.UDPAddr + rtpPort int + rtcpPort int + done chan struct{} + closed bool + mu sync.Mutex +} + +// multicastPortCounter provides monotonically increasing port numbers +// for multicast streams, starting from the configured base port. +var multicastPortCounter atomic.Int64 + +// InitMulticastPorts sets the starting port for multicast allocation. +func InitMulticastPorts(basePort int) { + multicastPortCounter.Store(int64(basePort)) +} + +// NewMulticastTransport creates a transport that sends to a multicast group. +func NewMulticastTransport(cfg config.MulticastConfig) (*MulticastTransport, error) { + groupIP := net.ParseIP(cfg.Address) + if groupIP == nil { + return nil, fmt.Errorf("invalid multicast address: %s", cfg.Address) + } + if !groupIP.IsMulticast() { + return nil, fmt.Errorf("not a multicast address: %s", cfg.Address) + } + + rtpPort := int(multicastPortCounter.Add(2) - 2) + rtcpPort := rtpPort + 1 + + ttl := cfg.TTL + if ttl <= 0 { + ttl = 16 + } + + var iface *net.Interface + if cfg.Interface != "" { + var err error + iface, err = net.InterfaceByName(cfg.Interface) + if err != nil { + return nil, fmt.Errorf("interface %s: %w", cfg.Interface, err) + } + } + + rtpAddr := &net.UDPAddr{IP: groupIP, Port: rtpPort} + rtcpAddr := &net.UDPAddr{IP: groupIP, Port: rtcpPort} + + rtpConn, err := setupMulticastSender(iface, ttl) + if err != nil { + return nil, fmt.Errorf("multicast RTP sender: %w", err) + } + + rtcpConn, err := setupMulticastSender(iface, ttl) + if err != nil { + rtpConn.Close() + return nil, fmt.Errorf("multicast RTCP sender: %w", err) + } + + return &MulticastTransport{ + rtpConn: rtpConn, + rtcpConn: rtcpConn, + rtpAddr: rtpAddr, + rtcpAddr: rtcpAddr, + rtpPort: rtpPort, + rtcpPort: rtcpPort, + done: make(chan struct{}), + }, nil +} + +func setupMulticastSender(iface *net.Interface, ttl int) (net.PacketConn, error) { + conn, err := net.ListenPacket("udp4", "0.0.0.0:0") + if err != nil { + return nil, err + } + + p := ipv4.NewPacketConn(conn) + if err := p.SetMulticastTTL(ttl); err != nil { + conn.Close() + return nil, fmt.Errorf("set TTL: %w", err) + } + if iface != nil { + if err := p.SetMulticastInterface(iface); err != nil { + conn.Close() + return nil, fmt.Errorf("set interface: %w", err) + } + } + + return conn, nil +} + +// SendRTP sends an RTP packet to the multicast group. +func (m *MulticastTransport) SendRTP(data []byte) error { + _, err := m.rtpConn.WriteTo(data, m.rtpAddr) + return err +} + +// SendRTCP sends an RTCP packet to the multicast group. +func (m *MulticastTransport) SendRTCP(data []byte) error { + _, err := m.rtcpConn.WriteTo(data, m.rtcpAddr) + return err +} + +// ServerPorts returns the multicast RTP and RTCP ports. +func (m *MulticastTransport) ServerPorts() (int, int) { + return m.rtpPort, m.rtcpPort +} + +// MulticastAddr returns the multicast group IP address. +func (m *MulticastTransport) MulticastAddr() net.IP { + return m.rtpAddr.IP +} + +// Close shuts down the multicast transport. +func (m *MulticastTransport) Close() { + m.mu.Lock() + defer m.mu.Unlock() + if m.closed { + return + } + m.closed = true + close(m.done) + m.rtpConn.Close() + m.rtcpConn.Close() +} diff --git a/module/rtsp/multicast_test.go b/module/rtsp/multicast_test.go new file mode 100644 index 0000000..4390ee6 --- /dev/null +++ b/module/rtsp/multicast_test.go @@ -0,0 +1,214 @@ +package rtsp + +import ( + "net" + "testing" + + "github.com/im-pingo/liveforge/config" +) + +func TestNewMulticastTransport(t *testing.T) { + InitMulticastPorts(40000) + cfg := config.MulticastConfig{ + Enabled: true, + Address: "239.0.0.1", + BasePort: 40000, + TTL: 4, + } + + mt, err := NewMulticastTransport(cfg) + if err != nil { + t.Fatalf("NewMulticastTransport: %v", err) + } + defer mt.Close() + + rtpPort, rtcpPort := mt.ServerPorts() + if rtpPort != 40000 { + t.Errorf("expected RTP port 40000, got %d", rtpPort) + } + if rtcpPort != 40001 { + t.Errorf("expected RTCP port 40001, got %d", rtcpPort) + } + + addr := mt.MulticastAddr() + if !addr.Equal(net.ParseIP("239.0.0.1")) { + t.Errorf("unexpected multicast addr: %v", addr) + } +} + +func TestMulticastTransportPortIncrement(t *testing.T) { + InitMulticastPorts(42000) + cfg := config.MulticastConfig{ + Enabled: true, + Address: "239.0.0.2", + TTL: 4, + } + + mt1, err := NewMulticastTransport(cfg) + if err != nil { + t.Fatalf("first transport: %v", err) + } + defer mt1.Close() + + mt2, err := NewMulticastTransport(cfg) + if err != nil { + t.Fatalf("second transport: %v", err) + } + defer mt2.Close() + + rtp1, _ := mt1.ServerPorts() + rtp2, _ := mt2.ServerPorts() + if rtp2 != rtp1+2 { + t.Errorf("expected port increment by 2: first=%d second=%d", rtp1, rtp2) + } +} + +func TestMulticastTransportInvalidAddress(t *testing.T) { + cfg := config.MulticastConfig{ + Enabled: true, + Address: "not-an-ip", + TTL: 4, + } + _, err := NewMulticastTransport(cfg) + if err == nil { + t.Error("expected error for invalid address") + } +} + +func TestMulticastTransportNonMulticast(t *testing.T) { + cfg := config.MulticastConfig{ + Enabled: true, + Address: "192.168.1.1", + TTL: 4, + } + _, err := NewMulticastTransport(cfg) + if err == nil { + t.Error("expected error for non-multicast address") + } +} + +func TestMulticastTransportSendRTP(t *testing.T) { + InitMulticastPorts(44000) + cfg := config.MulticastConfig{ + Enabled: true, + Address: "239.0.0.3", + TTL: 1, + } + + mt, err := NewMulticastTransport(cfg) + if err != nil { + t.Fatalf("NewMulticastTransport: %v", err) + } + defer mt.Close() + + data := []byte{0x80, 0x60, 0x00, 0x01, 0x00, 0x00, 0x00, 0xA0, 0x12, 0x34, 0x56, 0x78} + if err := mt.SendRTP(data); err != nil { + t.Errorf("SendRTP: %v", err) + } + if err := mt.SendRTCP(data); err != nil { + t.Errorf("SendRTCP: %v", err) + } +} + +func TestMulticastTransportDoubleClose(t *testing.T) { + InitMulticastPorts(46000) + cfg := config.MulticastConfig{ + Enabled: true, + Address: "239.0.0.4", + TTL: 1, + } + + mt, err := NewMulticastTransport(cfg) + if err != nil { + t.Fatalf("NewMulticastTransport: %v", err) + } + + mt.Close() + mt.Close() // should not panic +} + +func TestParseTransportHeaderMulticast(t *testing.T) { + tc := parseTransportHeader("RTP/AVP;multicast") + if !tc.IsMulticast { + t.Error("expected IsMulticast=true") + } + if tc.IsTCP { + t.Error("expected IsTCP=false") + } +} + +func TestHandleSetupMulticastEnabled(t *testing.T) { + mcastCfg := &config.MulticastConfig{ + Enabled: true, + Address: "239.0.0.10", + TTL: 8, + } + InitMulticastPorts(48000) + h := NewHandler(nil, nil, mcastCfg) + session := NewRTSPSession("mcast-test", "live/mcast") + session.Transition(StateDescribed) + + req := &Request{ + Method: "SETUP", + URL: "rtsp://host/live/mcast/trackID=0", + Headers: make(map[string][]string), + } + req.Headers.Set("CSeq", "3") + req.Headers.Set("Transport", "RTP/AVP;multicast") + + resp := h.HandleSetup(req, session, "192.168.1.100:0") + if resp.StatusCode != 200 { + t.Fatalf("expected 200, got %d %s", resp.StatusCode, resp.Reason) + } + + transport := resp.Headers.Get("Transport") + if transport == "" { + t.Fatal("missing Transport header") + } + if !contains(transport, "multicast") { + t.Errorf("expected multicast in Transport: %s", transport) + } + if !contains(transport, "239.0.0.10") { + t.Errorf("expected multicast address in Transport: %s", transport) + } + + if len(session.Tracks) != 1 { + t.Fatalf("expected 1 track, got %d", len(session.Tracks)) + } + if session.Tracks[0].Multicast == nil { + t.Error("expected Multicast transport on track") + } + session.Tracks[0].Multicast.Close() +} + +func TestHandleSetupMulticastDisabled(t *testing.T) { + h := NewHandler(nil, nil, nil) + session := NewRTSPSession("no-mcast", "live/nomcast") + session.Transition(StateDescribed) + + req := &Request{ + Method: "SETUP", + URL: "rtsp://host/live/nomcast/trackID=0", + Headers: make(map[string][]string), + } + req.Headers.Set("CSeq", "3") + req.Headers.Set("Transport", "RTP/AVP;multicast") + + resp := h.HandleSetup(req, session, "192.168.1.100:0") + if resp.StatusCode != 461 { + t.Errorf("expected 461 Unsupported Transport, got %d", resp.StatusCode) + } +} + +func contains(s, substr string) bool { + return len(s) >= len(substr) && (s == substr || len(s) > 0 && containsHelper(s, substr)) +} + +func containsHelper(s, substr string) bool { + for i := 0; i <= len(s)-len(substr); i++ { + if s[i:i+len(substr)] == substr { + return true + } + } + return false +} diff --git a/module/rtsp/server.go b/module/rtsp/server.go index f977cd6..b0107ce 100644 --- a/module/rtsp/server.go +++ b/module/rtsp/server.go @@ -10,6 +10,7 @@ import ( "sync" "time" + "github.com/im-pingo/liveforge/config" "github.com/im-pingo/liveforge/core" "github.com/im-pingo/liveforge/pkg/avframe" "github.com/im-pingo/liveforge/pkg/portalloc" @@ -48,7 +49,13 @@ func (m *Module) Init(s *core.Server) error { m.ports, _ = portalloc.New(30000, 40000) // default range } - m.handler = NewHandler(s, m.ports) + var mcast *config.MulticastConfig + if cfg.Multicast.Enabled { + InitMulticastPorts(cfg.Multicast.BasePort) + mcast = &cfg.Multicast + } + + m.handler = NewHandler(s, m.ports, mcast) ln, err := s.MakeListener(cfg.Listen, cfg.TLS) if err != nil { @@ -269,13 +276,16 @@ func (m *Module) runSubscriberLoop(conn net.Conn, session *RTSPSession) { return } - // Determine video/audio channels and optional UDP transports from track setup. + // Determine video/audio channels and optional UDP/multicast transports from track setup. var videoChannel, audioChannel uint8 var videoUDP, audioUDP *UDPTransport + var videoMcast, audioMcast *MulticastTransport for _, t := range session.Tracks { if t.Codec.IsVideo() { if t.Transport.IsTCP { videoChannel = uint8(t.Transport.Interleaved[0]) + } else if t.Multicast != nil { + videoMcast = t.Multicast } else { videoUDP = t.UDP } @@ -283,6 +293,8 @@ func (m *Module) runSubscriberLoop(conn net.Conn, session *RTSPSession) { if t.Codec.IsAudio() { if t.Transport.IsTCP { audioChannel = uint8(t.Transport.Interleaved[0]) + } else if t.Multicast != nil { + audioMcast = t.Multicast } else { audioUDP = t.UDP } @@ -296,6 +308,8 @@ func (m *Module) runSubscriberLoop(conn net.Conn, session *RTSPSession) { } sub.videoUDP = videoUDP sub.audioUDP = audioUDP + sub.videoMulticast = videoMcast + sub.audioMulticast = audioMcast session.Subscriber = sub if err := session.Stream.AddSubscriber("rtsp"); err != nil { slog.Warn("subscriber limit reached", "module", "rtsp", "session", session.ID, "error", err) diff --git a/module/rtsp/session.go b/module/rtsp/session.go index b451517..d0f4a60 100644 --- a/module/rtsp/session.go +++ b/module/rtsp/session.go @@ -40,7 +40,8 @@ type TrackSetup struct { TrackID int Codec avframe.CodecType Transport TransportConfig - UDP *UDPTransport // non-nil for UDP transport + UDP *UDPTransport // non-nil for UDP unicast + Multicast *MulticastTransport // non-nil for UDP multicast } // RTSPSession represents an RTSP session with state management. diff --git a/module/rtsp/subscriber.go b/module/rtsp/subscriber.go index db68c72..1aa6365 100644 --- a/module/rtsp/subscriber.go +++ b/module/rtsp/subscriber.go @@ -14,7 +14,7 @@ import ( pioncodecs "github.com/pion/rtp/v2/codecs" ) -// RTSPSubscriber implements RTSP playback via TCP interleaved or UDP transport. +// RTSPSubscriber implements RTSP playback via TCP interleaved, UDP unicast, or UDP multicast. type RTSPSubscriber struct { id string options core.SubscribeOptions @@ -29,10 +29,14 @@ type RTSPSubscriber struct { videoChannel uint8 // TCP interleaved channel for video RTP audioChannel uint8 // TCP interleaved channel for audio RTP - // UDP transport (non-nil when client negotiated UDP transport) + // UDP unicast transport (non-nil when client negotiated UDP transport) videoUDP *UDPTransport audioUDP *UDPTransport + // UDP multicast transport (non-nil when client negotiated multicast) + videoMulticast *MulticastTransport + audioMulticast *MulticastTransport + prevVideoDTS int64 prevAudioDTS int64 videoDTSInitialized bool @@ -184,7 +188,7 @@ func (s *RTSPSubscriber) sendVideo(frame *avframe.AVFrame) error { } pkts := s.videoPacketizer.Packetize(payload, samples) - return s.sendPackets(pkts, s.videoChannel, s.videoUDP) + return s.sendPackets(pkts, s.videoChannel, s.videoUDP, s.videoMulticast) } func (s *RTSPSubscriber) sendAudio(frame *avframe.AVFrame) error { @@ -197,7 +201,7 @@ func (s *RTSPSubscriber) sendAudio(frame *avframe.AVFrame) error { s.prevAudioDTS = frame.DTS s.audioDTSInitialized = true pkts := s.audioPacketizer.Packetize(frame.Payload, samples) - return s.sendPackets(pkts, s.audioChannel, s.audioUDP) + return s.sendPackets(pkts, s.audioChannel, s.audioUDP, s.audioMulticast) } // Fallback to custom packetizer @@ -209,10 +213,10 @@ func (s *RTSPSubscriber) sendAudio(frame *avframe.AVFrame) error { return err } wrapped := s.customAudioSes.WrapPackets(pkts, frame.DTS) - return s.sendPackets(wrapped, s.audioChannel, s.audioUDP) + return s.sendPackets(wrapped, s.audioChannel, s.audioUDP, s.audioMulticast) } -func (s *RTSPSubscriber) sendPackets(pkts []*pionrtp.Packet, channel uint8, udp *UDPTransport) error { +func (s *RTSPSubscriber) sendPackets(pkts []*pionrtp.Packet, channel uint8, udp *UDPTransport, mcast *MulticastTransport) error { for _, pkt := range pkts { data, err := pkt.Marshal() if err != nil { @@ -222,7 +226,11 @@ func (s *RTSPSubscriber) sendPackets(pkts []*pionrtp.Packet, channel uint8, udp s.packetCount.Add(1) s.octetCount.Add(uint32(len(pkt.Payload))) s.lastRTPTime.Store(pkt.Timestamp) - if udp != nil { + if mcast != nil { + if err := mcast.SendRTP(data); err != nil { + return err + } + } else if udp != nil { if err := udp.SendRTP(data); err != nil { return err } @@ -255,7 +263,9 @@ func (s *RTSPSubscriber) rtcpLoop() { ntpTime := toNTP(time.Now()) sr := pkgrtp.BuildSR(s.videoSSRC, ntpTime, rtpTime, pktCount, octCount) - if s.videoUDP != nil { + if s.videoMulticast != nil { + s.videoMulticast.SendRTCP(sr) + } else if s.videoUDP != nil { s.videoUDP.SendRTCP(sr) } else { rtcpChannel := s.videoChannel + 1 // odd channel = RTCP diff --git a/module/sip/dispatch_test.go b/module/sip/dispatch_test.go new file mode 100644 index 0000000..ee9751a --- /dev/null +++ b/module/sip/dispatch_test.go @@ -0,0 +1,614 @@ +package sip + +import ( + "sync" + "sync/atomic" + "testing" + + "github.com/emiago/sipgo/sip" +) + +// mockServerTx is a minimal mock implementing sip.ServerTransaction. +// It captures responses passed to Respond for later assertions. +type mockServerTx struct { + mu sync.Mutex + responses []*sip.Response +} + +func newMockServerTx() *mockServerTx { + return &mockServerTx{} +} + +func (m *mockServerTx) Respond(res *sip.Response) error { + m.mu.Lock() + defer m.mu.Unlock() + m.responses = append(m.responses, res) + return nil +} + +func (m *mockServerTx) Acks() <-chan *sip.Request { + ch := make(chan *sip.Request) + return ch +} + +func (m *mockServerTx) OnCancel(f sip.FnTxCancel) bool { + return false +} + +func (m *mockServerTx) Terminate() {} + +func (m *mockServerTx) OnTerminate(f sip.FnTxTerminate) bool { + return false +} + +func (m *mockServerTx) Done() <-chan struct{} { + ch := make(chan struct{}) + return ch +} + +func (m *mockServerTx) Err() error { + return nil +} + +// lastResponse returns the most recent response captured by Respond. +func (m *mockServerTx) lastResponse() *sip.Response { + m.mu.Lock() + defer m.mu.Unlock() + if len(m.responses) == 0 { + return nil + } + return m.responses[len(m.responses)-1] +} + +// responseCount returns how many times Respond was called. +func (m *mockServerTx) responseCount() int { + m.mu.Lock() + defer m.mu.Unlock() + return len(m.responses) +} + +// makeTestRequest creates a minimal SIP request for the given method. +func makeTestRequest(method sip.RequestMethod) *sip.Request { + return sip.NewRequest(method, sip.Uri{Host: "localhost", Port: 5060}) +} + +// --------------------------------------------------------------------------- +// Handler Registration +// --------------------------------------------------------------------------- + +func TestOnRegisterRegistersHandler(t *testing.T) { + s := newService() + called := false + s.OnRegister(func(req *sip.Request, tx sip.ServerTransaction) { + called = true + }) + + s.mu.RLock() + count := len(s.registerHandlers) + s.mu.RUnlock() + + if count != 1 { + t.Fatalf("expected 1 register handler, got %d", count) + } + + tx := newMockServerTx() + s.dispatchRegister(makeTestRequest(sip.REGISTER), tx) + if !called { + t.Fatal("register handler was not called") + } +} + +func TestOnInviteRegistersHandler(t *testing.T) { + s := newService() + called := false + s.OnInvite(func(req *sip.Request, tx sip.ServerTransaction) { + called = true + }) + + s.mu.RLock() + count := len(s.inviteHandlers) + s.mu.RUnlock() + + if count != 1 { + t.Fatalf("expected 1 invite handler, got %d", count) + } + + tx := newMockServerTx() + s.dispatchInvite(makeTestRequest(sip.INVITE), tx) + if !called { + t.Fatal("invite handler was not called") + } +} + +func TestOnByeRegistersHandler(t *testing.T) { + s := newService() + called := false + s.OnBye(func(req *sip.Request, tx sip.ServerTransaction) { + called = true + }) + + s.mu.RLock() + count := len(s.byeHandlers) + s.mu.RUnlock() + + if count != 1 { + t.Fatalf("expected 1 bye handler, got %d", count) + } + + tx := newMockServerTx() + s.dispatchBye(makeTestRequest(sip.BYE), tx) + if !called { + t.Fatal("bye handler was not called") + } +} + +func TestOnMessageRegistersHandler(t *testing.T) { + s := newService() + called := false + s.OnMessage(func(req *sip.Request, tx sip.ServerTransaction) { + called = true + }) + + s.mu.RLock() + count := len(s.messageHandlers) + s.mu.RUnlock() + + if count != 1 { + t.Fatalf("expected 1 message handler, got %d", count) + } + + tx := newMockServerTx() + s.dispatchMessage(makeTestRequest(sip.MESSAGE), tx) + if !called { + t.Fatal("message handler was not called") + } +} + +func TestOnSubscribeRegistersHandler(t *testing.T) { + s := newService() + called := false + s.OnSubscribe(func(req *sip.Request, tx sip.ServerTransaction) { + called = true + }) + + s.mu.RLock() + count := len(s.subscribeHandlers) + s.mu.RUnlock() + + if count != 1 { + t.Fatalf("expected 1 subscribe handler, got %d", count) + } + + tx := newMockServerTx() + s.dispatchSubscribe(makeTestRequest(sip.SUBSCRIBE), tx) + if !called { + t.Fatal("subscribe handler was not called") + } +} + +func TestOnNotifyRegistersHandler(t *testing.T) { + s := newService() + called := false + s.OnNotify(func(req *sip.Request, tx sip.ServerTransaction) { + called = true + }) + + s.mu.RLock() + count := len(s.notifyHandlers) + s.mu.RUnlock() + + if count != 1 { + t.Fatalf("expected 1 notify handler, got %d", count) + } + + tx := newMockServerTx() + s.dispatchNotify(makeTestRequest(sip.NOTIFY), tx) + if !called { + t.Fatal("notify handler was not called") + } +} + +// --------------------------------------------------------------------------- +// Dispatch Without Handlers (default responses) +// --------------------------------------------------------------------------- + +func TestDispatchWithoutHandlers(t *testing.T) { + tests := []struct { + name string + method sip.RequestMethod + dispatch func(s *service, req *sip.Request, tx sip.ServerTransaction) + wantStatus int + wantReason string + }{ + { + name: "REGISTER without handler returns 405", + method: sip.REGISTER, + dispatch: func(s *service, req *sip.Request, tx sip.ServerTransaction) { s.dispatchRegister(req, tx) }, + wantStatus: 405, + wantReason: "Method Not Allowed", + }, + { + name: "INVITE without handler returns 405", + method: sip.INVITE, + dispatch: func(s *service, req *sip.Request, tx sip.ServerTransaction) { s.dispatchInvite(req, tx) }, + wantStatus: 405, + wantReason: "Method Not Allowed", + }, + { + name: "BYE without handler returns 200", + method: sip.BYE, + dispatch: func(s *service, req *sip.Request, tx sip.ServerTransaction) { s.dispatchBye(req, tx) }, + wantStatus: 200, + wantReason: "OK", + }, + { + name: "MESSAGE without handler returns 200", + method: sip.MESSAGE, + dispatch: func(s *service, req *sip.Request, tx sip.ServerTransaction) { s.dispatchMessage(req, tx) }, + wantStatus: 200, + wantReason: "OK", + }, + { + name: "SUBSCRIBE without handler returns 405", + method: sip.SUBSCRIBE, + dispatch: func(s *service, req *sip.Request, tx sip.ServerTransaction) { s.dispatchSubscribe(req, tx) }, + wantStatus: 405, + wantReason: "Method Not Allowed", + }, + { + name: "NOTIFY without handler returns 200", + method: sip.NOTIFY, + dispatch: func(s *service, req *sip.Request, tx sip.ServerTransaction) { s.dispatchNotify(req, tx) }, + wantStatus: 200, + wantReason: "OK", + }, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + s := newService() + tx := newMockServerTx() + req := makeTestRequest(tc.method) + + tc.dispatch(s, req, tx) + + resp := tx.lastResponse() + if resp == nil { + t.Fatal("expected a response, got nil") + } + if resp.StatusCode != tc.wantStatus { + t.Errorf("status code = %d, want %d", resp.StatusCode, tc.wantStatus) + } + if resp.Reason != tc.wantReason { + t.Errorf("reason = %q, want %q", resp.Reason, tc.wantReason) + } + }) + } +} + +// --------------------------------------------------------------------------- +// Dispatch With Handlers +// --------------------------------------------------------------------------- + +func TestDispatchWithHandlers(t *testing.T) { + tests := []struct { + name string + method sip.RequestMethod + register func(s *service, called *bool) + dispatch func(s *service, req *sip.Request, tx sip.ServerTransaction) + }{ + { + name: "REGISTER dispatch calls handler", + method: sip.REGISTER, + register: func(s *service, called *bool) { + s.OnRegister(func(req *sip.Request, tx sip.ServerTransaction) { *called = true }) + }, + dispatch: func(s *service, req *sip.Request, tx sip.ServerTransaction) { s.dispatchRegister(req, tx) }, + }, + { + name: "INVITE dispatch calls handler", + method: sip.INVITE, + register: func(s *service, called *bool) { + s.OnInvite(func(req *sip.Request, tx sip.ServerTransaction) { *called = true }) + }, + dispatch: func(s *service, req *sip.Request, tx sip.ServerTransaction) { s.dispatchInvite(req, tx) }, + }, + { + name: "BYE dispatch calls handler", + method: sip.BYE, + register: func(s *service, called *bool) { + s.OnBye(func(req *sip.Request, tx sip.ServerTransaction) { *called = true }) + }, + dispatch: func(s *service, req *sip.Request, tx sip.ServerTransaction) { s.dispatchBye(req, tx) }, + }, + { + name: "MESSAGE dispatch calls handler", + method: sip.MESSAGE, + register: func(s *service, called *bool) { + s.OnMessage(func(req *sip.Request, tx sip.ServerTransaction) { *called = true }) + }, + dispatch: func(s *service, req *sip.Request, tx sip.ServerTransaction) { s.dispatchMessage(req, tx) }, + }, + { + name: "SUBSCRIBE dispatch calls handler", + method: sip.SUBSCRIBE, + register: func(s *service, called *bool) { + s.OnSubscribe(func(req *sip.Request, tx sip.ServerTransaction) { *called = true }) + }, + dispatch: func(s *service, req *sip.Request, tx sip.ServerTransaction) { s.dispatchSubscribe(req, tx) }, + }, + { + name: "NOTIFY dispatch calls handler", + method: sip.NOTIFY, + register: func(s *service, called *bool) { + s.OnNotify(func(req *sip.Request, tx sip.ServerTransaction) { *called = true }) + }, + dispatch: func(s *service, req *sip.Request, tx sip.ServerTransaction) { s.dispatchNotify(req, tx) }, + }, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + s := newService() + called := false + tc.register(s, &called) + + tx := newMockServerTx() + req := makeTestRequest(tc.method) + tc.dispatch(s, req, tx) + + if !called { + t.Fatal("handler was not called") + } + // When a handler is registered, the service should not send a default response. + if tx.responseCount() != 0 { + t.Errorf("expected no auto-response when handler is registered, got %d responses", tx.responseCount()) + } + }) + } +} + +// --------------------------------------------------------------------------- +// Multiple Handlers +// --------------------------------------------------------------------------- + +func TestMultipleHandlersAllCalled(t *testing.T) { + tests := []struct { + name string + method sip.RequestMethod + register func(s *service, counters []*int32) + dispatch func(s *service, req *sip.Request, tx sip.ServerTransaction) + }{ + { + name: "REGISTER multiple handlers", + method: sip.REGISTER, + register: func(s *service, counters []*int32) { + for _, c := range counters { + c := c + s.OnRegister(func(req *sip.Request, tx sip.ServerTransaction) { atomic.AddInt32(c, 1) }) + } + }, + dispatch: func(s *service, req *sip.Request, tx sip.ServerTransaction) { s.dispatchRegister(req, tx) }, + }, + { + name: "INVITE multiple handlers", + method: sip.INVITE, + register: func(s *service, counters []*int32) { + for _, c := range counters { + c := c + s.OnInvite(func(req *sip.Request, tx sip.ServerTransaction) { atomic.AddInt32(c, 1) }) + } + }, + dispatch: func(s *service, req *sip.Request, tx sip.ServerTransaction) { s.dispatchInvite(req, tx) }, + }, + { + name: "BYE multiple handlers", + method: sip.BYE, + register: func(s *service, counters []*int32) { + for _, c := range counters { + c := c + s.OnBye(func(req *sip.Request, tx sip.ServerTransaction) { atomic.AddInt32(c, 1) }) + } + }, + dispatch: func(s *service, req *sip.Request, tx sip.ServerTransaction) { s.dispatchBye(req, tx) }, + }, + { + name: "MESSAGE multiple handlers", + method: sip.MESSAGE, + register: func(s *service, counters []*int32) { + for _, c := range counters { + c := c + s.OnMessage(func(req *sip.Request, tx sip.ServerTransaction) { atomic.AddInt32(c, 1) }) + } + }, + dispatch: func(s *service, req *sip.Request, tx sip.ServerTransaction) { s.dispatchMessage(req, tx) }, + }, + { + name: "SUBSCRIBE multiple handlers", + method: sip.SUBSCRIBE, + register: func(s *service, counters []*int32) { + for _, c := range counters { + c := c + s.OnSubscribe(func(req *sip.Request, tx sip.ServerTransaction) { atomic.AddInt32(c, 1) }) + } + }, + dispatch: func(s *service, req *sip.Request, tx sip.ServerTransaction) { s.dispatchSubscribe(req, tx) }, + }, + { + name: "NOTIFY multiple handlers", + method: sip.NOTIFY, + register: func(s *service, counters []*int32) { + for _, c := range counters { + c := c + s.OnNotify(func(req *sip.Request, tx sip.ServerTransaction) { atomic.AddInt32(c, 1) }) + } + }, + dispatch: func(s *service, req *sip.Request, tx sip.ServerTransaction) { s.dispatchNotify(req, tx) }, + }, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + s := newService() + const handlerCount = 3 + counters := make([]*int32, handlerCount) + for i := range counters { + counters[i] = new(int32) + } + + tc.register(s, counters) + + tx := newMockServerTx() + req := makeTestRequest(tc.method) + tc.dispatch(s, req, tx) + + for i, c := range counters { + if v := atomic.LoadInt32(c); v != 1 { + t.Errorf("handler %d called %d times, want 1", i, v) + } + } + }) + } +} + +// --------------------------------------------------------------------------- +// Concurrent Handler Registration +// --------------------------------------------------------------------------- + +func TestConcurrentHandlerRegistration(t *testing.T) { + s := newService() + const goroutines = 50 + + var wg sync.WaitGroup + wg.Add(goroutines * 6) + + for i := 0; i < goroutines; i++ { + go func() { + defer wg.Done() + s.OnRegister(func(req *sip.Request, tx sip.ServerTransaction) {}) + }() + go func() { + defer wg.Done() + s.OnInvite(func(req *sip.Request, tx sip.ServerTransaction) {}) + }() + go func() { + defer wg.Done() + s.OnBye(func(req *sip.Request, tx sip.ServerTransaction) {}) + }() + go func() { + defer wg.Done() + s.OnMessage(func(req *sip.Request, tx sip.ServerTransaction) {}) + }() + go func() { + defer wg.Done() + s.OnSubscribe(func(req *sip.Request, tx sip.ServerTransaction) {}) + }() + go func() { + defer wg.Done() + s.OnNotify(func(req *sip.Request, tx sip.ServerTransaction) {}) + }() + } + + wg.Wait() + + s.mu.RLock() + defer s.mu.RUnlock() + + if got := len(s.registerHandlers); got != goroutines { + t.Errorf("register handlers = %d, want %d", got, goroutines) + } + if got := len(s.inviteHandlers); got != goroutines { + t.Errorf("invite handlers = %d, want %d", got, goroutines) + } + if got := len(s.byeHandlers); got != goroutines { + t.Errorf("bye handlers = %d, want %d", got, goroutines) + } + if got := len(s.messageHandlers); got != goroutines { + t.Errorf("message handlers = %d, want %d", got, goroutines) + } + if got := len(s.subscribeHandlers); got != goroutines { + t.Errorf("subscribe handlers = %d, want %d", got, goroutines) + } + if got := len(s.notifyHandlers); got != goroutines { + t.Errorf("notify handlers = %d, want %d", got, goroutines) + } +} + +// --------------------------------------------------------------------------- +// Service Accessors +// --------------------------------------------------------------------------- + +func TestServiceAccessors(t *testing.T) { + s := newService() + s.localAddr = "0.0.0.0:5060" + s.serverID = "34020000002000000001" + s.domain = "3402000000" + + if got := s.LocalAddr(); got != "0.0.0.0:5060" { + t.Errorf("LocalAddr() = %q, want %q", got, "0.0.0.0:5060") + } + if got := s.ServerID(); got != "34020000002000000001" { + t.Errorf("ServerID() = %q, want %q", got, "34020000002000000001") + } + if got := s.Domain(); got != "3402000000" { + t.Errorf("Domain() = %q, want %q", got, "3402000000") + } +} + +func TestServiceAccessorsDefault(t *testing.T) { + s := newService() + + if got := s.LocalAddr(); got != "" { + t.Errorf("LocalAddr() on new service = %q, want empty string", got) + } + if got := s.ServerID(); got != "" { + t.Errorf("ServerID() on new service = %q, want empty string", got) + } + if got := s.Domain(); got != "" { + t.Errorf("Domain() on new service = %q, want empty string", got) + } +} + +// --------------------------------------------------------------------------- +// Dispatch Does Not Send Auto-Response When Handler Present +// --------------------------------------------------------------------------- + +func TestDispatchHandlerReceivesCorrectRequest(t *testing.T) { + s := newService() + var received *sip.Request + + s.OnRegister(func(req *sip.Request, tx sip.ServerTransaction) { + received = req + }) + + tx := newMockServerTx() + req := makeTestRequest(sip.REGISTER) + s.dispatchRegister(req, tx) + + if received == nil { + t.Fatal("handler did not receive the request") + } + if received != req { + t.Error("handler received a different request object than what was dispatched") + } +} + +func TestDispatchHandlerReceivesCorrectTransaction(t *testing.T) { + s := newService() + var receivedTx sip.ServerTransaction + + s.OnInvite(func(req *sip.Request, tx sip.ServerTransaction) { + receivedTx = tx + }) + + tx := newMockServerTx() + req := makeTestRequest(sip.INVITE) + s.dispatchInvite(req, tx) + + if receivedTx == nil { + t.Fatal("handler did not receive the transaction") + } + if receivedTx != tx { + t.Error("handler received a different transaction object than what was dispatched") + } +} diff --git a/module/sip/transaction_test.go b/module/sip/transaction_test.go new file mode 100644 index 0000000..442a5ef --- /dev/null +++ b/module/sip/transaction_test.go @@ -0,0 +1,154 @@ +package sip + +import ( + "testing" + + "github.com/emiago/sipgo/sip" +) + +func makeTestInviteAndResponse() (*sip.Request, *sip.Response) { + invite := sip.NewRequest(sip.INVITE, sip.Uri{User: "bob", Host: "example.com"}) + invite.SipVersion = "SIP/2.0" + + fromHeader := &sip.FromHeader{ + DisplayName: "Alice", + Address: sip.Uri{User: "alice", Host: "example.com"}, + Params: sip.NewParams(), + } + fromHeader.Params.Add("tag", "from-tag-123") + invite.AppendHeader(fromHeader) + + toHeader := &sip.ToHeader{ + Address: sip.Uri{User: "bob", Host: "example.com"}, + } + invite.AppendHeader(toHeader) + + callID := sip.CallIDHeader("call-id-abc@example.com") + invite.AppendHeader(&callID) + + invite.AppendHeader(&sip.CSeqHeader{SeqNo: 1, MethodName: sip.INVITE}) + + resp := sip.NewResponseFromRequest(invite, 200, "OK", nil) + if to := resp.To(); to != nil { + to.Params.Add("tag", "to-tag-456") + } + + return invite, resp +} + +func TestBuildACKHeaders(t *testing.T) { + invite, resp := makeTestInviteAndResponse() + ack := buildACK(invite, resp) + + if ack.Method != sip.ACK { + t.Errorf("expected ACK method, got %s", ack.Method) + } + + from := ack.From() + if from == nil { + t.Fatal("ACK missing From header") + } + tag, _ := from.Params.Get("tag") + if tag != "from-tag-123" { + t.Errorf("expected from-tag-123, got %s", tag) + } + + to := ack.To() + if to == nil { + t.Fatal("ACK missing To header") + } + toTag, _ := to.Params.Get("tag") + if toTag != "to-tag-456" { + t.Errorf("expected to-tag-456, got %s", toTag) + } + + if ack.CallID() == nil { + t.Fatal("ACK missing Call-ID header") + } + if ack.CallID().Value() != "call-id-abc@example.com" { + t.Errorf("expected call-id-abc@example.com, got %s", ack.CallID().Value()) + } + + cseq := ack.CSeq() + if cseq == nil { + t.Fatal("ACK missing CSeq header") + } + if cseq.SeqNo != 1 { + t.Errorf("ACK CSeq should be 1, got %d", cseq.SeqNo) + } + if cseq.MethodName != sip.ACK { + t.Errorf("ACK CSeq method should be ACK, got %s", cseq.MethodName) + } +} + +func TestBuildBYEHeaders(t *testing.T) { + invite, resp := makeTestInviteAndResponse() + bye := buildBYE(invite, resp) + + if bye.Method != sip.BYE { + t.Errorf("expected BYE method, got %s", bye.Method) + } + + cseq := bye.CSeq() + if cseq == nil { + t.Fatal("BYE missing CSeq header") + } + if cseq.SeqNo != 2 { + t.Errorf("BYE CSeq should be INVITE CSeq+1=2, got %d", cseq.SeqNo) + } + if cseq.MethodName != sip.BYE { + t.Errorf("BYE CSeq method should be BYE, got %s", cseq.MethodName) + } +} + +func TestInviteTransactionClose(t *testing.T) { + tx := &InviteTransaction{ + done: make(chan struct{}), + } + + select { + case <-tx.Done(): + t.Fatal("done channel should not be closed yet") + default: + } + + tx.Close() + + select { + case <-tx.Done(): + default: + t.Fatal("done channel should be closed after Close()") + } + + // Double close should not panic + tx.Close() +} + +func TestInviteTransactionResponseNilBeforeSet(t *testing.T) { + tx := &InviteTransaction{ + done: make(chan struct{}), + } + if tx.Response() != nil { + t.Error("response should be nil before being set") + } +} + +func TestInviteTransactionSendACKNoResponse(t *testing.T) { + tx := &InviteTransaction{ + done: make(chan struct{}), + } + err := tx.SendACK(t.Context()) + if err == nil { + t.Error("expected error when no response") + } +} + +func TestInviteTransactionSendBYENoResponse(t *testing.T) { + tx := &InviteTransaction{ + done: make(chan struct{}), + } + err := tx.SendBYE(t.Context()) + if err == nil { + t.Error("expected error when no dialog established") + } +} diff --git a/module/sipgateway/call_session.go b/module/sipgateway/call_session.go new file mode 100644 index 0000000..20d4b41 --- /dev/null +++ b/module/sipgateway/call_session.go @@ -0,0 +1,256 @@ +package sipgateway + +import ( + "log/slog" + "net" + "sync" + "time" + + "github.com/im-pingo/liveforge/core" + "github.com/im-pingo/liveforge/pkg/avframe" + "github.com/im-pingo/liveforge/pkg/rtp" + pionrtp "github.com/pion/rtp/v2" +) + +// CallSession manages a single SIP call with RTP bridging to/from a stream. +type CallSession struct { + callID string + streamKey string + codec negotiatedCodec + direction string // "inbound" or "outbound" + + rtpPort int + rtcpPort int + conn *net.UDPConn + remoteAddr *net.UDPAddr + + stream *core.Stream + publisher *sipPublisher + + mu sync.Mutex + closed chan struct{} +} + +type sipPublisher struct { + id string + info *avframe.MediaInfo +} + +func (p *sipPublisher) ID() string { return p.id } +func (p *sipPublisher) MediaInfo() *avframe.MediaInfo { return p.info } +func (p *sipPublisher) Close() error { return nil } + +func newCallSession(callID, streamKey string, codec negotiatedCodec, direction string, rtpPort, rtcpPort int) *CallSession { + return &CallSession{ + callID: callID, + streamKey: streamKey, + codec: codec, + direction: direction, + rtpPort: rtpPort, + rtcpPort: rtcpPort, + closed: make(chan struct{}), + } +} + +func (cs *CallSession) startInbound(stream *core.Stream, remoteIP string, remotePort int) error { + cs.stream = stream + + cs.publisher = &sipPublisher{ + id: "sip-" + cs.callID, + info: &avframe.MediaInfo{ + AudioCodec: cs.codec.Codec, + SampleRate: cs.codec.ClockRate, + Channels: 1, + }, + } + stream.SetPublisher(cs.publisher) + + addr := &net.UDPAddr{Port: cs.rtpPort} + conn, err := net.ListenUDP("udp", addr) + if err != nil { + return err + } + cs.conn = conn + + if remoteIP != "" && remotePort > 0 { + cs.remoteAddr = &net.UDPAddr{IP: net.ParseIP(remoteIP), Port: remotePort} + } + + go cs.receiveLoop() + return nil +} + +func (cs *CallSession) startOutbound(stream *core.Stream, remoteIP string, remotePort int) error { + cs.stream = stream + cs.remoteAddr = &net.UDPAddr{IP: net.ParseIP(remoteIP), Port: remotePort} + + addr := &net.UDPAddr{Port: cs.rtpPort} + conn, err := net.ListenUDP("udp", addr) + if err != nil { + return err + } + cs.conn = conn + + go cs.sendLoop() + return nil +} + +func (cs *CallSession) receiveLoop() { + defer slog.Info("rtp receive loop stopped", "module", "sipgateway", "call", cs.callID) + + depacketizer := cs.newDepacketizer() + buf := make([]byte, 2048) + + for { + cs.conn.SetReadDeadline(time.Now().Add(30 * time.Second)) + n, _, err := cs.conn.ReadFromUDP(buf) + if err != nil { + select { + case <-cs.closed: + return + default: + } + if netErr, ok := err.(net.Error); ok && netErr.Timeout() { + continue + } + slog.Warn("rtp read error", "module", "sipgateway", "call", cs.callID, "error", err) + return + } + + if n < 12 { + continue + } + + var pkt pionrtp.Packet + if err := pkt.Unmarshal(buf[:n]); err != nil { + continue + } + + frame, err := depacketizer.depacketize(&pkt) + if err != nil || frame == nil { + continue + } + + cs.stream.WriteFrame(frame) + } +} + +func (cs *CallSession) sendLoop() { + defer slog.Info("rtp send loop stopped", "module", "sipgateway", "call", cs.callID) + + session := rtp.NewSession(uint8(cs.codec.PT), uint32(cs.codec.ClockRate)) + packetizer := cs.newPacketizer() + + cs.stream.AddSubscriber("sipgateway") + defer cs.stream.RemoveSubscriber("sipgateway") + + reader := cs.stream.RingBuffer().NewReader() + rb := cs.stream.RingBuffer() + + for { + select { + case <-cs.closed: + return + default: + } + + frame, ok := reader.TryRead() + if !ok { + if rb.IsClosed() { + return + } + select { + case <-cs.closed: + return + case <-rb.Signal(): + } + continue + } + + if frame.MediaType != avframe.MediaTypeAudio { + continue + } + + pkts, err := packetizer.packetize(frame) + if err != nil || len(pkts) == 0 { + continue + } + + session.WrapPackets(pkts, frame.DTS) + + for _, pkt := range pkts { + data, err := pkt.Marshal() + if err != nil { + continue + } + cs.conn.WriteToUDP(data, cs.remoteAddr) + } + } +} + +func (cs *CallSession) Close() { + cs.mu.Lock() + defer cs.mu.Unlock() + select { + case <-cs.closed: + return + default: + close(cs.closed) + } + if cs.conn != nil { + cs.conn.Close() + } +} + +type audioDepacketizer struct { + codec avframe.CodecType + inner interface { + Depacketize(pkt *pionrtp.Packet) (*avframe.AVFrame, error) + } +} + +func (cs *CallSession) newDepacketizer() *audioDepacketizer { + d := &audioDepacketizer{codec: cs.codec.Codec} + switch cs.codec.Codec { + case avframe.CodecG711U: + d.inner = &rtp.G711Depacketizer{Codec: avframe.CodecG711U} + case avframe.CodecG711A: + d.inner = &rtp.G711Depacketizer{Codec: avframe.CodecG711A} + case avframe.CodecOpus: + d.inner = &rtp.OpusDepacketizer{} + case avframe.CodecG722: + d.inner = &rtp.G722Depacketizer{} + default: + d.inner = &rtp.G711Depacketizer{Codec: avframe.CodecG711U} + } + return d +} + +func (d *audioDepacketizer) depacketize(pkt *pionrtp.Packet) (*avframe.AVFrame, error) { + return d.inner.Depacketize(pkt) +} + +type audioPacketizer struct { + inner interface { + Packetize(frame *avframe.AVFrame, mtu int) ([]*pionrtp.Packet, error) + } +} + +func (cs *CallSession) newPacketizer() *audioPacketizer { + p := &audioPacketizer{} + switch cs.codec.Codec { + case avframe.CodecG711U, avframe.CodecG711A: + p.inner = &rtp.G711Packetizer{} + case avframe.CodecOpus: + p.inner = &rtp.OpusPacketizer{} + case avframe.CodecG722: + p.inner = &rtp.G722Packetizer{} + default: + p.inner = &rtp.G711Packetizer{} + } + return p +} + +func (p *audioPacketizer) packetize(frame *avframe.AVFrame) ([]*pionrtp.Packet, error) { + return p.inner.Packetize(frame, 1400) +} diff --git a/module/sipgateway/codec.go b/module/sipgateway/codec.go new file mode 100644 index 0000000..de4e202 --- /dev/null +++ b/module/sipgateway/codec.go @@ -0,0 +1,203 @@ +package sipgateway + +import ( + "strings" + + "github.com/im-pingo/liveforge/pkg/avframe" + "github.com/im-pingo/liveforge/pkg/sdp" +) + +type codecInfo struct { + Codec avframe.CodecType + ClockRate int + PT int +} + +var encodingToCodec = map[string]codecInfo{ + "PCMU": {avframe.CodecG711U, 8000, 0}, + "PCMA": {avframe.CodecG711A, 8000, 8}, + "G722": {avframe.CodecG722, 8000, 9}, + "G729": {avframe.CodecG729, 8000, 18}, + "opus": {avframe.CodecOpus, 48000, 111}, + "speex": {avframe.CodecSpeex, 16000, 102}, + "MPEG4-GENERIC": {avframe.CodecAAC, 44100, 101}, +} + +type negotiatedCodec struct { + Codec avframe.CodecType + PT int + ClockRate int + EncodingName string +} + +func negotiateCodec(offer *sdp.MediaDescription, preferred []string) (negotiatedCodec, bool) { + priorityMap := make(map[string]int) + for i, name := range preferred { + priorityMap[strings.ToUpper(name)] = len(preferred) - i + } + + var best negotiatedCodec + bestPriority := -1 + + for _, pt := range offer.Formats { + rm := offer.RTPMap(pt) + if rm == nil { + if info, ok := staticPT(pt); ok { + rm = &sdp.RTPMapInfo{ + PayloadType: pt, + EncodingName: info.EncodingName, + ClockRate: info.ClockRate, + } + } else { + continue + } + } + + nameUpper := strings.ToUpper(rm.EncodingName) + info, supported := encodingToCodec[rm.EncodingName] + if !supported { + info, supported = encodingToCodec[nameUpper] + if !supported { + continue + } + } + + prio, inPreferred := priorityMap[nameUpper] + if !inPreferred { + prio = 0 + } + + if prio > bestPriority { + bestPriority = prio + best = negotiatedCodec{ + Codec: info.Codec, + PT: pt, + ClockRate: rm.ClockRate, + EncodingName: rm.EncodingName, + } + } + } + + return best, bestPriority >= 0 +} + +func staticPT(pt int) (struct{ EncodingName string; ClockRate int }, bool) { + type ptInfo struct{ EncodingName string; ClockRate int } + m := map[int]ptInfo{ + 0: {"PCMU", 8000}, + 8: {"PCMA", 8000}, + 9: {"G722", 8000}, + 18: {"G729", 8000}, + } + info, ok := m[pt] + return info, ok +} + +func buildAnswerSDP(serverAddr string, rtpPort int, nc negotiatedCodec) []byte { + sd := &sdp.SessionDescription{ + Version: 0, + Origin: sdp.Origin{ + Username: "-", + SessionID: "1", + SessionVersion: "1", + NetType: "IN", + AddrType: "IP4", + Address: serverAddr, + }, + Name: "LiveForge SIP Gateway", + Connection: &sdp.Connection{ + NetType: "IN", + AddrType: "IP4", + Address: serverAddr, + }, + Timing: sdp.Timing{Start: 0, Stop: 0}, + } + + md := &sdp.MediaDescription{ + Type: "audio", + Port: rtpPort, + Proto: "RTP/AVP", + Formats: []int{nc.PT}, + Attributes: []sdp.Attribute{ + {Key: "rtpmap", Value: rtpmapValue(nc)}, + {Key: "sendrecv"}, + {Key: "ptime", Value: "20"}, + }, + } + + sd.Media = append(sd.Media, md) + return sd.Marshal() +} + +func buildOfferSDP(serverAddr string, rtpPort int, codecs []negotiatedCodec) []byte { + sd := &sdp.SessionDescription{ + Version: 0, + Origin: sdp.Origin{ + Username: "-", + SessionID: "1", + SessionVersion: "1", + NetType: "IN", + AddrType: "IP4", + Address: serverAddr, + }, + Name: "LiveForge SIP Gateway", + Connection: &sdp.Connection{ + NetType: "IN", + AddrType: "IP4", + Address: serverAddr, + }, + Timing: sdp.Timing{Start: 0, Stop: 0}, + } + + var formats []int + var attrs []sdp.Attribute + for _, nc := range codecs { + formats = append(formats, nc.PT) + attrs = append(attrs, sdp.Attribute{Key: "rtpmap", Value: rtpmapValue(nc)}) + } + attrs = append(attrs, sdp.Attribute{Key: "sendrecv"}) + attrs = append(attrs, sdp.Attribute{Key: "ptime", Value: "20"}) + + md := &sdp.MediaDescription{ + Type: "audio", + Port: rtpPort, + Proto: "RTP/AVP", + Formats: formats, + Attributes: attrs, + } + sd.Media = append(sd.Media, md) + return sd.Marshal() +} + +func rtpmapValue(nc negotiatedCodec) string { + if nc.Codec == avframe.CodecOpus { + return strings.Join([]string{ + itoa(nc.PT), " ", nc.EncodingName, "/", itoa(nc.ClockRate), "/2", + }, "") + } + return strings.Join([]string{ + itoa(nc.PT), " ", nc.EncodingName, "/", itoa(nc.ClockRate), + }, "") +} + +func itoa(n int) string { + buf := [20]byte{} + pos := len(buf) + if n == 0 { + return "0" + } + neg := n < 0 + if neg { + n = -n + } + for n > 0 { + pos-- + buf[pos] = byte('0' + n%10) + n /= 10 + } + if neg { + pos-- + buf[pos] = '-' + } + return string(buf[pos:]) +} diff --git a/module/sipgateway/gateway.go b/module/sipgateway/gateway.go new file mode 100644 index 0000000..6b01418 --- /dev/null +++ b/module/sipgateway/gateway.go @@ -0,0 +1,394 @@ +package sipgateway + +import ( + "context" + "fmt" + "log/slog" + "net" + "sync" + + "github.com/emiago/sipgo/sip" + "github.com/im-pingo/liveforge/config" + "github.com/im-pingo/liveforge/core" + sipmod "github.com/im-pingo/liveforge/module/sip" + "github.com/im-pingo/liveforge/pkg/portalloc" + "github.com/im-pingo/liveforge/pkg/sdp" +) + +// Gateway manages SIP-to-stream call bridging. +type Gateway struct { + sipService sipmod.SIPService + hub *core.StreamHub + eventBus *core.EventBus + portAlloc *portalloc.PortAllocator + prefix string + maxCalls int + codecs []string + localIP string + + mu sync.Mutex + sessions map[string]*CallSession +} + +// NewGateway creates and starts a SIP gateway. +func NewGateway(cfg config.SIPGatewayConfig, sipSvc sipmod.SIPService, hub *core.StreamHub, bus *core.EventBus) (*Gateway, error) { + if len(cfg.RTPPortRange) != 2 { + return nil, fmt.Errorf("sipgateway: rtp_port_range must have exactly 2 elements [min, max]") + } + + pa, err := portalloc.New(cfg.RTPPortRange[0], cfg.RTPPortRange[1]) + if err != nil { + return nil, fmt.Errorf("sipgateway: port allocator: %w", err) + } + + prefix := cfg.StreamPrefix + if prefix == "" { + prefix = "sip" + } + + maxCalls := cfg.MaxCalls + if maxCalls <= 0 { + maxCalls = 100 + } + + codecs := cfg.Codecs + if len(codecs) == 0 { + codecs = []string{"opus", "PCMA", "PCMU"} + } + + localIP := localAddress(sipSvc.LocalAddr()) + + gw := &Gateway{ + sipService: sipSvc, + hub: hub, + eventBus: bus, + portAlloc: pa, + prefix: prefix, + maxCalls: maxCalls, + codecs: codecs, + localIP: localIP, + sessions: make(map[string]*CallSession), + } + + sipSvc.OnInvite(gw.handleInvite) + sipSvc.OnBye(gw.handleBye) + + slog.Info("sip gateway enabled", "module", "sipgateway", + "prefix", prefix, "max_calls", maxCalls, "codecs", codecs) + + return gw, nil +} + +func (gw *Gateway) handleInvite(req *sip.Request, tx sip.ServerTransaction) { + callID := req.CallID().Value() + + gw.mu.Lock() + if _, exists := gw.sessions[callID]; exists { + gw.mu.Unlock() + resp := sip.NewResponseFromRequest(req, 486, "Busy Here", nil) + tx.Respond(resp) + return + } + if len(gw.sessions) >= gw.maxCalls { + gw.mu.Unlock() + resp := sip.NewResponseFromRequest(req, 503, "Service Unavailable", nil) + tx.Respond(resp) + return + } + gw.mu.Unlock() + + body := req.Body() + if len(body) == 0 { + resp := sip.NewResponseFromRequest(req, 400, "Bad Request", nil) + tx.Respond(resp) + return + } + + offerSDP, err := sdp.Parse(body) + if err != nil { + slog.Warn("invalid SDP in INVITE", "module", "sipgateway", "call", callID, "error", err) + resp := sip.NewResponseFromRequest(req, 400, "Bad Request", nil) + tx.Respond(resp) + return + } + + var audioMedia *sdp.MediaDescription + for _, m := range offerSDP.Media { + if m.Type == "audio" { + audioMedia = m + break + } + } + if audioMedia == nil { + resp := sip.NewResponseFromRequest(req, 488, "Not Acceptable Here", nil) + tx.Respond(resp) + return + } + + nc, ok := negotiateCodec(audioMedia, gw.codecs) + if !ok { + resp := sip.NewResponseFromRequest(req, 488, "Not Acceptable Here", nil) + tx.Respond(resp) + return + } + + rtpPort, rtcpPort, err := gw.portAlloc.AllocatePair() + if err != nil { + slog.Error("port allocation failed", "module", "sipgateway", "error", err) + resp := sip.NewResponseFromRequest(req, 503, "Service Unavailable", nil) + tx.Respond(resp) + return + } + + streamKey := gw.streamKeyFromRequest(req) + stream, _ := gw.hub.GetOrCreate(streamKey) + + cs := newCallSession(callID, streamKey, nc, "inbound", rtpPort, rtcpPort) + + remoteIP := remoteAddress(offerSDP) + if err := cs.startInbound(stream, remoteIP, audioMedia.Port); err != nil { + gw.portAlloc.Free(rtpPort, rtcpPort) + slog.Error("failed to start inbound session", "module", "sipgateway", + "call", callID, "error", err) + resp := sip.NewResponseFromRequest(req, 500, "Server Error", nil) + tx.Respond(resp) + return + } + + gw.mu.Lock() + gw.sessions[callID] = cs + gw.mu.Unlock() + + answerBody := buildAnswerSDP(gw.localIP, rtpPort, nc) + resp := sip.NewResponseFromRequest(req, 200, "OK", answerBody) + resp.AppendHeader(sip.NewHeader("Content-Type", "application/sdp")) + tx.Respond(resp) + + slog.Info("call established", "module", "sipgateway", + "call", callID, "stream", streamKey, "codec", nc.EncodingName, + "local_port", rtpPort, "remote", fmt.Sprintf("%s:%d", remoteIP, audioMedia.Port)) +} + +func (gw *Gateway) handleBye(req *sip.Request, tx sip.ServerTransaction) { + callID := req.CallID().Value() + + gw.mu.Lock() + cs, ok := gw.sessions[callID] + if ok { + delete(gw.sessions, callID) + } + gw.mu.Unlock() + + if !ok { + return + } + + cs.Close() + gw.portAlloc.Free(cs.rtpPort, cs.rtcpPort) + + if cs.stream != nil { + cs.stream.RemovePublisher() + } + + resp := sip.NewResponseFromRequest(req, 200, "OK", nil) + tx.Respond(resp) + + slog.Info("call ended", "module", "sipgateway", "call", callID, "stream", cs.streamKey) +} + +// Dial initiates an outbound call from a stream to a SIP URI. +func (gw *Gateway) Dial(ctx context.Context, targetURI, streamKey string) (string, error) { + stream, ok := gw.hub.Find(streamKey) + if !ok { + return "", fmt.Errorf("stream %q not found", streamKey) + } + + rtpPort, rtcpPort, err := gw.portAlloc.AllocatePair() + if err != nil { + return "", fmt.Errorf("port allocation: %w", err) + } + + var offerCodecs []negotiatedCodec + for _, name := range gw.codecs { + if info, ok := encodingToCodec[name]; ok { + offerCodecs = append(offerCodecs, negotiatedCodec{ + Codec: info.Codec, + PT: info.PT, + ClockRate: info.ClockRate, + EncodingName: name, + }) + } + } + + offerBody := buildOfferSDP(gw.localIP, rtpPort, offerCodecs) + + reqURI := sip.Uri{User: targetURI, Host: gw.sipService.Domain()} + fromURI := sip.Uri{User: gw.sipService.ServerID(), Host: gw.sipService.Domain()} + + inviteReq := sip.NewRequest(sip.INVITE, reqURI) + inviteReq.SetBody(offerBody) + inviteReq.AppendHeader(sip.NewHeader("Content-Type", "application/sdp")) + inviteReq.AppendHeader(&sip.FromHeader{Address: fromURI, Params: sip.NewParams()}) + inviteReq.AppendHeader(&sip.ToHeader{Address: reqURI}) + + invTx, err := gw.sipService.SendInvite(ctx, inviteReq) + if err != nil { + gw.portAlloc.Free(rtpPort, rtcpPort) + return "", fmt.Errorf("send INVITE: %w", err) + } + + // Wait for final response + select { + case <-ctx.Done(): + gw.portAlloc.Free(rtpPort, rtcpPort) + invTx.Close() + return "", ctx.Err() + case <-invTx.Done(): + } + + resp := invTx.Response() + if resp == nil || resp.StatusCode != 200 { + gw.portAlloc.Free(rtpPort, rtcpPort) + if resp != nil { + return "", fmt.Errorf("INVITE rejected: %d %s", resp.StatusCode, resp.Reason) + } + return "", fmt.Errorf("INVITE failed: no response") + } + + answerSDP, err := sdp.Parse(resp.Body()) + if err != nil { + gw.portAlloc.Free(rtpPort, rtcpPort) + return "", fmt.Errorf("parse answer SDP: %w", err) + } + + var audioMedia *sdp.MediaDescription + for _, m := range answerSDP.Media { + if m.Type == "audio" { + audioMedia = m + break + } + } + if audioMedia == nil { + gw.portAlloc.Free(rtpPort, rtcpPort) + return "", fmt.Errorf("no audio in answer SDP") + } + + nc, ok := negotiateCodec(audioMedia, gw.codecs) + if !ok { + gw.portAlloc.Free(rtpPort, rtcpPort) + return "", fmt.Errorf("no common codec in answer") + } + + if err := invTx.SendACK(ctx); err != nil { + gw.portAlloc.Free(rtpPort, rtcpPort) + return "", fmt.Errorf("send ACK: %w", err) + } + + callID := inviteReq.CallID().Value() + cs := newCallSession(callID, streamKey, nc, "outbound", rtpPort, rtcpPort) + + remoteIP := remoteAddress(answerSDP) + if err := cs.startOutbound(stream, remoteIP, audioMedia.Port); err != nil { + gw.portAlloc.Free(rtpPort, rtcpPort) + return "", fmt.Errorf("start outbound: %w", err) + } + + gw.mu.Lock() + gw.sessions[callID] = cs + gw.mu.Unlock() + + slog.Info("outbound call established", "module", "sipgateway", + "call", callID, "target", targetURI, "stream", streamKey, "codec", nc.EncodingName) + + return callID, nil +} + +// Hangup terminates a call by its call-ID. +func (gw *Gateway) Hangup(callID string) error { + gw.mu.Lock() + cs, ok := gw.sessions[callID] + if ok { + delete(gw.sessions, callID) + } + gw.mu.Unlock() + + if !ok { + return fmt.Errorf("call %q not found", callID) + } + + cs.Close() + gw.portAlloc.Free(cs.rtpPort, cs.rtcpPort) + + if cs.direction == "inbound" && cs.stream != nil { + cs.stream.RemovePublisher() + } + + return nil +} + +// ActiveCalls returns the number of active calls. +func (gw *Gateway) ActiveCalls() int { + gw.mu.Lock() + defer gw.mu.Unlock() + return len(gw.sessions) +} + +// Close stops all active calls and the gateway. +func (gw *Gateway) Close() { + gw.mu.Lock() + sessions := make(map[string]*CallSession, len(gw.sessions)) + for k, v := range gw.sessions { + sessions[k] = v + } + gw.sessions = make(map[string]*CallSession) + gw.mu.Unlock() + + for _, cs := range sessions { + cs.Close() + gw.portAlloc.Free(cs.rtpPort, cs.rtcpPort) + } +} + +func (gw *Gateway) streamKeyFromRequest(req *sip.Request) string { + user := req.Recipient.User + if user == "" { + user = req.CallID().Value() + } + return gw.prefix + "/" + user +} + +func remoteAddress(sd *sdp.SessionDescription) string { + if sd.Connection != nil { + return sd.Connection.Address + } + for _, m := range sd.Media { + if m.Connection != nil { + return m.Connection.Address + } + } + return "" +} + +func localAddress(listenAddr string) string { + host, _, err := net.SplitHostPort(listenAddr) + if err != nil { + return "0.0.0.0" + } + if host == "" || host == "0.0.0.0" || host == "::" { + if ip := preferredOutboundIP(); ip != "" { + return ip + } + return "0.0.0.0" + } + return host +} + +func preferredOutboundIP() string { + conn, err := net.Dial("udp4", "8.8.8.8:80") + if err != nil { + return "" + } + defer conn.Close() + addr := conn.LocalAddr().(*net.UDPAddr) + return addr.IP.String() +} diff --git a/module/sipgateway/gateway_test.go b/module/sipgateway/gateway_test.go new file mode 100644 index 0000000..23d3a7d --- /dev/null +++ b/module/sipgateway/gateway_test.go @@ -0,0 +1,552 @@ +package sipgateway + +import ( + "context" + "sync" + "testing" + "time" + + "github.com/emiago/sipgo/sip" + "github.com/im-pingo/liveforge/config" + "github.com/im-pingo/liveforge/core" + sipmod "github.com/im-pingo/liveforge/module/sip" + "github.com/im-pingo/liveforge/pkg/sdp" +) + +type mockSIPService struct { + mu sync.Mutex + inviteHandlers []sipmod.InviteHandler + byeHandlers []sipmod.ByeHandler + localAddr string + serverID string + domain string +} + +func (m *mockSIPService) OnRegister(h sipmod.RegisterHandler) {} +func (m *mockSIPService) OnInvite(h sipmod.InviteHandler) { m.mu.Lock(); m.inviteHandlers = append(m.inviteHandlers, h); m.mu.Unlock() } +func (m *mockSIPService) OnBye(h sipmod.ByeHandler) { m.mu.Lock(); m.byeHandlers = append(m.byeHandlers, h); m.mu.Unlock() } +func (m *mockSIPService) OnMessage(h sipmod.MessageHandler) {} +func (m *mockSIPService) OnSubscribe(h sipmod.SubscribeHandler) {} +func (m *mockSIPService) OnNotify(h sipmod.NotifyHandler) {} +func (m *mockSIPService) SendRequest(ctx context.Context, req *sip.Request) (*sip.Response, error) { + return nil, nil +} +func (m *mockSIPService) SendInvite(ctx context.Context, req *sip.Request) (*sipmod.InviteTransaction, error) { + return nil, nil +} +func (m *mockSIPService) LocalAddr() string { return m.localAddr } +func (m *mockSIPService) ServerID() string { return m.serverID } +func (m *mockSIPService) Domain() string { return m.domain } + +type mockServerTx struct { + mu sync.Mutex + response *sip.Response +} + +func (tx *mockServerTx) Respond(resp *sip.Response) error { + tx.mu.Lock() + tx.response = resp + tx.mu.Unlock() + return nil +} + +func (tx *mockServerTx) Acks() <-chan *sip.Request { return nil } +func (tx *mockServerTx) Done() <-chan struct{} { return nil } +func (tx *mockServerTx) Terminate() {} +func (tx *mockServerTx) Err() error { return nil } +func (tx *mockServerTx) OnTerminate(f sip.FnTxTerminate) bool { return true } +func (tx *mockServerTx) OnCancel(f sip.FnTxCancel) bool { return true } + +func (tx *mockServerTx) getResponse() *sip.Response { + tx.mu.Lock() + defer tx.mu.Unlock() + return tx.response +} + +func newTestHub() *core.StreamHub { + bus := core.NewEventBus() + return core.NewStreamHub(config.StreamConfig{RingBufferSize: 256}, config.LimitsConfig{}, bus) +} + +func newTestGatewayConfig() config.SIPGatewayConfig { + return config.SIPGatewayConfig{ + Enabled: true, + StreamPrefix: "sip", + RTPPortRange: []int{40000, 40100}, + Codecs: []string{"PCMA", "PCMU"}, + MaxCalls: 10, + } +} + +func TestNewGateway(t *testing.T) { + sipSvc := &mockSIPService{localAddr: "127.0.0.1:5060", serverID: "test", domain: "test.local"} + hub := newTestHub() + bus := core.NewEventBus() + + gw, err := NewGateway(newTestGatewayConfig(), sipSvc, hub, bus) + if err != nil { + t.Fatalf("NewGateway: %v", err) + } + defer gw.Close() + + if gw.ActiveCalls() != 0 { + t.Errorf("ActiveCalls = %d, want 0", gw.ActiveCalls()) + } +} + +func TestNewGatewayBadPortRange(t *testing.T) { + sipSvc := &mockSIPService{localAddr: "127.0.0.1:5060"} + hub := newTestHub() + bus := core.NewEventBus() + + cfg := newTestGatewayConfig() + cfg.RTPPortRange = []int{100} + + _, err := NewGateway(cfg, sipSvc, hub, bus) + if err == nil { + t.Error("expected error for bad port range") + } +} + +func TestNegotiateCodec(t *testing.T) { + tests := []struct { + name string + offer *sdp.MediaDescription + preferred []string + wantCodec string + wantOK bool + }{ + { + name: "PCMA offered and preferred", + offer: &sdp.MediaDescription{ + Type: "audio", + Formats: []int{8, 0}, + Attributes: []sdp.Attribute{ + {Key: "rtpmap", Value: "8 PCMA/8000"}, + {Key: "rtpmap", Value: "0 PCMU/8000"}, + }, + }, + preferred: []string{"PCMA", "PCMU"}, + wantCodec: "PCMA", + wantOK: true, + }, + { + name: "PCMU preferred over PCMA", + offer: &sdp.MediaDescription{ + Type: "audio", + Formats: []int{8, 0}, + Attributes: []sdp.Attribute{ + {Key: "rtpmap", Value: "8 PCMA/8000"}, + {Key: "rtpmap", Value: "0 PCMU/8000"}, + }, + }, + preferred: []string{"PCMU", "PCMA"}, + wantCodec: "PCMU", + wantOK: true, + }, + { + name: "static PT without rtpmap", + offer: &sdp.MediaDescription{ + Type: "audio", + Formats: []int{0}, + }, + preferred: []string{"PCMU"}, + wantCodec: "PCMU", + wantOK: true, + }, + { + name: "opus dynamic PT", + offer: &sdp.MediaDescription{ + Type: "audio", + Formats: []int{111}, + Attributes: []sdp.Attribute{ + {Key: "rtpmap", Value: "111 opus/48000/2"}, + }, + }, + preferred: []string{"opus"}, + wantCodec: "opus", + wantOK: true, + }, + { + name: "no common codec", + offer: &sdp.MediaDescription{ + Type: "audio", + Formats: []int{96}, + Attributes: []sdp.Attribute{ + {Key: "rtpmap", Value: "96 CUSTOM/16000"}, + }, + }, + preferred: []string{"PCMA", "PCMU"}, + wantOK: false, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + nc, ok := negotiateCodec(tt.offer, tt.preferred) + if ok != tt.wantOK { + t.Fatalf("ok = %v, want %v", ok, tt.wantOK) + } + if ok && nc.EncodingName != tt.wantCodec { + t.Errorf("codec = %q, want %q", nc.EncodingName, tt.wantCodec) + } + }) + } +} + +func TestBuildAnswerSDP(t *testing.T) { + nc := negotiatedCodec{ + Codec: 8, // CodecG711A + PT: 8, + ClockRate: 8000, + EncodingName: "PCMA", + } + + body := buildAnswerSDP("192.168.1.1", 40000, nc) + + sd, err := sdp.Parse(body) + if err != nil { + t.Fatalf("Parse answer SDP: %v", err) + } + + if len(sd.Media) != 1 { + t.Fatalf("expected 1 media section, got %d", len(sd.Media)) + } + + m := sd.Media[0] + if m.Type != "audio" { + t.Errorf("media type = %q, want audio", m.Type) + } + if m.Port != 40000 { + t.Errorf("port = %d, want 40000", m.Port) + } + if len(m.Formats) != 1 || m.Formats[0] != 8 { + t.Errorf("formats = %v, want [8]", m.Formats) + } +} + +func TestBuildOfferSDP(t *testing.T) { + codecs := []negotiatedCodec{ + {Codec: 8, PT: 8, ClockRate: 8000, EncodingName: "PCMA"}, + {Codec: 0, PT: 0, ClockRate: 8000, EncodingName: "PCMU"}, + } + + body := buildOfferSDP("10.0.0.1", 40002, codecs) + + sd, err := sdp.Parse(body) + if err != nil { + t.Fatalf("Parse offer SDP: %v", err) + } + + if len(sd.Media) != 1 { + t.Fatalf("expected 1 media, got %d", len(sd.Media)) + } + + m := sd.Media[0] + if len(m.Formats) != 2 { + t.Errorf("formats = %v, want 2 entries", m.Formats) + } +} + +func TestGatewayHandleInviteSuccess(t *testing.T) { + sipSvc := &mockSIPService{localAddr: "127.0.0.1:5060", serverID: "test", domain: "test.local"} + hub := newTestHub() + bus := core.NewEventBus() + + gw, err := NewGateway(newTestGatewayConfig(), sipSvc, hub, bus) + if err != nil { + t.Fatalf("NewGateway: %v", err) + } + defer gw.Close() + + offerSDP := "v=0\r\no=- 1 1 IN IP4 192.168.1.100\r\ns=-\r\nc=IN IP4 192.168.1.100\r\nt=0 0\r\nm=audio 40000 RTP/AVP 8 0\r\na=rtpmap:8 PCMA/8000\r\na=rtpmap:0 PCMU/8000\r\n" + + req := sip.NewRequest(sip.INVITE, sip.Uri{User: "teststream", Host: "test.local"}) + req.AppendHeader(sip.NewHeader("Call-ID", "test-call-123")) + req.SetBody([]byte(offerSDP)) + + tx := &mockServerTx{} + + sipSvc.mu.Lock() + handlers := make([]sipmod.InviteHandler, len(sipSvc.inviteHandlers)) + copy(handlers, sipSvc.inviteHandlers) + sipSvc.mu.Unlock() + + for _, h := range handlers { + h(req, tx) + } + + resp := tx.getResponse() + if resp == nil { + t.Fatal("no response sent") + } + if resp.StatusCode != 200 { + t.Errorf("status = %d, want 200", resp.StatusCode) + } + + if gw.ActiveCalls() != 1 { + t.Errorf("ActiveCalls = %d, want 1", gw.ActiveCalls()) + } + + _, ok := hub.Find("sip/teststream") + if !ok { + t.Error("stream sip/teststream not found in hub") + } +} + +func TestGatewayHandleInviteNoAudio(t *testing.T) { + sipSvc := &mockSIPService{localAddr: "127.0.0.1:5060", serverID: "test", domain: "test.local"} + hub := newTestHub() + bus := core.NewEventBus() + + gw, err := NewGateway(newTestGatewayConfig(), sipSvc, hub, bus) + if err != nil { + t.Fatalf("NewGateway: %v", err) + } + defer gw.Close() + + offerSDP := "v=0\r\no=- 1 1 IN IP4 192.168.1.100\r\ns=-\r\nc=IN IP4 192.168.1.100\r\nt=0 0\r\nm=video 40000 RTP/AVP 96\r\na=rtpmap:96 H264/90000\r\n" + + req := sip.NewRequest(sip.INVITE, sip.Uri{User: "teststream", Host: "test.local"}) + req.AppendHeader(sip.NewHeader("Call-ID", "test-call-video")) + req.SetBody([]byte(offerSDP)) + + tx := &mockServerTx{} + + sipSvc.mu.Lock() + handlers := make([]sipmod.InviteHandler, len(sipSvc.inviteHandlers)) + copy(handlers, sipSvc.inviteHandlers) + sipSvc.mu.Unlock() + + for _, h := range handlers { + h(req, tx) + } + + resp := tx.getResponse() + if resp == nil || resp.StatusCode != 488 { + status := 0 + if resp != nil { + status = resp.StatusCode + } + t.Errorf("status = %d, want 488", status) + } +} + +func TestGatewayHandleBye(t *testing.T) { + sipSvc := &mockSIPService{localAddr: "127.0.0.1:5060", serverID: "test", domain: "test.local"} + hub := newTestHub() + bus := core.NewEventBus() + + gw, err := NewGateway(newTestGatewayConfig(), sipSvc, hub, bus) + if err != nil { + t.Fatalf("NewGateway: %v", err) + } + defer gw.Close() + + // First establish a call via INVITE + offerSDP := "v=0\r\no=- 1 1 IN IP4 192.168.1.100\r\ns=-\r\nc=IN IP4 192.168.1.100\r\nt=0 0\r\nm=audio 40000 RTP/AVP 8\r\na=rtpmap:8 PCMA/8000\r\n" + + req := sip.NewRequest(sip.INVITE, sip.Uri{User: "byetest", Host: "test.local"}) + req.AppendHeader(sip.NewHeader("Call-ID", "bye-call-456")) + req.SetBody([]byte(offerSDP)) + + inviteTx := &mockServerTx{} + + sipSvc.mu.Lock() + inviteHandlers := make([]sipmod.InviteHandler, len(sipSvc.inviteHandlers)) + copy(inviteHandlers, sipSvc.inviteHandlers) + sipSvc.mu.Unlock() + for _, h := range inviteHandlers { + h(req, inviteTx) + } + + if gw.ActiveCalls() != 1 { + t.Fatalf("ActiveCalls after INVITE = %d, want 1", gw.ActiveCalls()) + } + + // Now send BYE + byeReq := sip.NewRequest(sip.BYE, sip.Uri{User: "byetest", Host: "test.local"}) + byeReq.AppendHeader(sip.NewHeader("Call-ID", "bye-call-456")) + + byeTx := &mockServerTx{} + + sipSvc.mu.Lock() + byeHandlers := make([]sipmod.ByeHandler, len(sipSvc.byeHandlers)) + copy(byeHandlers, sipSvc.byeHandlers) + sipSvc.mu.Unlock() + for _, h := range byeHandlers { + h(byeReq, byeTx) + } + + time.Sleep(50 * time.Millisecond) + + if gw.ActiveCalls() != 0 { + t.Errorf("ActiveCalls after BYE = %d, want 0", gw.ActiveCalls()) + } + + byeResp := byeTx.getResponse() + if byeResp == nil || byeResp.StatusCode != 200 { + t.Error("expected 200 OK for BYE") + } +} + +func TestGatewayMaxCalls(t *testing.T) { + sipSvc := &mockSIPService{localAddr: "127.0.0.1:5060", serverID: "test", domain: "test.local"} + hub := newTestHub() + bus := core.NewEventBus() + + cfg := newTestGatewayConfig() + cfg.MaxCalls = 1 + + gw, err := NewGateway(cfg, sipSvc, hub, bus) + if err != nil { + t.Fatalf("NewGateway: %v", err) + } + defer gw.Close() + + offerSDP := "v=0\r\no=- 1 1 IN IP4 192.168.1.100\r\ns=-\r\nc=IN IP4 192.168.1.100\r\nt=0 0\r\nm=audio 40000 RTP/AVP 8\r\na=rtpmap:8 PCMA/8000\r\n" + + sipSvc.mu.Lock() + inviteHandlers := make([]sipmod.InviteHandler, len(sipSvc.inviteHandlers)) + copy(inviteHandlers, sipSvc.inviteHandlers) + sipSvc.mu.Unlock() + + // First call + req1 := sip.NewRequest(sip.INVITE, sip.Uri{User: "call1", Host: "test.local"}) + req1.AppendHeader(sip.NewHeader("Call-ID", "max-call-1")) + req1.SetBody([]byte(offerSDP)) + tx1 := &mockServerTx{} + for _, h := range inviteHandlers { + h(req1, tx1) + } + + // Second call should be rejected + req2 := sip.NewRequest(sip.INVITE, sip.Uri{User: "call2", Host: "test.local"}) + req2.AppendHeader(sip.NewHeader("Call-ID", "max-call-2")) + req2.SetBody([]byte(offerSDP)) + tx2 := &mockServerTx{} + for _, h := range inviteHandlers { + h(req2, tx2) + } + + resp2 := tx2.getResponse() + if resp2 == nil || resp2.StatusCode != 503 { + status := 0 + if resp2 != nil { + status = resp2.StatusCode + } + t.Errorf("second call status = %d, want 503", status) + } +} + +func TestGatewayClose(t *testing.T) { + sipSvc := &mockSIPService{localAddr: "127.0.0.1:5060", serverID: "test", domain: "test.local"} + hub := newTestHub() + bus := core.NewEventBus() + + gw, err := NewGateway(newTestGatewayConfig(), sipSvc, hub, bus) + if err != nil { + t.Fatalf("NewGateway: %v", err) + } + + // Establish a call + offerSDP := "v=0\r\no=- 1 1 IN IP4 192.168.1.100\r\ns=-\r\nc=IN IP4 192.168.1.100\r\nt=0 0\r\nm=audio 40000 RTP/AVP 0\r\na=rtpmap:0 PCMU/8000\r\n" + + req := sip.NewRequest(sip.INVITE, sip.Uri{User: "closetest", Host: "test.local"}) + req.AppendHeader(sip.NewHeader("Call-ID", "close-call-789")) + req.SetBody([]byte(offerSDP)) + + tx := &mockServerTx{} + sipSvc.mu.Lock() + handlers := make([]sipmod.InviteHandler, len(sipSvc.inviteHandlers)) + copy(handlers, sipSvc.inviteHandlers) + sipSvc.mu.Unlock() + for _, h := range handlers { + h(req, tx) + } + + if gw.ActiveCalls() != 1 { + t.Fatalf("ActiveCalls = %d, want 1", gw.ActiveCalls()) + } + + gw.Close() + + if gw.ActiveCalls() != 0 { + t.Errorf("ActiveCalls after close = %d, want 0", gw.ActiveCalls()) + } +} + +func TestModuleName(t *testing.T) { + m := NewModule(&mockSIPService{}) + if m.Name() != "sipgateway" { + t.Errorf("Name = %q, want sipgateway", m.Name()) + } +} + +func TestModuleCloseNilGateway(t *testing.T) { + m := NewModule(&mockSIPService{}) + if err := m.Close(); err != nil { + t.Fatalf("Close: %v", err) + } +} + +func TestHangupUnknownCall(t *testing.T) { + sipSvc := &mockSIPService{localAddr: "127.0.0.1:5060", serverID: "test", domain: "test.local"} + hub := newTestHub() + bus := core.NewEventBus() + + gw, err := NewGateway(newTestGatewayConfig(), sipSvc, hub, bus) + if err != nil { + t.Fatalf("NewGateway: %v", err) + } + defer gw.Close() + + if err := gw.Hangup("nonexistent"); err == nil { + t.Error("expected error for unknown call") + } +} + +func TestRemoteAddress(t *testing.T) { + tests := []struct { + name string + sd *sdp.SessionDescription + want string + }{ + { + name: "session-level connection", + sd: &sdp.SessionDescription{ + Connection: &sdp.Connection{Address: "10.0.0.1"}, + }, + want: "10.0.0.1", + }, + { + name: "media-level connection", + sd: &sdp.SessionDescription{ + Media: []*sdp.MediaDescription{ + {Connection: &sdp.Connection{Address: "10.0.0.2"}}, + }, + }, + want: "10.0.0.2", + }, + { + name: "no connection", + sd: &sdp.SessionDescription{}, + want: "", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got := remoteAddress(tt.sd) + if got != tt.want { + t.Errorf("remoteAddress = %q, want %q", got, tt.want) + } + }) + } +} + +func TestLocalAddress(t *testing.T) { + if got := localAddress("192.168.1.1:5060"); got != "192.168.1.1" { + t.Errorf("localAddress(192.168.1.1:5060) = %q", got) + } +} diff --git a/module/sipgateway/module.go b/module/sipgateway/module.go new file mode 100644 index 0000000..e81eb05 --- /dev/null +++ b/module/sipgateway/module.go @@ -0,0 +1,51 @@ +package sipgateway + +import ( + "github.com/im-pingo/liveforge/core" + sipmod "github.com/im-pingo/liveforge/module/sip" +) + +// Module implements core.Module for the SIP-to-stream gateway. +type Module struct { + gw *Gateway + sipService sipmod.SIPService +} + +// NewModule creates a new SIP gateway module. +func NewModule(sipService sipmod.SIPService) *Module { + return &Module{sipService: sipService} +} + +// Name returns the module name. +func (m *Module) Name() string { return "sipgateway" } + +// Init initializes the gateway if enabled in config. +func (m *Module) Init(s *core.Server) error { + cfg := s.Config().SIP.Gateway + if !cfg.Enabled { + return nil + } + + gw, err := NewGateway(cfg, m.sipService, s.StreamHub(), s.GetEventBus()) + if err != nil { + return err + } + m.gw = gw + return nil +} + +// Hooks returns empty hooks — gateway uses SIP event dispatch. +func (m *Module) Hooks() []core.HookRegistration { return nil } + +// Close stops the gateway. +func (m *Module) Close() error { + if m.gw != nil { + m.gw.Close() + } + return nil +} + +// Gateway returns the gateway instance, or nil if disabled. +func (m *Module) Gateway() *Gateway { + return m.gw +} diff --git a/module/webrtc/webrtc_test.go b/module/webrtc/webrtc_test.go index c59de9c..1f20dbd 100644 --- a/module/webrtc/webrtc_test.go +++ b/module/webrtc/webrtc_test.go @@ -13,6 +13,7 @@ import ( "github.com/im-pingo/liveforge/core" "github.com/im-pingo/liveforge/pkg/avframe" "github.com/pion/rtcp" + "github.com/pion/sdp/v3" "github.com/pion/webrtc/v4" "github.com/pion/webrtc/v4/pkg/media" ) @@ -565,3 +566,117 @@ func createMinimalOffer(t *testing.T) string { } return offer.SDP } + +func TestOfferSupportsCodec(t *testing.T) { + offer := createMinimalOffer(t) + var parsed sdp.SessionDescription + if err := parsed.UnmarshalString(offer); err != nil { + t.Fatal(err) + } + + tests := []struct { + media string + mime string + want bool + }{ + {"video", webrtc.MimeTypeH264, true}, + {"video", webrtc.MimeTypeVP8, true}, + {"video", "video/NONEXISTENT", false}, + {"audio", "audio/NONEXISTENT", false}, + } + + for _, tt := range tests { + t.Run(tt.mime, func(t *testing.T) { + got := offerSupportsCodec(&parsed, tt.media, tt.mime) + if got != tt.want { + t.Errorf("offerSupportsCodec(%s, %s) = %v, want %v", tt.media, tt.mime, got, tt.want) + } + }) + } +} + +func TestICEServersFromConfigWithTURN(t *testing.T) { + cfg := &config.Config{ + Stream: config.StreamConfig{ + RingBufferSize: 256, + GOPCache: true, + GOPCacheNum: 1, + IdleTimeout: 30 * time.Second, + NoPublisherTimeout: 15 * time.Second, + }, + WebRTC: config.WebRTCConfig{ + Enabled: true, + Listen: ":0", + ICELite: false, + TLS: func() *bool { b := false; return &b }(), + UDPPortRange: []int{20000, 20100}, + ICEServers: []config.ICEServer{ + {URLs: []string{"stun:stun.l.google.com:19302"}}, + {URLs: []string{"turn:turn.example.com:3478"}, Username: "user1", Credential: "pass1"}, + {URLs: []string{"turns:turn.example.com:5349"}, Username: "user2", Credential: "pass2"}, + }, + }, + } + s := core.NewServer(cfg) + m := NewModule() + if err := m.Init(s); err != nil { + t.Fatalf("Init failed: %v", err) + } + defer m.Close() + + servers := m.iceServersFromConfig() + if len(servers) != 3 { + t.Fatalf("expected 3 ICE servers, got %d", len(servers)) + } + + // STUN server (no credentials) + if servers[0].URLs[0] != "stun:stun.l.google.com:19302" { + t.Errorf("unexpected STUN URL: %s", servers[0].URLs[0]) + } + + // TURN server with credentials + if servers[1].URLs[0] != "turn:turn.example.com:3478" { + t.Errorf("unexpected TURN URL: %s", servers[1].URLs[0]) + } + if servers[1].Username != "user1" || servers[1].Credential != "pass1" { + t.Errorf("TURN credentials not passed: user=%s cred=%v", servers[1].Username, servers[1].Credential) + } + + // TURNS server + if servers[2].URLs[0] != "turns:turn.example.com:5349" { + t.Errorf("unexpected TURNS URL: %s", servers[2].URLs[0]) + } +} + +func TestICEServersSkippedWithICELite(t *testing.T) { + cfg := &config.Config{ + Stream: config.StreamConfig{ + RingBufferSize: 256, + GOPCache: true, + GOPCacheNum: 1, + IdleTimeout: 30 * time.Second, + NoPublisherTimeout: 15 * time.Second, + }, + WebRTC: config.WebRTCConfig{ + Enabled: true, + Listen: ":0", + ICELite: true, + TLS: func() *bool { b := false; return &b }(), + UDPPortRange: []int{20000, 20100}, + ICEServers: []config.ICEServer{ + {URLs: []string{"turn:turn.example.com:3478"}, Username: "u", Credential: "p"}, + }, + }, + } + s := core.NewServer(cfg) + m := NewModule() + if err := m.Init(s); err != nil { + t.Fatalf("Init failed: %v", err) + } + defer m.Close() + + servers := m.iceServersFromConfig() + if servers != nil { + t.Errorf("ICE Lite should return nil ICE servers, got %d", len(servers)) + } +} diff --git a/module/webrtc/whep.go b/module/webrtc/whep.go index 949cb95..bbfa069 100644 --- a/module/webrtc/whep.go +++ b/module/webrtc/whep.go @@ -145,6 +145,13 @@ func (m *Module) handleWHEP(w http.ResponseWriter, r *http.Request) { if info.HasVideo() && offerHasVideo { mime := codecToMime(info.VideoCodec) + if mime != "" && !offerSupportsCodec(&parsedSDP, "video", mime) { + // Publisher codec not in offer; H.265 publishers with H.264-only peers + // would need video transcoding (deferred). Log and skip video. + slog.Debug("WHEP video codec mismatch", "module", "webrtc", + "publisher", mime, "stream", streamKey) + mime = "" + } if mime != "" { vt, err := webrtc.NewTrackLocalStaticSample( webrtc.RTPCodecCapability{MimeType: mime, ClockRate: 90000}, @@ -363,3 +370,22 @@ func normalizeH264Offer(offerSDP string) string { return string(result) } +// offerSupportsCodec checks whether the remote offer contains the given codec +// MIME type (e.g., "video/H265") in its rtpmap attributes. +func offerSupportsCodec(parsed *sdp.SessionDescription, media, mime string) bool { + codecName := strings.ToUpper(mime) + if idx := strings.LastIndex(codecName, "/"); idx >= 0 { + codecName = codecName[idx+1:] + } + for _, md := range parsed.MediaDescriptions { + if md.MediaName.Media != media { + continue + } + for _, attr := range md.Attributes { + if attr.Key == "rtpmap" && strings.Contains(strings.ToUpper(attr.Value), codecName+"/") { + return true + } + } + } + return false +} \ No newline at end of file diff --git a/pkg/muxer/mp4/box.go b/pkg/muxer/mp4/box.go new file mode 100644 index 0000000..96d46b0 --- /dev/null +++ b/pkg/muxer/mp4/box.go @@ -0,0 +1,42 @@ +package mp4 + +import ( + "encoding/binary" + "io" +) + +func writeBox(w io.Writer, boxType [4]byte, payload []byte) { + size := uint32(8 + len(payload)) + var hdr [8]byte + binary.BigEndian.PutUint32(hdr[0:4], size) + copy(hdr[4:8], boxType[:]) + w.Write(hdr[:]) + w.Write(payload) +} + +func writeFullBox(w io.Writer, boxType [4]byte, version byte, flags uint32, payload []byte) { + size := uint32(12 + len(payload)) + var hdr [12]byte + binary.BigEndian.PutUint32(hdr[0:4], size) + copy(hdr[4:8], boxType[:]) + hdr[8] = version + binary.BigEndian.PutUint32(hdr[8:12], (uint32(version)<<24)|(flags&0x00FFFFFF)) + w.Write(hdr[:]) + w.Write(payload) +} + +func putU16(buf []byte, v uint16) { + binary.BigEndian.PutUint16(buf, v) +} + +func putU32(buf []byte, v uint32) { + binary.BigEndian.PutUint32(buf, v) +} + +func putU64(buf []byte, v uint64) { + binary.BigEndian.PutUint64(buf, v) +} + +func boxSize(boxType [4]byte, payload []byte) int { + return 8 + len(payload) +} diff --git a/pkg/muxer/mp4/muxer.go b/pkg/muxer/mp4/muxer.go new file mode 100644 index 0000000..b28017c --- /dev/null +++ b/pkg/muxer/mp4/muxer.go @@ -0,0 +1,634 @@ +package mp4 + +import ( + "bytes" + "encoding/binary" + "io" + + "github.com/im-pingo/liveforge/pkg/avframe" + "github.com/im-pingo/liveforge/pkg/codec/aac" + "github.com/im-pingo/liveforge/pkg/codec/h264" +) + +// Muxer produces a classic MP4 file (ftyp + mdat + moov). +// Frames are written sequentially; moov is built and written on Finalize(). +type Muxer struct { + videoCodec avframe.CodecType + audioCodec avframe.CodecType + timescale uint32 + + videoSeqHeader []byte // raw AVC/HEVC decoder config + audioSeqHeader []byte // raw AudioSpecificConfig + sps, pps []byte // extracted from AVC record + + videoSamples []sampleEntry + audioSamples []sampleEntry + + mdatOffset int64 // byte offset where mdat data starts (after 8-byte header) + mdatSize int64 // total bytes in mdat (excluding box header) + + audioSampleRate uint32 + audioChannels uint16 +} + +type sampleEntry struct { + size uint32 + duration uint32 + offset int64 + isSync bool + cts int32 // composition time offset (PTS - DTS) +} + +// NewMuxer creates an MP4 muxer. +func NewMuxer(videoCodec, audioCodec avframe.CodecType) *Muxer { + return &Muxer{ + videoCodec: videoCodec, + audioCodec: audioCodec, + timescale: 90000, + audioSampleRate: 44100, + audioChannels: 2, + } +} + +// SetAudioParams sets audio sample rate and channels from stream metadata. +func (m *Muxer) SetAudioParams(sampleRate uint32, channels uint16) { + if sampleRate > 0 { + m.audioSampleRate = sampleRate + } + if channels > 0 { + m.audioChannels = channels + } +} + +// WriteFtyp writes the ftyp box. +func (m *Muxer) WriteFtyp(w io.Writer) error { + var buf bytes.Buffer + buf.Write([]byte("isom")) // major brand + putU32Buf(&buf, 0x00000200) // minor version + buf.Write([]byte("isomiso2")) // compatible brands + if m.videoCodec == avframe.CodecH264 { + buf.Write([]byte("avc1")) + } + buf.Write([]byte("mp41")) + + writeBox(w, [4]byte{'f', 't', 'y', 'p'}, buf.Bytes()) + return nil +} + +// WriteMdatHeader writes the mdat box header (size placeholder, filled by Finalize). +// Returns the offset of the mdat size field for later fixup. +func (m *Muxer) WriteMdatHeader(w io.WriteSeeker) (int64, error) { + offset, _ := w.Seek(0, io.SeekCurrent) + m.mdatOffset = offset + 8 // data starts after the 8-byte box header + + var hdr [8]byte + putU32(hdr[0:4], 0) // placeholder — will be fixed up + copy(hdr[4:8], []byte("mdat")) + _, err := w.Write(hdr[:]) + return offset, err +} + +// WriteFrame appends a frame to the mdat region and records sample metadata. +// Returns the number of bytes written. +func (m *Muxer) WriteFrame(w io.WriteSeeker, frame *avframe.AVFrame, prevDTS int64) (int, error) { + if frame.FrameType == avframe.FrameTypeSequenceHeader { + if frame.MediaType.IsVideo() { + m.videoCodec = frame.Codec + m.videoSeqHeader = append([]byte(nil), frame.Payload...) + if m.videoCodec == avframe.CodecH264 && len(m.videoSeqHeader) > 5 { + sps, pps, err := h264.ExtractSPSPPSFromAVCRecord(m.videoSeqHeader) + if err == nil { + m.sps = sps + m.pps = pps + } + } + } else if frame.MediaType.IsAudio() { + m.audioCodec = frame.Codec + m.audioSeqHeader = append([]byte(nil), frame.Payload...) + if m.audioCodec == avframe.CodecAAC && len(m.audioSeqHeader) >= 2 { + info, err := aac.ParseAudioSpecificConfig(m.audioSeqHeader) + if err == nil { + m.audioSampleRate = uint32(info.SampleRate) + m.audioChannels = uint16(info.Channels) + } + } + } + return 0, nil + } + + offset, _ := w.Seek(0, io.SeekCurrent) + data := frame.Payload + n, err := w.Write(data) + if err != nil { + return 0, err + } + + m.mdatSize += int64(n) + + duration := uint32(0) + if prevDTS >= 0 { + d := frame.DTS - prevDTS + if d > 0 { + duration = uint32(d * int64(m.timescale) / 1000) + } + } + + entry := sampleEntry{ + size: uint32(n), + offset: offset, + isSync: frame.FrameType.IsKeyframe() || frame.FrameType == avframe.FrameTypeSequenceHeader, + cts: int32((frame.PTS - frame.DTS) * int64(m.timescale) / 1000), + } + + if frame.MediaType.IsVideo() { + if len(m.videoSamples) > 0 { + m.videoSamples[len(m.videoSamples)-1].duration = duration + } + entry.isSync = frame.FrameType.IsKeyframe() + m.videoSamples = append(m.videoSamples, entry) + } else if frame.MediaType.IsAudio() { + if len(m.audioSamples) > 0 { + m.audioSamples[len(m.audioSamples)-1].duration = duration + } + m.audioSamples = append(m.audioSamples, entry) + } + + return n, nil +} + +// Finalize writes the moov box and fixes up the mdat size. +func (m *Muxer) Finalize(w io.WriteSeeker) error { + // Fix up mdat box size + end, _ := w.Seek(0, io.SeekCurrent) + mdatTotalSize := end - (m.mdatOffset - 8) + w.Seek(m.mdatOffset-8, io.SeekStart) + var sizeBytes [4]byte + binary.BigEndian.PutUint32(sizeBytes[:], uint32(mdatTotalSize)) + w.Write(sizeBytes[:]) + w.Seek(end, io.SeekStart) + + // Set last sample durations if missing + if len(m.videoSamples) > 0 && m.videoSamples[len(m.videoSamples)-1].duration == 0 { + if len(m.videoSamples) > 1 { + m.videoSamples[len(m.videoSamples)-1].duration = m.videoSamples[len(m.videoSamples)-2].duration + } else { + m.videoSamples[0].duration = 3000 // ~33ms at 90kHz + } + } + if len(m.audioSamples) > 0 && m.audioSamples[len(m.audioSamples)-1].duration == 0 { + if len(m.audioSamples) > 1 { + m.audioSamples[len(m.audioSamples)-1].duration = m.audioSamples[len(m.audioSamples)-2].duration + } else { + m.audioSamples[0].duration = 1024 // typical AAC frame + } + } + + moov := m.buildMoov() + _, err := w.Write(moov) + return err +} + +func (m *Muxer) buildMoov() []byte { + var buf bytes.Buffer + + mvhd := m.buildMvhd() + writeBox(&buf, [4]byte{'m', 'v', 'h', 'd'}, mvhd) + + if len(m.videoSamples) > 0 { + trak := m.buildTrak(true) + writeBox(&buf, [4]byte{'t', 'r', 'a', 'k'}, trak) + } + + if len(m.audioSamples) > 0 { + trak := m.buildTrak(false) + writeBox(&buf, [4]byte{'t', 'r', 'a', 'k'}, trak) + } + + var out bytes.Buffer + writeBox(&out, [4]byte{'m', 'o', 'o', 'v'}, buf.Bytes()) + return out.Bytes() +} + +func (m *Muxer) totalDuration(samples []sampleEntry) uint64 { + var total uint64 + for _, s := range samples { + total += uint64(s.duration) + } + return total +} + +func (m *Muxer) buildMvhd() []byte { + d := m.totalDuration(m.videoSamples) + if len(m.audioSamples) > 0 { + ad := m.totalDuration(m.audioSamples) + if ad > d { + d = ad + } + } + + buf := make([]byte, 100) + putU32(buf[0:4], 0) // version + flags + putU32(buf[4:8], 0) // creation time + putU32(buf[8:12], 0) // modification time + putU32(buf[12:16], m.timescale) // timescale + putU32(buf[16:20], uint32(d)) // duration + putU32(buf[20:24], 0x00010000) // rate 1.0 + putU16(buf[24:26], 0x0100) // volume 1.0 + + // reserved + matrix + predefined + copy(buf[26:], make([]byte, 10+36+24)) + + // next_track_id + nextTrack := uint32(2) + if len(m.audioSamples) > 0 { + nextTrack = 3 + } + putU32(buf[96:100], nextTrack) + + return buf +} + +func (m *Muxer) buildTrak(isVideo bool) []byte { + var buf bytes.Buffer + + trackID := uint32(1) + if !isVideo { + trackID = 2 + } + + var samples []sampleEntry + ts := m.timescale + if isVideo { + samples = m.videoSamples + } else { + samples = m.audioSamples + ts = m.audioSampleRate + } + + dur := m.totalDuration(samples) + + tkhd := m.buildTkhd(trackID, uint32(dur), isVideo) + writeFullBox(&buf, [4]byte{'t', 'k', 'h', 'd'}, 0, 3, tkhd) + + mdia := m.buildMdia(isVideo, ts, samples) + writeBox(&buf, [4]byte{'m', 'd', 'i', 'a'}, mdia) + + return buf.Bytes() +} + +func (m *Muxer) buildTkhd(trackID, duration uint32, isVideo bool) []byte { + buf := make([]byte, 80) + putU32(buf[0:4], 0) // creation time + putU32(buf[4:8], 0) // modification time + putU32(buf[8:12], trackID) + // reserved 4 bytes + putU32(buf[16:20], duration) + // reserved 8 bytes + layer, alt group, volume + if !isVideo { + putU16(buf[32:34], 0x0100) // volume + } + // reserved 2 bytes + matrix (36 bytes) + // identity matrix + putU32(buf[36:40], 0x00010000) + putU32(buf[48:52], 0x00010000) + putU32(buf[60:64], 0x40000000) + + if isVideo { + putU32(buf[64:68], 1920<<16) // width + putU32(buf[68:72], 1080<<16) // height + } + + return buf[:72] +} + +func (m *Muxer) buildMdia(isVideo bool, timescale uint32, samples []sampleEntry) []byte { + var buf bytes.Buffer + + dur := m.totalDuration(samples) + mdhd := make([]byte, 24) + putU32(mdhd[8:12], timescale) + putU32(mdhd[12:16], uint32(dur)) + putU32(mdhd[16:20], 0x55C40000) // und language + writeFullBox(&buf, [4]byte{'m', 'd', 'h', 'd'}, 0, 0, mdhd) + + hdlr := m.buildHdlr(isVideo) + writeFullBox(&buf, [4]byte{'h', 'd', 'l', 'r'}, 0, 0, hdlr) + + minf := m.buildMinf(isVideo, samples) + writeBox(&buf, [4]byte{'m', 'i', 'n', 'f'}, minf) + + return buf.Bytes() +} + +func (m *Muxer) buildHdlr(isVideo bool) []byte { + var buf bytes.Buffer + buf.Write(make([]byte, 4)) // pre-defined + if isVideo { + buf.Write([]byte("vide")) + } else { + buf.Write([]byte("soun")) + } + buf.Write(make([]byte, 12)) // reserved + if isVideo { + buf.Write([]byte("VideoHandler\x00")) + } else { + buf.Write([]byte("SoundHandler\x00")) + } + return buf.Bytes() +} + +func (m *Muxer) buildMinf(isVideo bool, samples []sampleEntry) []byte { + var buf bytes.Buffer + + if isVideo { + vmhd := make([]byte, 8) + writeFullBox(&buf, [4]byte{'v', 'm', 'h', 'd'}, 0, 1, vmhd) + } else { + smhd := make([]byte, 4) + writeFullBox(&buf, [4]byte{'s', 'm', 'h', 'd'}, 0, 0, smhd) + } + + // dinf + dref (data reference: url self-contained) + var dref bytes.Buffer + putU32Buf(&dref, 1) // entry count + writeFullBox(&dref, [4]byte{'u', 'r', 'l', ' '}, 0, 1, nil) + var dinf bytes.Buffer + writeFullBox(&dinf, [4]byte{'d', 'r', 'e', 'f'}, 0, 0, dref.Bytes()) + writeBox(&buf, [4]byte{'d', 'i', 'n', 'f'}, dinf.Bytes()) + + stbl := m.buildStbl(isVideo, samples) + writeBox(&buf, [4]byte{'s', 't', 'b', 'l'}, stbl) + + return buf.Bytes() +} + +func (m *Muxer) buildStbl(isVideo bool, samples []sampleEntry) []byte { + var buf bytes.Buffer + + stsd := m.buildStsd(isVideo) + writeFullBox(&buf, [4]byte{'s', 't', 's', 'd'}, 0, 0, stsd) + + stts := buildStts(samples) + writeFullBox(&buf, [4]byte{'s', 't', 't', 's'}, 0, 0, stts) + + if isVideo { + ctts := buildCtts(samples) + if ctts != nil { + writeFullBox(&buf, [4]byte{'c', 't', 't', 's'}, 0, 0, ctts) + } + + stss := buildStss(samples) + if stss != nil { + writeFullBox(&buf, [4]byte{'s', 't', 's', 's'}, 0, 0, stss) + } + } + + stsc := buildStsc(len(samples)) + writeFullBox(&buf, [4]byte{'s', 't', 's', 'c'}, 0, 0, stsc) + + stsz := buildStsz(samples) + writeFullBox(&buf, [4]byte{'s', 't', 's', 'z'}, 0, 0, stsz) + + stco := buildStco(samples) + writeFullBox(&buf, [4]byte{'s', 't', 'c', 'o'}, 0, 0, stco) + + return buf.Bytes() +} + +func (m *Muxer) buildStsd(isVideo bool) []byte { + var buf bytes.Buffer + putU32Buf(&buf, 1) // entry count + + if isVideo && m.videoCodec == avframe.CodecH264 { + buf.Write(m.buildAvc1()) + } else if !isVideo && m.audioCodec == avframe.CodecAAC { + buf.Write(m.buildMp4a()) + } + + return buf.Bytes() +} + +func (m *Muxer) buildAvc1() []byte { + var buf bytes.Buffer + // SampleEntry header (78 bytes for video) + hdr := make([]byte, 78) + // data_reference_index = 1 + putU16(hdr[6:8], 1) + // width/height at offset 24/26 + putU16(hdr[24:26], 1920) + putU16(hdr[26:28], 1080) + // horiz/vert resolution + putU32(hdr[28:32], 0x00480000) // 72 dpi + putU32(hdr[32:36], 0x00480000) + // frame count + putU16(hdr[42:44], 1) + // depth + putU16(hdr[74:76], 0x0018) + + hdr[76] = 0xFF + hdr[77] = 0xFF + + buf.Write(hdr) + + // avcC box + if m.videoSeqHeader != nil { + writeBox(&buf, [4]byte{'a', 'v', 'c', 'C'}, m.videoSeqHeader) + } + + data := buf.Bytes() + var out bytes.Buffer + size := uint32(8 + len(data)) + var boxHdr [8]byte + binary.BigEndian.PutUint32(boxHdr[0:4], size) + copy(boxHdr[4:8], []byte("avc1")) + out.Write(boxHdr[:]) + out.Write(data) + return out.Bytes() +} + +func (m *Muxer) buildMp4a() []byte { + var buf bytes.Buffer + // SampleEntry header (28 bytes for audio) + hdr := make([]byte, 28) + putU16(hdr[6:8], 1) // data_reference_index + putU16(hdr[16:18], m.audioChannels) + putU16(hdr[18:20], 16) // sample size bits + putU32(hdr[24:28], m.audioSampleRate<<16) + + buf.Write(hdr) + + // esds box + esds := buildEsds(m.audioSeqHeader, m.audioSampleRate) + writeFullBox(&buf, [4]byte{'e', 's', 'd', 's'}, 0, 0, esds) + + data := buf.Bytes() + var out bytes.Buffer + size := uint32(8 + len(data)) + var boxHdr [8]byte + binary.BigEndian.PutUint32(boxHdr[0:4], size) + copy(boxHdr[4:8], []byte("mp4a")) + out.Write(boxHdr[:]) + out.Write(data) + return out.Bytes() +} + +func buildEsds(asc []byte, sampleRate uint32) []byte { + var buf bytes.Buffer + + // ES_Descriptor + ascLen := len(asc) + decConfigLen := 13 + 2 + ascLen + esLen := 3 + 2 + decConfigLen + 2 + 1 + + buf.WriteByte(0x03) // ES_DescrTag + buf.WriteByte(byte(esLen)) // length + putU16Buf(&buf, 1) // ES_ID + buf.WriteByte(0) // flags + + // DecoderConfigDescriptor + buf.WriteByte(0x04) // DecoderConfigDescrTag + buf.WriteByte(byte(decConfigLen)) + buf.WriteByte(0x40) // objectTypeIndication (AAC) + buf.WriteByte(0x15) // streamType (audio) + buf.Write([]byte{0x00, 0x00, 0x00}) // bufferSizeDB + putU32Buf(&buf, 0) // maxBitrate + putU32Buf(&buf, 0) // avgBitrate + + // DecoderSpecificInfo + buf.WriteByte(0x05) + buf.WriteByte(byte(ascLen)) + buf.Write(asc) + + // SLConfigDescriptor + buf.WriteByte(0x06) + buf.WriteByte(1) + buf.WriteByte(0x02) + + return buf.Bytes() +} + +func buildStts(samples []sampleEntry) []byte { + if len(samples) == 0 { + buf := make([]byte, 4) + return buf + } + + type sttsEntry struct { + count uint32 + duration uint32 + } + + var entries []sttsEntry + for _, s := range samples { + if len(entries) > 0 && entries[len(entries)-1].duration == s.duration { + entries[len(entries)-1].count++ + } else { + entries = append(entries, sttsEntry{count: 1, duration: s.duration}) + } + } + + buf := make([]byte, 4+len(entries)*8) + putU32(buf[0:4], uint32(len(entries))) + for i, e := range entries { + off := 4 + i*8 + putU32(buf[off:off+4], e.count) + putU32(buf[off+4:off+8], e.duration) + } + return buf +} + +func buildCtts(samples []sampleEntry) []byte { + hasCTS := false + for _, s := range samples { + if s.cts != 0 { + hasCTS = true + break + } + } + if !hasCTS { + return nil + } + + type cttsEntry struct { + count uint32 + offset int32 + } + + var entries []cttsEntry + for _, s := range samples { + if len(entries) > 0 && entries[len(entries)-1].offset == s.cts { + entries[len(entries)-1].count++ + } else { + entries = append(entries, cttsEntry{count: 1, offset: s.cts}) + } + } + + buf := make([]byte, 4+len(entries)*8) + putU32(buf[0:4], uint32(len(entries))) + for i, e := range entries { + off := 4 + i*8 + putU32(buf[off:off+4], e.count) + putU32(buf[off+4:off+8], uint32(e.offset)) + } + return buf +} + +func buildStss(samples []sampleEntry) []byte { + var syncIndices []uint32 + for i, s := range samples { + if s.isSync { + syncIndices = append(syncIndices, uint32(i+1)) + } + } + if len(syncIndices) == len(samples) { + return nil // all sync, no need for stss + } + + buf := make([]byte, 4+len(syncIndices)*4) + putU32(buf[0:4], uint32(len(syncIndices))) + for i, idx := range syncIndices { + putU32(buf[4+i*4:8+i*4], idx) + } + return buf +} + +func buildStsc(sampleCount int) []byte { + // One chunk per sample (simplest approach) + buf := make([]byte, 4+12) + putU32(buf[0:4], 1) // entry count + putU32(buf[4:8], 1) // first chunk + putU32(buf[8:12], 1) // samples per chunk + putU32(buf[12:16], 1) // sample description index + return buf +} + +func buildStsz(samples []sampleEntry) []byte { + buf := make([]byte, 8+len(samples)*4) + putU32(buf[0:4], 0) // sample_size = 0 (variable) + putU32(buf[4:8], uint32(len(samples))) + for i, s := range samples { + putU32(buf[8+i*4:12+i*4], s.size) + } + return buf +} + +func buildStco(samples []sampleEntry) []byte { + buf := make([]byte, 4+len(samples)*4) + putU32(buf[0:4], uint32(len(samples))) + for i, s := range samples { + putU32(buf[4+i*4:8+i*4], uint32(s.offset)) + } + return buf +} + +func putU32Buf(buf *bytes.Buffer, v uint32) { + var b [4]byte + binary.BigEndian.PutUint32(b[:], v) + buf.Write(b[:]) +} + +func putU16Buf(buf *bytes.Buffer, v uint16) { + var b [2]byte + binary.BigEndian.PutUint16(b[:], v) + buf.Write(b[:]) +} diff --git a/pkg/muxer/mp4/muxer_test.go b/pkg/muxer/mp4/muxer_test.go new file mode 100644 index 0000000..6e2e3de --- /dev/null +++ b/pkg/muxer/mp4/muxer_test.go @@ -0,0 +1,154 @@ +package mp4 + +import ( + "bytes" + "encoding/binary" + "io" + "testing" + + "github.com/im-pingo/liveforge/pkg/avframe" +) + +type memSeeker struct { + bytes.Buffer + pos int64 +} + +func (m *memSeeker) Write(p []byte) (int, error) { + // Ensure we write at the current position + if m.pos < int64(m.Len()) { + // Overwrite existing data + data := m.Bytes() + n := copy(data[m.pos:], p) + if n < len(p) { + m.Buffer.Write(p[n:]) + } + m.pos += int64(len(p)) + return len(p), nil + } + // Extend + if m.pos > int64(m.Len()) { + padding := make([]byte, m.pos-int64(m.Len())) + m.Buffer.Write(padding) + } + n, err := m.Buffer.Write(p) + m.pos += int64(n) + return n, err +} + +func (m *memSeeker) Seek(offset int64, whence int) (int64, error) { + switch whence { + case io.SeekStart: + m.pos = offset + case io.SeekCurrent: + m.pos += offset + case io.SeekEnd: + m.pos = int64(m.Len()) + offset + } + return m.pos, nil +} + +func TestMuxerFtypBox(t *testing.T) { + m := NewMuxer(avframe.CodecH264, avframe.CodecAAC) + var buf bytes.Buffer + m.WriteFtyp(&buf) + + data := buf.Bytes() + if len(data) < 8 { + t.Fatal("ftyp too short") + } + boxType := string(data[4:8]) + if boxType != "ftyp" { + t.Errorf("expected ftyp box, got %s", boxType) + } +} + +func TestMuxerWriteAndFinalize(t *testing.T) { + m := NewMuxer(avframe.CodecH264, avframe.CodecAAC) + w := &memSeeker{} + + m.WriteFtyp(w) + m.WriteMdatHeader(w) + + // Write video sequence header + seqFrame := avframe.NewAVFrame( + avframe.MediaTypeVideo, avframe.CodecH264, avframe.FrameTypeSequenceHeader, + 0, 0, []byte{0x01, 0x64, 0x00, 0x28, 0xFF, 0xE1, 0x00, 0x04, 0x67, 0x64, 0x00, 0x28, 0x01, 0x00, 0x04, 0x68, 0xEE, 0x3C, 0x80}, + ) + m.WriteFrame(w, seqFrame, -1) + + // Write some video frames + var prevDTS int64 = -1 + for i := range 5 { + ft := avframe.FrameTypeInterframe + if i == 0 { + ft = avframe.FrameTypeKeyframe + } + dts := int64(i * 33) + pts := dts + frame := avframe.NewAVFrame( + avframe.MediaTypeVideo, avframe.CodecH264, ft, + pts, dts, []byte{0x00, 0x00, 0x00, byte(i + 1), 0x65, 0x88}, + ) + m.WriteFrame(w, frame, prevDTS) + prevDTS = dts + } + + if err := m.Finalize(w); err != nil { + t.Fatalf("Finalize: %v", err) + } + + data := w.Bytes() + if len(data) < 100 { + t.Fatalf("output too small: %d bytes", len(data)) + } + + // Verify we have ftyp, mdat, and moov boxes + foundFtyp := false + foundMdat := false + foundMoov := false + pos := 0 + for pos < len(data)-8 { + size := binary.BigEndian.Uint32(data[pos : pos+4]) + boxType := string(data[pos+4 : pos+8]) + + switch boxType { + case "ftyp": + foundFtyp = true + case "mdat": + foundMdat = true + case "moov": + foundMoov = true + } + + if size < 8 || int(size) > len(data)-pos { + break + } + pos += int(size) + } + + if !foundFtyp { + t.Error("missing ftyp box") + } + if !foundMdat { + t.Error("missing mdat box") + } + if !foundMoov { + t.Error("missing moov box") + } + + if len(m.videoSamples) != 5 { + t.Errorf("expected 5 video samples, got %d", len(m.videoSamples)) + } +} + +func TestMuxerEmptyFinalize(t *testing.T) { + m := NewMuxer(avframe.CodecH264, avframe.CodecAAC) + w := &memSeeker{} + m.WriteFtyp(w) + m.WriteMdatHeader(w) + + if err := m.Finalize(w); err != nil { + t.Fatalf("Finalize on empty: %v", err) + } +} diff --git a/pkg/muxer/ts/muxer.go b/pkg/muxer/ts/muxer.go index 7bed4e9..8f8e291 100644 --- a/pkg/muxer/ts/muxer.go +++ b/pkg/muxer/ts/muxer.go @@ -109,7 +109,9 @@ func (m *Muxer) WriteFrame(frame *avframe.AVFrame) []byte { } func (m *Muxer) writeVideoFrame(frame *avframe.AVFrame) []byte { - var result []byte + // Estimate output size: PAT+PMT (376 bytes) + PES header + payload + TS overhead + estSize := len(frame.Payload)*2 + 1024 + result := make([]byte, 0, estSize) // Prepend PAT+PMT before keyframes if frame.FrameType.IsKeyframe() { @@ -123,18 +125,18 @@ func (m *Muxer) writeVideoFrame(frame *avframe.AVFrame) []byte { case avframe.CodecH264: annexB := h264.AVCCToAnnexB(frame.Payload) if frame.FrameType.IsKeyframe() && m.videoSeqHeader != nil { - // Prepend SPS/PPS before keyframe + payload = make([]byte, 0, len(m.videoSeqHeader)+len(annexB)) payload = append(payload, m.videoSeqHeader...) } payload = append(payload, annexB...) case avframe.CodecH265: annexB := h265.HVCCToAnnexB(frame.Payload) if frame.FrameType.IsKeyframe() && m.videoSeqHeader != nil { + payload = make([]byte, 0, len(m.videoSeqHeader)+len(annexB)) payload = append(payload, m.videoSeqHeader...) } payload = append(payload, annexB...) case avframe.CodecAV1: - // AV1 OBUs passed through directly payload = frame.Payload default: payload = frame.Payload @@ -142,21 +144,19 @@ func (m *Muxer) writeVideoFrame(frame *avframe.AVFrame) []byte { // Build PES pesHeader := BuildPESHeader(0xE0, frame.PTS, frame.DTS, len(payload)) - pesData := append(pesHeader, payload...) + pesData := make([]byte, 0, len(pesHeader)+len(payload)) + pesData = append(pesData, pesHeader...) + pesData = append(pesData, payload...) - // Build packetization options: embed PCR in first TS packet if needed, - // set random_access_indicator for keyframes. This avoids separate - // PCR-only packets that cause continuity counter issues with demuxers. opts := &PESPacketizeOptions{ RandomAccess: frame.FrameType.IsKeyframe(), - PCR: -1, // no PCR by default + PCR: -1, } if m.shouldInsertPCR(frame.DTS) { opts.PCR = frame.DTS m.lastPCR = frame.DTS } - // Packetize into TS tsPackets := PacketizePES(PIDVideo, pesData, &m.videoContinuity, opts) result = append(result, tsPackets...) @@ -169,13 +169,12 @@ func (m *Muxer) writeAudioFrame(frame *avframe.AVFrame) []byte { switch m.audioCodec { case avframe.CodecAAC: if m.aacInfo != nil { - // Prepend ADTS header adts := aac.BuildADTSHeader(m.aacInfo, len(frame.Payload)) + payload = make([]byte, 0, len(adts)+len(frame.Payload)) payload = append(payload, adts...) } payload = append(payload, frame.Payload...) case avframe.CodecMP3: - // MP3 is self-delimiting, pass through payload = frame.Payload case avframe.CodecOpus: payload = frame.Payload @@ -183,9 +182,10 @@ func (m *Muxer) writeAudioFrame(frame *avframe.AVFrame) []byte { payload = frame.Payload } - // Build PES (audio uses PTS only, no DTS) pesHeader := BuildPESHeader(0xC0, frame.PTS, frame.PTS, len(payload)) - pesData := append(pesHeader, payload...) + pesData := make([]byte, 0, len(pesHeader)+len(payload)) + pesData = append(pesData, pesHeader...) + pesData = append(pesData, payload...) tsPackets := PacketizePES(PIDAudio, pesData, &m.audioContinuity, nil) return tsPackets diff --git a/pkg/pool/pool.go b/pkg/pool/pool.go new file mode 100644 index 0000000..ab91afe --- /dev/null +++ b/pkg/pool/pool.go @@ -0,0 +1,61 @@ +package pool + +import ( + "sync" +) + +const ( + smallBufSize = 256 + mediumBufSize = 4096 + largeBufSize = 65536 +) + +var ( + smallPool = sync.Pool{ + New: func() any { b := make([]byte, 0, smallBufSize); return &b }, + } + mediumPool = sync.Pool{ + New: func() any { b := make([]byte, 0, mediumBufSize); return &b }, + } + largePool = sync.Pool{ + New: func() any { b := make([]byte, 0, largeBufSize); return &b }, + } +) + +// GetBuffer returns a pooled byte slice with at least the given capacity. +// The returned slice has length 0 and the caller must append to it. +// Call PutBuffer when done to return it to the pool. +func GetBuffer(minCap int) *[]byte { + if minCap <= smallBufSize { + return smallPool.Get().(*[]byte) + } + if minCap <= mediumBufSize { + return mediumPool.Get().(*[]byte) + } + if minCap <= largeBufSize { + return largePool.Get().(*[]byte) + } + b := make([]byte, 0, minCap) + return &b +} + +// PutBuffer returns a buffer to the pool. The buffer is reset to zero length. +// Nil pointers and oversized buffers (>128KB) are silently discarded. +func PutBuffer(b *[]byte) { + if b == nil { + return + } + c := cap(*b) + *b = (*b)[:0] + + if c > largeBufSize*2 { + return + } + if c >= largeBufSize { + largePool.Put(b) + } else if c >= mediumBufSize { + mediumPool.Put(b) + } else { + smallPool.Put(b) + } +} diff --git a/pkg/pool/pool_test.go b/pkg/pool/pool_test.go new file mode 100644 index 0000000..897f3dd --- /dev/null +++ b/pkg/pool/pool_test.go @@ -0,0 +1,85 @@ +package pool + +import ( + "testing" +) + +func TestGetPutBuffer(t *testing.T) { + tests := []struct { + name string + minCap int + wantMinCap int + }{ + {"small", 100, smallBufSize}, + {"medium", 1000, mediumBufSize}, + {"large", 10000, largeBufSize}, + {"oversized", 200000, 200000}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + b := GetBuffer(tt.minCap) + if b == nil { + t.Fatal("GetBuffer returned nil") + } + if len(*b) != 0 { + t.Errorf("expected length 0, got %d", len(*b)) + } + if cap(*b) < tt.wantMinCap { + t.Errorf("expected cap >= %d, got %d", tt.wantMinCap, cap(*b)) + } + *b = append(*b, 1, 2, 3) + PutBuffer(b) + }) + } +} + +func TestPutNilBuffer(t *testing.T) { + PutBuffer(nil) +} + +func TestBufferReuse(t *testing.T) { + b1 := GetBuffer(100) + ptr1 := &(*b1)[0:cap(*b1)][0] + PutBuffer(b1) + + b2 := GetBuffer(100) + ptr2 := &(*b2)[0:cap(*b2)][0] + + if ptr1 != ptr2 { + t.Log("buffer was not reused (may happen under contention)") + } + PutBuffer(b2) +} + +func BenchmarkGetPutSmall(b *testing.B) { + for b.Loop() { + buf := GetBuffer(100) + *buf = append(*buf, make([]byte, 100)...) + PutBuffer(buf) + } +} + +func BenchmarkGetPutMedium(b *testing.B) { + for b.Loop() { + buf := GetBuffer(2000) + *buf = append(*buf, make([]byte, 2000)...) + PutBuffer(buf) + } +} + +func BenchmarkAllocSmall(b *testing.B) { + for b.Loop() { + buf := make([]byte, 0, 256) + buf = append(buf, make([]byte, 100)...) + _ = buf + } +} + +func BenchmarkAllocMedium(b *testing.B) { + for b.Loop() { + buf := make([]byte, 0, 4096) + buf = append(buf, make([]byte, 2000)...) + _ = buf + } +} diff --git a/pkg/sdp/builder.go b/pkg/sdp/builder.go index 3fbb9e3..e776d03 100644 --- a/pkg/sdp/builder.go +++ b/pkg/sdp/builder.go @@ -106,6 +106,9 @@ func buildMediaDesc(mediaType string, codec avframe.CodecType, sampleRate, chann if clockRate == 0 { clockRate = sampleRate } + if clockRate == 0 && codec == avframe.CodecAAC { + clockRate = 44100 + } md := &MediaDescription{ Type: mediaType, diff --git a/test/integration/protocol_test.go b/test/integration/protocol_test.go new file mode 100644 index 0000000..fb80797 --- /dev/null +++ b/test/integration/protocol_test.go @@ -0,0 +1,241 @@ +package integration + +import ( + "testing" + "time" + + "github.com/im-pingo/liveforge/config" + "github.com/im-pingo/liveforge/core" + "github.com/im-pingo/liveforge/pkg/avframe" + rtspmod "github.com/im-pingo/liveforge/module/rtsp" + srtmod "github.com/im-pingo/liveforge/module/srt" +) + +func TestRTSPServerStartStop(t *testing.T) { + cfg := &config.Config{} + cfg.Server.Name = "rtsp-integration" + cfg.RTSP.Enabled = true + cfg.RTSP.Listen = "127.0.0.1:0" + cfg.RTSP.RTPPortRange = []int{30000, 30100} + cfg.Stream.GOPCache = true + cfg.Stream.GOPCacheNum = 1 + cfg.Stream.RingBufferSize = 256 + cfg.Stream.NoPublisherTimeout = 5 * time.Second + + s := core.NewServer(cfg) + s.RegisterModule(rtspmod.NewModule()) + + if err := s.Init(); err != nil { + t.Fatalf("server init: %v", err) + } + + s.Shutdown() +} + +func TestSRTServerStartStop(t *testing.T) { + cfg := &config.Config{} + cfg.Server.Name = "srt-integration" + cfg.SRT.Enabled = true + cfg.SRT.Listen = "127.0.0.1:0" + cfg.SRT.Latency = 120 + cfg.Stream.GOPCache = true + cfg.Stream.GOPCacheNum = 1 + cfg.Stream.RingBufferSize = 256 + cfg.Stream.NoPublisherTimeout = 5 * time.Second + + s := core.NewServer(cfg) + s.RegisterModule(srtmod.NewModule()) + + if err := s.Init(); err != nil { + t.Fatalf("server init: %v", err) + } + + s.Shutdown() +} + +func TestMultiProtocolPublishSubscribe(t *testing.T) { + bus := core.NewEventBus() + cfg := config.StreamConfig{ + GOPCache: true, + GOPCacheNum: 2, + AudioCacheMs: 1000, + RingBufferSize: 512, + NoPublisherTimeout: 5 * time.Second, + } + + hub := core.NewStreamHub(cfg, config.LimitsConfig{}, bus) + + stream, err := hub.GetOrCreate("live/multiproto-test") + if err != nil { + t.Fatalf("GetOrCreate: %v", err) + } + + pub := &testPublisher{ + id: "rtsp-pub", + info: &avframe.MediaInfo{ + VideoCodec: avframe.CodecH264, + AudioCodec: avframe.CodecAAC, + }, + } + if err := stream.SetPublisher(pub); err != nil { + t.Fatalf("SetPublisher: %v", err) + } + + // Write sequence headers + stream.WriteFrame(avframe.NewAVFrame( + avframe.MediaTypeVideo, avframe.CodecH264, avframe.FrameTypeSequenceHeader, + 0, 0, []byte{0x01, 0x64, 0x00, 0x28}, + )) + stream.WriteFrame(avframe.NewAVFrame( + avframe.MediaTypeAudio, avframe.CodecAAC, avframe.FrameTypeSequenceHeader, + 0, 0, []byte{0x12, 0x10}, + )) + + // Write two GOPs + for gop := 0; gop < 2; gop++ { + baseTS := int64(gop * 2000) + stream.WriteFrame(avframe.NewAVFrame( + avframe.MediaTypeVideo, avframe.CodecH264, avframe.FrameTypeKeyframe, + baseTS, baseTS, []byte{0x65, 0x88, byte(gop)}, + )) + for i := 1; i <= 4; i++ { + ts := baseTS + int64(i*40) + stream.WriteFrame(avframe.NewAVFrame( + avframe.MediaTypeVideo, avframe.CodecH264, avframe.FrameTypeInterframe, + ts, ts, []byte{0x41, byte(i)}, + )) + stream.WriteFrame(avframe.NewAVFrame( + avframe.MediaTypeAudio, avframe.CodecAAC, avframe.FrameTypeInterframe, + ts, ts, []byte{0xFF, byte(i)}, + )) + } + } + + // Verify GOP cache has 2 GOPs + gop := stream.GOPCache() + keyframeCount := 0 + for _, f := range gop { + if f.FrameType == avframe.FrameTypeKeyframe { + keyframeCount++ + } + } + if keyframeCount != 2 { + t.Errorf("expected 2 keyframes in GOP cache, got %d", keyframeCount) + } + + // Multiple concurrent readers + readers := make([]*avframe.AVFrame, 0) + reader := stream.RingBuffer().NewReader() + for { + frame, ok := reader.TryRead() + if !ok { + break + } + readers = append(readers, frame) + } + + // 2 seq headers + 2 GOPs * (1 keyframe + 4 inter + 4 audio) = 2 + 18 = 20 + if len(readers) != 20 { + t.Errorf("expected 20 frames, got %d", len(readers)) + } + + stream.RemovePublisher() + hub.Remove("live/multiproto-test") +} + +func TestHLSMuxerLifecycle(t *testing.T) { + bus := core.NewEventBus() + cfg := config.StreamConfig{ + GOPCache: true, + GOPCacheNum: 1, + RingBufferSize: 256, + NoPublisherTimeout: 5 * time.Second, + } + + hub := core.NewStreamHub(cfg, config.LimitsConfig{}, bus) + + stream, err := hub.GetOrCreate("live/hls-test") + if err != nil { + t.Fatalf("GetOrCreate: %v", err) + } + + pub := &testPublisher{ + id: "hls-pub", + info: &avframe.MediaInfo{VideoCodec: avframe.CodecH264, AudioCodec: avframe.CodecAAC}, + } + if err := stream.SetPublisher(pub); err != nil { + t.Fatalf("SetPublisher: %v", err) + } + + mm := stream.MuxerManager() + + // Register a test muxer to verify lifecycle + started := make(chan struct{}, 1) + mm.RegisterMuxerStart("flv", func(inst *core.MuxerInstance, s *core.Stream) { + started <- struct{}{} + }) + + // Request the muxer + reader, inst := mm.GetOrCreateMuxer("flv") + if inst == nil { + t.Fatal("expected muxer instance") + } + _ = reader + + select { + case <-started: + case <-time.After(time.Second): + t.Fatal("muxer start callback not called within 1s") + } + + stream.RemovePublisher() + hub.Remove("live/hls-test") +} + +func TestStreamStatsIntegration(t *testing.T) { + bus := core.NewEventBus() + cfg := config.StreamConfig{ + GOPCache: true, + GOPCacheNum: 1, + RingBufferSize: 256, + NoPublisherTimeout: 5 * time.Second, + } + + hub := core.NewStreamHub(cfg, config.LimitsConfig{}, bus) + + stream, err := hub.GetOrCreate("live/stats-test") + if err != nil { + t.Fatalf("GetOrCreate: %v", err) + } + + pub := &testPublisher{ + id: "stats-pub", + info: &avframe.MediaInfo{VideoCodec: avframe.CodecH264}, + } + if err := stream.SetPublisher(pub); err != nil { + t.Fatalf("SetPublisher: %v", err) + } + + // Write frames and verify stats + stream.WriteFrame(avframe.NewAVFrame( + avframe.MediaTypeVideo, avframe.CodecH264, avframe.FrameTypeSequenceHeader, + 0, 0, []byte{0x01, 0x64, 0x00, 0x28}, + )) + + payload := make([]byte, 1000) + stream.WriteFrame(avframe.NewAVFrame( + avframe.MediaTypeVideo, avframe.CodecH264, avframe.FrameTypeKeyframe, + 0, 0, payload, + )) + + stats := stream.Stats() + if stats.BytesIn == 0 { + t.Error("expected non-zero bytes in") + } + if stats.VideoFrames == 0 { + t.Error("expected non-zero video frames") + } + + stream.RemovePublisher() + hub.Remove("live/stats-test") +}