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 fullSyncBatchSize = 200 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 } 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, }) 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"` } _ = 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) // Must run async so readLoop can receive FULL_LIBRARY_SYNC_ACK 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": 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"` } if json.Unmarshal(payload, &p) == nil { c.applyScannerSettings(p.FullScanInterval, p.FullScanAt, 0) } case "ERROR": log.Printf("Master error: %s", string(payload)) } } 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.mu.Unlock() c.flushLibraryEvents(events) } 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) 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) } }() 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) 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 }