354 lines
8.4 KiB
Go
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
|
|
}
|