stremio/node/media-node/internal/control/client.go
Jos Vooges | STH 0f7b566626 Add Plex-like live streams with viewer labels and progress, plus admin UI redesign.
Co-authored-by: Cursor <cursoragent@cursor.com>
2026-08-29 01:24:55 +02:00

763 lines
19 KiB
Go

package control
import (
"context"
"encoding/json"
"fmt"
"log"
"math/rand"
"net"
"net/http"
"os"
"runtime"
"strings"
"sync"
"time"
"github.com/google/uuid"
"github.com/gorilla/websocket"
"github.com/sthmedia/media-node/internal/config"
"github.com/sthmedia/media-node/internal/database"
"github.com/sthmedia/media-node/internal/health"
"github.com/sthmedia/media-node/internal/scanner"
"github.com/sthmedia/media-node/internal/streaming"
)
const (
protocolVersion = 1
fullSyncBatchSize = 200
libraryEventBatchMax = 50
libraryEventFlushWait = 200 * time.Millisecond
writeDeadlineShort = 15 * time.Second
writeDeadlineLong = 60 * time.Second
syncAckTimeout = 120 * time.Second
)
type syncAck struct {
SyncID string
BatchIndex int
}
type Client struct {
cfg *config.Config
configPath string
creds *config.Credentials
store *database.Store
scanner *scanner.Scanner
streamer *streaming.Server
version string
conn *websocket.Conn
mu sync.Mutex
connected bool
startTime time.Time
stopCh chan struct{}
syncAckCh chan syncAck
fullSyncActive bool
pendingEvents []scanner.LibraryEvent
eventBatch []scanner.LibraryEvent
eventFlushTimer *time.Timer
}
func NewClient(
cfg *config.Config,
creds *config.Credentials,
store *database.Store,
sc *scanner.Scanner,
streamer *streaming.Server,
version string,
configPath string,
) *Client {
return &Client{
cfg: cfg,
configPath: configPath,
creds: creds,
store: store,
scanner: sc,
streamer: streamer,
version: version,
startTime: time.Now(),
stopCh: make(chan struct{}),
syncAckCh: make(chan syncAck, 16),
}
}
func (c *Client) Run() {
backoff := time.Second
maxBackoff := 60 * time.Second
for {
select {
case <-c.stopCh:
return
default:
}
if err := c.connect(); err != nil {
log.Printf("Master connection failed: %v", err)
}
jitter := time.Duration(rand.Intn(1000)) * time.Millisecond
select {
case <-c.stopCh:
return
case <-time.After(backoff + jitter):
}
backoff *= 2
if backoff > maxBackoff {
backoff = maxBackoff
}
}
}
func (c *Client) Stop() {
select {
case <-c.stopCh:
default:
close(c.stopCh)
}
c.forceClose()
}
func (c *Client) forceClose() {
c.mu.Lock()
conn := c.conn
c.conn = nil
c.connected = false
c.mu.Unlock()
if conn != nil {
_ = conn.Close()
}
}
func (c *Client) connect() error {
wsURL := toWebSocketURL(c.cfg.Master.URL) + "/api/v1/node/connect"
dialer := websocket.Dialer{HandshakeTimeout: 15 * time.Second}
conn, _, err := dialer.Dial(wsURL, http.Header{})
if err != nil {
return fmt.Errorf("dial: %w", err)
}
c.mu.Lock()
c.conn = conn
c.connected = true
c.fullSyncActive = false
c.pendingEvents = nil
c.mu.Unlock()
// Drain stale ACKs from a previous connection
for {
select {
case <-c.syncAckCh:
default:
goto drained
}
}
drained:
defer func() {
c.mu.Lock()
c.connected = false
c.conn = nil
c.fullSyncActive = false
c.mu.Unlock()
_ = conn.Close()
}()
if err := c.sendHello(); err != nil {
return err
}
done := make(chan struct{})
go func() {
c.readLoop()
close(done)
}()
go c.heartbeatLoop(done)
<-done
return fmt.Errorf("connection closed")
}
func (c *Client) sendHello() error {
hostname, _ := os.Hostname()
msg := controlMessage("HELLO", map[string]interface{}{
"nodeId": c.creds.NodeID,
"apiKey": c.creds.APIKey,
"softwareVersion": c.version,
"architecture": runtime.GOARCH,
"hostname": hostname,
"publicStreamUrl": c.cfg.Stream.PublicURL,
"moviesPaths": c.cfg.Media.Movies,
"seriesPaths": c.cfg.Media.Series,
})
return c.write(msg)
}
func (c *Client) heartbeatLoop(done <-chan struct{}) {
ticker := time.NewTicker(30 * time.Second)
defer ticker.Stop()
for {
select {
case <-done:
return
case <-ticker.C:
if err := c.sendHeartbeat(); err != nil {
log.Printf("heartbeat failed: %v", err)
c.forceClose()
return
}
}
}
}
func (c *Client) sendHeartbeat() error {
paths := append(append([]string{}, c.cfg.Media.Movies...), c.cfg.Media.Series...)
total, free := health.DiskUsage(paths)
fileCount, _ := c.store.FileCount()
msg := controlMessage("HEARTBEAT", map[string]interface{}{
"nodeId": c.creds.NodeID,
"softwareVersion": c.version,
"uptimeSeconds": int(time.Since(c.startTime).Seconds()),
"cpuUsagePercent": 0,
"memoryUsagePercent": health.MemoryUsagePercent(),
"storageTotalBytes": total,
"storageFreeBytes": free,
"activeStreams": c.streamer.ActiveStreams(),
"currentBytesPerSec": c.streamer.BytesPerSecond(),
"libraryFileCount": fileCount,
"scanStatus": c.scanner.Status(),
"errors": []string{},
})
return c.write(msg)
}
func (c *Client) readLoop() {
for {
c.mu.Lock()
conn := c.conn
c.mu.Unlock()
if conn == nil {
return
}
_, data, err := conn.ReadMessage()
if err != nil {
log.Printf("WebSocket read error: %v", err)
return
}
var msg struct {
ProtocolVersion int `json:"protocolVersion"`
Type string `json:"type"`
Payload json.RawMessage `json:"payload"`
}
if err := json.Unmarshal(data, &msg); err != nil {
continue
}
c.handleMessage(msg.Type, msg.Payload)
}
}
func (c *Client) handleMessage(msgType string, payload json.RawMessage) {
switch msgType {
case "HELLO_ACK":
var ack struct {
Accepted bool `json:"accepted"`
Reason string `json:"reason"`
FullScanInterval string `json:"fullScanInterval"`
FullScanAt string `json:"fullScanAt"`
FullScanIntervalSeconds int `json:"fullScanIntervalSeconds"`
MoviesPaths []string `json:"moviesPaths"`
SeriesPaths []string `json:"seriesPaths"`
ScanRoots []struct {
Path string `json:"path"`
ShelfID string `json:"shelfId"`
Kind string `json:"kind"`
} `json:"scanRoots"`
}
_ = json.Unmarshal(payload, &ack)
if !ack.Accepted {
log.Printf("Master rejected connection: %s", ack.Reason)
return
}
log.Println("Connected to Master")
c.applyScannerSettings(ack.FullScanInterval, ack.FullScanAt, ack.FullScanIntervalSeconds)
rootsChanged := false
if len(ack.ScanRoots) > 0 {
rootsChanged = c.applyScanRoots(ack.ScanRoots)
} else if len(ack.MoviesPaths) > 0 || len(ack.SeriesPaths) > 0 {
c.applyMediaPaths(ack.MoviesPaths, ack.SeriesPaths)
rootsChanged = true
}
// If roots changed, FullScan already started and will sync when done.
// Otherwise push current inventory (with shelfIds) to Master.
if !rootsChanged {
go c.sendFullSync()
}
case "FULL_LIBRARY_SYNC_ACK":
var p struct {
Accepted bool `json:"accepted"`
SyncID string `json:"syncId"`
BatchIndex int `json:"batchIndex"`
}
if json.Unmarshal(payload, &p) == nil && p.Accepted {
select {
case c.syncAckCh <- syncAck{SyncID: p.SyncID, BatchIndex: p.BatchIndex}:
default:
}
}
case "CREATE_PLAYBACK_SESSION":
var p struct {
SessionID string `json:"sessionId"`
TokenHash string `json:"tokenHash"`
LocalFileID string `json:"localFileId"`
IdleTimeoutSeconds int `json:"idleTimeoutSeconds"`
AbsoluteExpiresAt string `json:"absoluteExpiresAt"`
}
if json.Unmarshal(payload, &p) == nil {
expires, err := time.Parse(time.RFC3339Nano, p.AbsoluteExpiresAt)
if err != nil {
expires, err = time.Parse(time.RFC3339, p.AbsoluteExpiresAt)
}
ok := true
errMsg := ""
if err != nil {
ok = false
errMsg = "invalid absoluteExpiresAt"
log.Printf("Playback session rejected: bad expiry for %s", p.SessionID)
} else if _, ferr := c.store.GetFile(p.LocalFileID); ferr != nil {
ok = false
errMsg = "file not found on node"
log.Printf("Playback session rejected: missing file %s for %s", p.LocalFileID, p.SessionID)
} else if err := c.streamer.CreateSession(database.PlaybackSession{
SessionID: p.SessionID,
TokenHash: p.TokenHash,
LocalFileID: p.LocalFileID,
IdleTimeoutSeconds: p.IdleTimeoutSeconds,
AbsoluteExpiresAt: expires,
LastActivity: time.Now().UTC(),
}); err != nil {
ok = false
errMsg = err.Error()
log.Printf("Playback session create failed: %s (%v)", p.SessionID, err)
} else {
log.Printf("Playback session created: %s file=%s", p.SessionID, p.LocalFileID)
}
ack := controlMessage("CREATE_PLAYBACK_SESSION_ACK", map[string]interface{}{
"sessionId": p.SessionID,
"ok": ok,
"error": errMsg,
})
if werr := c.write(ack); werr != nil {
log.Printf("failed to send session ACK: %v", werr)
}
}
case "REVOKE_PLAYBACK_SESSION":
var p struct {
SessionID string `json:"sessionId"`
}
if json.Unmarshal(payload, &p) == nil {
_ = c.streamer.RevokeSession(p.SessionID)
log.Printf("Playback session revoked: %s", p.SessionID)
}
case "RESCAN":
log.Println("RESCAN requested by master")
go c.scanner.FullScan()
case "FULL_LIBRARY_SYNC":
go c.sendFullSync()
case "CONFIG_UPDATE":
var p struct {
FullScanInterval string `json:"fullScanInterval"`
FullScanAt string `json:"fullScanAt"`
MoviesPaths []string `json:"moviesPaths"`
SeriesPaths []string `json:"seriesPaths"`
ScanRoots []struct {
Path string `json:"path"`
ShelfID string `json:"shelfId"`
Kind string `json:"kind"`
} `json:"scanRoots"`
}
if json.Unmarshal(payload, &p) == nil {
c.applyScannerSettings(p.FullScanInterval, p.FullScanAt, 0)
if p.ScanRoots != nil {
c.applyScanRoots(p.ScanRoots)
} else {
c.applyMediaPaths(p.MoviesPaths, p.SeriesPaths)
}
}
case "RESTART":
go c.doRestart()
case "ERROR":
log.Printf("Master error: %s", string(payload))
}
}
func (c *Client) applyMediaPaths(movies, series []string) {
// nil means "not provided" — keep current. Empty slice means clear (rare).
if movies == nil && series == nil {
return
}
nextMovies := c.cfg.Media.Movies
nextSeries := c.cfg.Media.Series
if movies != nil {
nextMovies = normalizePaths(movies)
}
if series != nil {
nextSeries = normalizePaths(series)
}
if pathsEqual(nextMovies, c.cfg.Media.Movies) && pathsEqual(nextSeries, c.cfg.Media.Series) && len(c.cfg.Media.Roots) == 0 {
return
}
c.scanner.UpdateMediaPaths(nextMovies, nextSeries)
c.cfg.Media.Movies = nextMovies
c.cfg.Media.Series = nextSeries
c.cfg.Media.Roots = nil
if c.configPath != "" {
if err := config.Save(c.configPath, c.cfg); err != nil {
log.Printf("Failed to save media paths: %v", err)
}
}
}
func (c *Client) applyScanRoots(rows []struct {
Path string `json:"path"`
ShelfID string `json:"shelfId"`
Kind string `json:"kind"`
}) bool {
roots := make([]config.MediaRoot, 0, len(rows))
for _, r := range rows {
path := strings.TrimSpace(r.Path)
path = strings.TrimRight(path, "/")
if path == "" {
continue
}
kind := strings.ToLower(strings.TrimSpace(r.Kind))
if kind == "series" || kind == "episode" {
kind = "episode"
} else {
kind = "movie"
}
roots = append(roots, config.MediaRoot{Path: path, ShelfID: r.ShelfID, Kind: kind})
}
changed := c.scanner.UpdateMediaRoots(roots)
c.cfg.Media.Roots = roots
var movies, series []string
for _, r := range roots {
if r.Kind == "episode" {
series = append(series, r.Path)
} else {
movies = append(movies, r.Path)
}
}
c.cfg.Media.Movies = movies
c.cfg.Media.Series = series
if changed && c.configPath != "" {
if err := config.Save(c.configPath, c.cfg); err != nil {
log.Printf("Failed to save media roots: %v", err)
}
}
return changed
}
func normalizePaths(in []string) []string {
seen := map[string]bool{}
var out []string
for _, p := range in {
p = strings.TrimSpace(p)
p = strings.TrimRight(p, "/")
if p == "" || seen[p] {
continue
}
seen[p] = true
out = append(out, p)
}
return out
}
func pathsEqual(a, b []string) bool {
if len(a) != len(b) {
return false
}
for i := range a {
if a[i] != b[i] {
return false
}
}
return true
}
func (c *Client) doRestart() {
log.Println("RESTART requested by master")
name := strings.TrimSpace(os.Getenv("MEDIA_NODE_CONTAINER_NAME"))
if name == "" {
name = "media-node"
}
if _, err := os.Stat("/var/run/docker.sock"); err == nil {
if err := dockerRestartContainer(name); err != nil {
log.Printf("Docker restart via sock failed (%v); exiting for restart policy", err)
} else {
log.Printf("Docker restart requested for %s", name)
time.Sleep(2 * time.Second)
}
}
os.Exit(0)
}
func dockerRestartContainer(name string) error {
httpc := http.Client{
Transport: &http.Transport{
DialContext: func(_ context.Context, _, _ string) (net.Conn, error) {
return net.Dial("unix", "/var/run/docker.sock")
},
},
Timeout: 15 * time.Second,
}
url := "http://localhost/containers/" + name + "/restart?t=5"
req, err := http.NewRequest(http.MethodPost, url, nil)
if err != nil {
return err
}
res, err := httpc.Do(req)
if err != nil {
return err
}
defer res.Body.Close()
if res.StatusCode >= 300 {
return fmt.Errorf("docker API status %d", res.StatusCode)
}
return nil
}
func (c *Client) applyScannerSettings(interval, at string, legacySeconds int) {
if interval == "" && at == "" && legacySeconds > 0 {
interval = fmt.Sprintf("%ds", legacySeconds)
}
if interval == "" && at == "" {
return
}
if interval == "" {
interval = c.cfg.Scanner.FullScanInterval
}
if at == "" {
at = c.cfg.Scanner.FullScanAt
}
c.scanner.UpdateSchedule(interval, at)
c.cfg.Scanner.FullScanInterval = interval
c.cfg.Scanner.FullScanAt = at
if c.configPath != "" {
if err := config.Save(c.configPath, c.cfg); err != nil {
log.Printf("Failed to save scanner settings: %v", err)
}
}
}
func (c *Client) SendLibraryEvents(events []scanner.LibraryEvent) {
if len(events) == 0 {
return
}
c.mu.Lock()
if c.fullSyncActive {
c.pendingEvents = append(c.pendingEvents, events...)
c.mu.Unlock()
return
}
c.eventBatch = append(c.eventBatch, events...)
if len(c.eventBatch) >= libraryEventBatchMax {
batch := c.eventBatch
c.eventBatch = nil
if c.eventFlushTimer != nil {
c.eventFlushTimer.Stop()
c.eventFlushTimer = nil
}
c.mu.Unlock()
c.flushLibraryEvents(batch)
return
}
if c.eventFlushTimer == nil {
c.eventFlushTimer = time.AfterFunc(libraryEventFlushWait, func() {
c.mu.Lock()
batch := c.eventBatch
c.eventBatch = nil
c.eventFlushTimer = nil
c.mu.Unlock()
if len(batch) > 0 {
c.flushLibraryEvents(batch)
}
})
}
c.mu.Unlock()
}
func (c *Client) flushLibraryEvents(events []scanner.LibraryEvent) {
msg := controlMessage("LIBRARY_EVENT", map[string]interface{}{
"events": events,
})
if err := c.write(msg); err != nil {
log.Printf("failed to send library events: %v", err)
c.forceClose()
}
}
func (c *Client) RequestFullLibrarySync() {
go c.sendFullSync()
}
func (c *Client) sendFullSync() {
c.mu.Lock()
if c.fullSyncActive {
c.mu.Unlock()
return
}
c.fullSyncActive = true
c.mu.Unlock()
defer func() {
c.mu.Lock()
c.fullSyncActive = false
pending := c.pendingEvents
c.pendingEvents = nil
c.mu.Unlock()
if len(pending) > 0 {
c.flushLibraryEvents(pending)
}
}()
// Never sync a half-built inventory — Master would prune the rest as missing.
log.Println("Waiting for scanner to finish before full library sync...")
c.scanner.WaitUntilIdle()
files := c.scanner.AllMediaFiles()
rev, _ := c.store.MaxRevision()
syncID := uuid.New().String()
batchCount := (len(files) + fullSyncBatchSize - 1) / fullSyncBatchSize
if batchCount == 0 {
batchCount = 1
}
log.Printf("Starting full library sync (%d files, %d batches)", len(files), batchCount)
for i := 0; i < batchCount; i++ {
start := i * fullSyncBatchSize
end := start + fullSyncBatchSize
if end > len(files) {
end = len(files)
}
batch := []scanner.MediaFileInfo{}
if start < len(files) {
batch = files[start:end]
}
// Drain stale ACKs before waiting for this batch
for {
select {
case <-c.syncAckCh:
default:
goto drained
}
}
drained:
msg := controlMessage("FULL_LIBRARY_SYNC", map[string]interface{}{
"nodeRevision": rev,
"files": batch,
"syncId": syncID,
"batchIndex": i,
"batchCount": batchCount,
"isLast": i == batchCount-1,
})
if err := c.writeWithDeadline(msg, writeDeadlineLong); err != nil {
log.Printf("full library sync batch %d/%d failed: %v", i+1, batchCount, err)
c.forceClose()
return
}
deadline := time.After(syncAckTimeout)
acked := false
for !acked {
select {
case ack := <-c.syncAckCh:
if ack.SyncID == syncID && ack.BatchIndex == i {
acked = true
}
case <-deadline:
log.Printf("full library sync batch %d/%d ACK timeout", i+1, batchCount)
c.forceClose()
return
case <-c.stopCh:
return
}
}
if (i+1)%10 == 0 || i == batchCount-1 {
log.Printf("Full library sync progress: batch %d/%d", i+1, batchCount)
}
}
log.Printf("Full library sync complete (%d files)", len(files))
}
func (c *Client) NotifySessionEnded(sessionID, reason string) {
msg := controlMessage("PLAYBACK_SESSION_ENDED", map[string]interface{}{
"sessionId": sessionID,
"reason": reason,
})
_ = c.write(msg)
}
func (c *Client) NotifySessionProgress(sessionID string, bytesSent, lastByteOffset, fileSizeBytes int64) {
msg := controlMessage("PLAYBACK_SESSION_PROGRESS", map[string]interface{}{
"sessionId": sessionID,
"lastActivity": time.Now().UTC().Format(time.RFC3339Nano),
"bytesSent": bytesSent,
"lastByteOffset": lastByteOffset,
"fileSizeBytes": fileSizeBytes,
})
_ = c.write(msg)
}
func (c *Client) write(msg interface{}) error {
return c.writeWithDeadline(msg, writeDeadlineShort)
}
func (c *Client) writeWithDeadline(msg interface{}, deadline time.Duration) error {
c.mu.Lock()
defer c.mu.Unlock()
if c.conn == nil {
return fmt.Errorf("not connected")
}
data, err := json.Marshal(msg)
if err != nil {
return err
}
_ = c.conn.SetWriteDeadline(time.Now().Add(deadline))
return c.conn.WriteMessage(websocket.TextMessage, data)
}
func controlMessage(msgType string, payload interface{}) map[string]interface{} {
return map[string]interface{}{
"protocolVersion": protocolVersion,
"type": msgType,
"messageId": uuid.New().String(),
"timestamp": time.Now().UTC().Format(time.RFC3339),
"payload": payload,
}
}
func toWebSocketURL(url string) string {
url = strings.TrimRight(url, "/")
if strings.HasPrefix(url, "https://") {
return "wss://" + strings.TrimPrefix(url, "https://")
}
if strings.HasPrefix(url, "http://") {
return "ws://" + strings.TrimPrefix(url, "http://")
}
if strings.HasPrefix(url, "wss://") || strings.HasPrefix(url, "ws://") {
return url
}
return "wss://" + url
}