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")
+}