stremio/node/media-node/internal/streaming/server.go
Jos Vooges | STH 359a3c6883 Dedupe live streams to one row per viewer and ignore probe sessions.
Co-authored-by: Cursor <cursoragent@cursor.com>
2026-08-29 01:35:20 +02:00

481 lines
12 KiB
Go

package streaming
import (
"crypto/sha256"
"encoding/hex"
"fmt"
"io"
"log"
"net"
"net/http"
"os"
"path/filepath"
"strconv"
"strings"
"sync"
"sync/atomic"
"time"
"github.com/sthmedia/media-node/internal/database"
)
type Server struct {
store *database.Store
listen string
maxPerIP int
mu sync.RWMutex
activeCount int
bytesWindow int64
lastWindow time.Time
bytesPerSecond int64
ipConns map[string]int
onSessionEnd func(sessionID, reason string)
onSessionProgress func(sessionID string, bytesSent, lastByteOffset, fileSizeBytes int64)
progressMu sync.Mutex
progressState map[string]*progressTracker
}
type progressTracker struct {
bytesSent int64
lastOffset int64
fileSize int64
lastReport time.Time
}
func New(store *database.Store, listen string, maxPerIP int) *Server {
// Stremio/VLC open many parallel Range requests; behind NPM they often share one IP.
if maxPerIP <= 0 {
maxPerIP = 64
}
s := &Server{
store: store,
listen: listen,
maxPerIP: maxPerIP,
lastWindow: time.Now(),
ipConns: make(map[string]int),
progressState: make(map[string]*progressTracker),
}
go s.cleanupLoop()
go s.bandwidthLoop()
return s
}
func (s *Server) SetSessionEndHandler(fn func(sessionID, reason string)) {
s.onSessionEnd = fn
}
func (s *Server) SetSessionProgressHandler(fn func(sessionID string, bytesSent, lastByteOffset, fileSizeBytes int64)) {
s.onSessionProgress = fn
}
func (s *Server) ListenAndServe() error {
mux := http.NewServeMux()
mux.HandleFunc("/play/", s.handlePlay)
mux.HandleFunc("/health", s.handleHealth)
mux.HandleFunc("/", func(w http.ResponseWriter, r *http.Request) {
if r.URL.Path != "/" {
http.NotFound(w, r)
return
}
http.NotFound(w, r)
})
server := &http.Server{
Addr: s.listen,
Handler: corsMiddleware(mux),
ReadHeaderTimeout: 10 * time.Second,
IdleTimeout: 120 * time.Second,
}
log.Printf("Stream server listening on %s", s.listen)
return server.ListenAndServe()
}
func corsMiddleware(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Access-Control-Allow-Origin", "*")
w.Header().Set("Access-Control-Allow-Methods", "GET, HEAD, OPTIONS")
w.Header().Set("Access-Control-Allow-Headers", "Range, Content-Type")
w.Header().Set("Access-Control-Expose-Headers", "Content-Length, Content-Range, Accept-Ranges")
if r.Method == http.MethodOptions {
w.WriteHeader(http.StatusNoContent)
return
}
next.ServeHTTP(w, r)
})
}
func (s *Server) handleHealth(w http.ResponseWriter, r *http.Request) {
host, _, _ := net.SplitHostPort(r.RemoteAddr)
if host != "127.0.0.1" && host != "::1" && host != "localhost" {
// Prefer internal-only health; still allow LAN NPM if needed via 200 for GET from private ranges
ip := net.ParseIP(host)
if ip == nil || !ip.IsLoopback() && !ip.IsPrivate() {
http.NotFound(w, r)
return
}
}
w.Header().Set("Content-Type", "application/json")
w.WriteHeader(http.StatusOK)
_, _ = w.Write([]byte(`{"status":"ok"}`))
}
func (s *Server) handlePlay(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodGet && r.Method != http.MethodHead {
http.NotFound(w, r)
return
}
token := strings.TrimPrefix(r.URL.Path, "/play/")
if token == "" || strings.Contains(token, "/") {
http.NotFound(w, r)
return
}
ip := clientIP(r)
log.Printf("GET /play request from %s method=%s range=%q", ip, r.Method, r.Header.Get("Range"))
if !s.acquireIP(ip) {
log.Printf("GET /play/[REDACTED] rejected: too many connections from %s", ip)
http.Error(w, "Too Many Requests", http.StatusTooManyRequests)
return
}
defer s.releaseIP(ip)
tokenHash := hashToken(token)
sess, err := s.store.GetSessionByTokenHash(tokenHash)
if err != nil {
// Brief grace: Master may have redirected before WS session landed.
for i := 0; i < 40 && err != nil; i++ {
time.Sleep(50 * time.Millisecond)
sess, err = s.store.GetSessionByTokenHash(tokenHash)
}
}
if err != nil {
log.Printf("GET /play/[REDACTED] invalid session from %s", ip)
http.NotFound(w, r)
return
}
now := time.Now().UTC()
if now.After(sess.AbsoluteExpiresAt) {
_ = s.store.DeleteSession(sess.SessionID)
s.clearProgress(sess.SessionID)
if s.onSessionEnd != nil {
s.onSessionEnd(sess.SessionID, "absolute_expiry")
}
http.NotFound(w, r)
return
}
idleDeadline := sess.LastActivity.Add(time.Duration(sess.IdleTimeoutSeconds) * time.Second)
if now.After(idleDeadline) {
_ = s.store.DeleteSession(sess.SessionID)
s.clearProgress(sess.SessionID)
if s.onSessionEnd != nil {
s.onSessionEnd(sess.SessionID, "idle_timeout")
}
http.NotFound(w, r)
return
}
file, err := s.store.GetFile(sess.LocalFileID)
if err != nil {
log.Printf("GET /play/[REDACTED] unknown localFileId=%s session=%s", sess.LocalFileID, sess.SessionID)
http.NotFound(w, r)
return
}
f, err := os.Open(file.Path)
if err != nil {
log.Printf("GET /play/[REDACTED] file open failed session=%s path=%s err=%v", sess.SessionID, file.Path, err)
http.NotFound(w, r)
return
}
defer f.Close()
stat, err := f.Stat()
if err != nil {
http.NotFound(w, r)
return
}
fileSize := stat.Size()
_ = s.store.TouchSession(sess.SessionID)
s.noteProgress(sess.SessionID, 0, fileSize, 0, true)
s.mu.Lock()
s.activeCount++
s.mu.Unlock()
defer func() {
s.mu.Lock()
s.activeCount--
s.mu.Unlock()
}()
w.Header().Set("Accept-Ranges", "bytes")
w.Header().Set("Content-Type", contentType(file.Path))
w.Header().Set("Cache-Control", "no-store")
w.Header().Set("Access-Control-Allow-Origin", "*")
w.Header().Set("Access-Control-Expose-Headers", "Content-Length, Content-Range, Accept-Ranges")
rangeHeader := r.Header.Get("Range")
if rangeHeader == "" {
w.Header().Set("Content-Length", strconv.FormatInt(fileSize, 10))
if r.Method == http.MethodHead {
w.WriteHeader(http.StatusOK)
return
}
s.streamRange(w, f, 0, fileSize-1, sess.SessionID, fileSize)
return
}
start, end, err := parseRange(rangeHeader, fileSize)
if err != nil {
w.Header().Set("Content-Range", fmt.Sprintf("bytes */%d", fileSize))
http.Error(w, "Invalid Range", http.StatusRequestedRangeNotSatisfiable)
return
}
contentLength := end - start + 1
w.Header().Set("Content-Range", fmt.Sprintf("bytes %d-%d/%d", start, end, fileSize))
w.Header().Set("Content-Length", strconv.FormatInt(contentLength, 10))
w.WriteHeader(http.StatusPartialContent)
if r.Method == http.MethodHead {
return
}
s.noteProgress(sess.SessionID, start, fileSize, 0, true)
s.streamRange(w, f, start, end, sess.SessionID, fileSize)
}
func (s *Server) streamRange(w http.ResponseWriter, f *os.File, start, end int64, sessionID string, fileSize int64) {
if _, err := f.Seek(start, io.SeekStart); err != nil {
return
}
remaining := end - start + 1
buf := make([]byte, 256*1024)
flusher, canFlush := w.(http.Flusher)
pos := start
for remaining > 0 {
toRead := int64(len(buf))
if toRead > remaining {
toRead = remaining
}
n, err := f.Read(buf[:toRead])
if n > 0 {
atomic.AddInt64(&s.bytesWindow, int64(n))
if _, werr := w.Write(buf[:n]); werr != nil {
return
}
if canFlush {
flusher.Flush()
}
remaining -= int64(n)
pos += int64(n)
s.noteProgress(sessionID, pos, fileSize, int64(n), false)
}
if err != nil {
return
}
}
}
func (s *Server) noteProgress(sessionID string, offset, fileSize, deltaBytes int64, force bool) {
if sessionID == "" {
return
}
_ = s.store.TouchSession(sessionID)
s.progressMu.Lock()
st := s.progressState[sessionID]
if st == nil {
st = &progressTracker{}
s.progressState[sessionID] = st
}
if deltaBytes > 0 {
st.bytesSent += deltaBytes
// Alleen offset bijwerken bij echte data — voorkomt 100% door 1-byte end-of-file probes.
if offset > st.lastOffset {
st.lastOffset = offset
}
}
if fileSize > 0 {
st.fileSize = fileSize
}
shouldReport := (deltaBytes > 0 || force) &&
st.bytesSent > 0 &&
(force || time.Since(st.lastReport) >= 3*time.Second)
bytesSent := st.bytesSent
lastOffset := st.lastOffset
size := st.fileSize
if shouldReport {
st.lastReport = time.Now()
}
s.progressMu.Unlock()
if shouldReport && s.onSessionProgress != nil {
s.onSessionProgress(sessionID, bytesSent, lastOffset, size)
}
}
func (s *Server) ActiveStreams() int {
s.mu.RLock()
defer s.mu.RUnlock()
return s.activeCount
}
func (s *Server) BytesPerSecond() int64 {
return atomic.LoadInt64(&s.bytesPerSecond)
}
func (s *Server) CurrentBandwidth() int64 {
return s.BytesPerSecond()
}
func (s *Server) CreateSession(sess database.PlaybackSession) error {
return s.store.SaveSession(sess)
}
func (s *Server) RevokeSession(sessionID string) error {
s.clearProgress(sessionID)
return s.store.DeleteSession(sessionID)
}
func (s *Server) acquireIP(ip string) bool {
s.mu.Lock()
defer s.mu.Unlock()
if s.ipConns[ip] >= s.maxPerIP {
return false
}
s.ipConns[ip]++
return true
}
func (s *Server) releaseIP(ip string) {
s.mu.Lock()
defer s.mu.Unlock()
if s.ipConns[ip] <= 1 {
delete(s.ipConns, ip)
return
}
s.ipConns[ip]--
}
func (s *Server) bandwidthLoop() {
ticker := time.NewTicker(time.Second)
for range ticker.C {
n := atomic.SwapInt64(&s.bytesWindow, 0)
atomic.StoreInt64(&s.bytesPerSecond, n)
}
}
func (s *Server) cleanupLoop() {
ticker := time.NewTicker(60 * time.Second)
for range ticker.C {
sessions, _ := s.store.ActiveSessions()
now := time.Now().UTC()
for _, sess := range sessions {
if now.After(sess.AbsoluteExpiresAt) {
_ = s.store.DeleteSession(sess.SessionID)
s.clearProgress(sess.SessionID)
if s.onSessionEnd != nil {
s.onSessionEnd(sess.SessionID, "absolute_expiry")
}
continue
}
idleDeadline := sess.LastActivity.Add(time.Duration(sess.IdleTimeoutSeconds) * time.Second)
if now.After(idleDeadline) {
_ = s.store.DeleteSession(sess.SessionID)
s.clearProgress(sess.SessionID)
if s.onSessionEnd != nil {
s.onSessionEnd(sess.SessionID, "idle_timeout")
}
}
}
}
}
func (s *Server) clearProgress(sessionID string) {
s.progressMu.Lock()
delete(s.progressState, sessionID)
s.progressMu.Unlock()
}
func parseRange(rangeHeader string, size int64) (int64, int64, error) {
if !strings.HasPrefix(rangeHeader, "bytes=") {
return 0, 0, fmt.Errorf("invalid range")
}
parts := strings.Split(strings.TrimPrefix(rangeHeader, "bytes="), "-")
if len(parts) != 2 {
return 0, 0, fmt.Errorf("invalid range")
}
var start, end int64
if parts[0] == "" {
suffix, err := strconv.ParseInt(parts[1], 10, 64)
if err != nil {
return 0, 0, err
}
start = size - suffix
if start < 0 {
start = 0
}
end = size - 1
} else {
var err error
start, err = strconv.ParseInt(parts[0], 10, 64)
if err != nil {
return 0, 0, err
}
if parts[1] == "" {
end = size - 1
} else {
end, err = strconv.ParseInt(parts[1], 10, 64)
if err != nil {
return 0, 0, err
}
}
}
if start > end || start >= size {
return 0, 0, fmt.Errorf("invalid range")
}
if end >= size {
end = size - 1
}
return start, end, nil
}
func contentType(path string) string {
switch strings.ToLower(filepath.Ext(path)) {
case ".mkv":
return "video/x-matroska"
case ".mp4", ".m4v":
return "video/mp4"
case ".avi":
return "video/x-msvideo"
case ".ts", ".m2ts":
return "video/mp2t"
default:
return "application/octet-stream"
}
}
func hashToken(token string) string {
sum := sha256.Sum256([]byte(token))
return hex.EncodeToString(sum[:])
}
func clientIP(r *http.Request) string {
if xff := r.Header.Get("X-Forwarded-For"); xff != "" {
parts := strings.Split(xff, ",")
return strings.TrimSpace(parts[0])
}
if xri := r.Header.Get("X-Real-IP"); xri != "" {
return xri
}
host, _, err := net.SplitHostPort(r.RemoteAddr)
if err != nil {
return r.RemoteAddr
}
return host
}