stremio/node/media-node/internal/control/client.go
Jos Vooges | STH 7741500e73 Fix Stremio playback: lazy sessions, remove notWebReady, node ACK.
Co-authored-by: Cursor <cursoragent@cursor.com>
2026-08-25 01:47:54 +02:00

354 lines
8.4 KiB
Go

package control
import (
"encoding/json"
"fmt"
"log"
"math/rand"
"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
type Client struct {
cfg *config.Config
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{}
}
func NewClient(
cfg *config.Config,
creds *config.Credentials,
store *database.Store,
sc *scanner.Scanner,
streamer *streaming.Server,
version string,
) *Client {
return &Client{
cfg: cfg,
creds: creds,
store: store,
scanner: sc,
streamer: streamer,
version: version,
startTime: time.Now(),
stopCh: make(chan struct{}),
}
}
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.mu.Lock()
if c.conn != nil {
_ = c.conn.Close()
}
c.mu.Unlock()
}
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.mu.Unlock()
defer func() {
c.mu.Lock()
c.connected = false
c.conn = nil
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,
})
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)
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 {
_, data, err := c.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"`
}
_ = json.Unmarshal(payload, &ack)
if !ack.Accepted {
log.Printf("Master rejected connection: %s", ack.Reason)
return
}
log.Println("Connected to Master")
c.sendFullSync()
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 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", p.SessionID)
}
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":
go c.scanner.FullScan()
case "FULL_LIBRARY_SYNC":
c.sendFullSync()
case "CONFIG_UPDATE":
log.Println("CONFIG_UPDATE received (apply on next restart)")
case "ERROR":
log.Printf("Master error: %s", string(payload))
}
}
func (c *Client) SendLibraryEvents(events []scanner.LibraryEvent) {
if len(events) == 0 {
return
}
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)
}
}
func (c *Client) sendFullSync() {
files := c.scanner.AllMediaFiles()
rev, _ := c.store.MaxRevision()
msg := controlMessage("FULL_LIBRARY_SYNC", map[string]interface{}{
"nodeRevision": rev,
"files": files,
})
if err := c.write(msg); err != nil {
log.Printf("full library sync failed: %v", err)
return
}
log.Printf("Full library sync sent (%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) write(msg interface{}) 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(15 * time.Second))
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
}