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, _ := time.Parse(time.RFC3339, p.AbsoluteExpiresAt) _ = c.streamer.CreateSession(database.PlaybackSession{ SessionID: p.SessionID, TokenHash: p.TokenHash, LocalFileID: p.LocalFileID, IdleTimeoutSeconds: p.IdleTimeoutSeconds, AbsoluteExpiresAt: expires, LastActivity: time.Now().UTC(), }) log.Printf("Playback session created: %s", p.SessionID) } 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 }