406 lines
9.6 KiB
Go
406 lines
9.6 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)
|
|
}
|
|
|
|
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),
|
|
}
|
|
go s.cleanupLoop()
|
|
go s.bandwidthLoop()
|
|
return s
|
|
}
|
|
|
|
func (s *Server) SetSessionEndHandler(fn func(sessionID, reason string)) {
|
|
s.onSessionEnd = 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)
|
|
http.NotFound(w, r)
|
|
return
|
|
}
|
|
idleDeadline := sess.LastActivity.Add(time.Duration(sess.IdleTimeoutSeconds) * time.Second)
|
|
if now.After(idleDeadline) {
|
|
_ = s.store.DeleteSession(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.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)
|
|
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.streamRange(w, f, start, end)
|
|
}
|
|
|
|
func (s *Server) streamRange(w http.ResponseWriter, f *os.File, start, end 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)
|
|
|
|
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)
|
|
}
|
|
if err != nil {
|
|
return
|
|
}
|
|
}
|
|
}
|
|
|
|
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 {
|
|
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)
|
|
continue
|
|
}
|
|
idleDeadline := sess.LastActivity.Add(time.Duration(sess.IdleTimeoutSeconds) * time.Second)
|
|
if now.After(idleDeadline) {
|
|
_ = s.store.DeleteSession(sess.SessionID)
|
|
if s.onSessionEnd != nil {
|
|
s.onSessionEnd(sess.SessionID, "idle_timeout")
|
|
}
|
|
}
|
|
}
|
|
}
|
|
}
|
|
|
|
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
|
|
}
|