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 } else if st.bytesSent == 0 && force { // Mark as live even before first chunk (Range open). st.bytesSent = 1 } if offset >= st.lastOffset { st.lastOffset = offset } if fileSize > 0 { st.fileSize = fileSize } shouldReport := 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 }