diff --git a/apps/master-api/src/app.ts b/apps/master-api/src/app.ts index 8627f34..0b3849e 100644 --- a/apps/master-api/src/app.ts +++ b/apps/master-api/src/app.ts @@ -38,7 +38,12 @@ async function main() { timeWindow: "1 minute", }); - await app.register(websocket); + await app.register(websocket, { + options: { + // Large library sync batches (still chunked on the node side) + maxPayload: 16 * 1024 * 1024, + }, + }); app.addHook("onSend", async (request, reply, payload) => { if (!request.url.startsWith("/stremio/")) { diff --git a/apps/master-api/src/media/sync.ts b/apps/master-api/src/media/sync.ts index 0996848..b0311d9 100644 --- a/apps/master-api/src/media/sync.ts +++ b/apps/master-api/src/media/sync.ts @@ -1,61 +1,87 @@ -import type { LibraryEvent } from "@media-cluster/shared-types"; import { prisma } from "../database/client"; +import type { LibraryEvent } from "@media-cluster/shared-types"; import { MetadataService } from "../metadata/service"; +const PRUNE_CHUNK = 500; + export class LibrarySyncService { - constructor(private readonly metadata: MetadataService) {} - - async processEvents(nodeId: string, events: LibraryEvent[]): Promise { - let maxRevision = 0; + constructor(private metadata: MetadataService) {} + async processEvents(nodeId: string, events: LibraryEvent[]): Promise { for (const event of events) { - maxRevision = Math.max(maxRevision, event.nodeRevision); - - if (event.type === "FILE_REMOVED") { - await prisma.mediaFile.updateMany({ - where: { nodeId, localFileId: event.file.localFileId }, - data: { available: false }, - }); - continue; + switch (event.type) { + case "ADDED": + case "UPDATED": + await this.metadata.matchMediaFile(nodeId, event.file); + break; + case "REMOVED": + await prisma.mediaFile.updateMany({ + where: { nodeId, localFileId: event.file.localFileId }, + data: { available: false }, + }); + break; } - - await this.metadata.matchMediaFile(nodeId, event.file); } - if (maxRevision > 0) { + const lastEvent = events[events.length - 1]; + if (lastEvent) { + const count = await prisma.mediaFile.count({ + where: { nodeId, available: true }, + }); await prisma.node.update({ where: { id: nodeId }, - data: { lastRevision: maxRevision }, + data: { + lastRevision: lastEvent.nodeRevision, + libraryFileCount: count, + }, }); } - - return maxRevision; } + /** + * Process one FULL_LIBRARY_SYNC message (optionally a batch). + * Pruning of files missing from the node only runs when `isLast` is true. + */ async processFullSync( nodeId: string, nodeRevision: number, - files: LibraryEvent["file"][] + files: LibraryEvent["file"][], + options: { + isLast: boolean; + fileIdsSoFar: Set; + } ): Promise { - const localFileIds = files.map((f) => f.localFileId); - - await prisma.mediaFile.updateMany({ - where: { - nodeId, - localFileId: { notIn: localFileIds }, - }, - data: { available: false }, - }); - for (const file of files) { + options.fileIdsSoFar.add(file.localFileId); await this.metadata.matchMediaFile(nodeId, file); } + if (!options.isLast) { + return; + } + + const keepIds = options.fileIdsSoFar; + const existing = await prisma.mediaFile.findMany({ + where: { nodeId }, + select: { id: true, localFileId: true }, + }); + const toDisable = existing + .filter((row) => !keepIds.has(row.localFileId)) + .map((row) => row.id); + + for (let i = 0; i < toDisable.length; i += PRUNE_CHUNK) { + const chunk = toDisable.slice(i, i + PRUNE_CHUNK); + await prisma.mediaFile.updateMany({ + where: { id: { in: chunk } }, + data: { available: false }, + }); + } + await prisma.node.update({ where: { id: nodeId }, data: { lastRevision: nodeRevision, - libraryFileCount: files.length, + libraryFileCount: keepIds.size, }, }); } diff --git a/apps/master-api/src/websocket/manager.ts b/apps/master-api/src/websocket/manager.ts index bf304d1..e640b98 100644 --- a/apps/master-api/src/websocket/manager.ts +++ b/apps/master-api/src/websocket/manager.ts @@ -27,12 +27,18 @@ interface PendingAck { timer: ReturnType; } +interface SyncState { + syncId: string; + fileIds: Set; +} + class NodeConnectionManager { private connections = new Map(); private librarySync: LibrarySyncService | null = null; private config: Config | null = null; private offlineCheckInterval: ReturnType | null = null; private pendingSessionAcks = new Map(); + private syncState = new Map(); init(config: Config): void { this.config = config; @@ -69,6 +75,7 @@ class NodeConnectionManager { socket.on("close", () => { if (nodeId) { this.connections.delete(nodeId); + this.syncState.delete(nodeId); void prisma.node.update({ where: { id: nodeId }, data: { status: "OFFLINE" }, @@ -147,6 +154,7 @@ class NodeConnectionManager { return; } + this.syncState.delete(payload.nodeId); this.connections.set(payload.nodeId, { nodeId: payload.nodeId, socket, @@ -210,6 +218,13 @@ class NodeConnectionManager { const conn = [...this.connections.values()].find((c) => c.socket === socket); if (!conn || !this.librarySync) return; + // Don't request another full sync while one is already in progress + if (this.syncState.has(conn.nodeId)) { + await this.librarySync.processEvents(conn.nodeId, payload.events); + this.send(socket, createMessage("LIBRARY_EVENT_ACK", { accepted: true })); + return; + } + const node = await prisma.node.findUnique({ where: { id: conn.nodeId } }); if (!node) return; @@ -232,12 +247,49 @@ class NodeConnectionManager { const conn = [...this.connections.values()].find((c) => c.socket === socket); if (!conn || !this.librarySync) return; - await this.librarySync.processFullSync( - conn.nodeId, - payload.nodeRevision, - payload.files - ); - this.send(socket, createMessage("FULL_LIBRARY_SYNC_ACK", { accepted: true })); + const syncId = payload.syncId ?? `legacy-${Date.now()}`; + const isLast = payload.isLast ?? true; + const batchIndex = payload.batchIndex ?? 0; + + let state = this.syncState.get(conn.nodeId); + if (!state || state.syncId !== syncId) { + state = { syncId, fileIds: new Set() }; + this.syncState.set(conn.nodeId, state); + } + + try { + await this.librarySync.processFullSync( + conn.nodeId, + payload.nodeRevision, + payload.files, + { isLast, fileIdsSoFar: state.fileIds } + ); + + this.send( + socket, + createMessage("FULL_LIBRARY_SYNC_ACK", { + accepted: true, + syncId: payload.syncId, + batchIndex: payload.batchIndex, + }) + ); + + if (isLast) { + console.log( + `Full library sync complete for ${conn.nodeId}: ${state.fileIds.size} files` + + (payload.batchCount != null ? ` (${payload.batchCount} batches)` : "") + ); + this.syncState.delete(conn.nodeId); + } else if (batchIndex % 10 === 0) { + console.log( + `Full library sync progress for ${conn.nodeId}: batch ${(batchIndex ?? 0) + 1}/${payload.batchCount ?? "?"}` + ); + } + } catch (err) { + console.error(`Full sync failed for ${conn.nodeId}:`, err); + this.syncState.delete(conn.nodeId); + this.send(socket, createMessage("ERROR", { message: "Full sync failed" })); + } } private handlePlaybackSessionAck(payload: CreatePlaybackSessionAckPayload): void { @@ -309,6 +361,7 @@ class NodeConnectionManager { if (conn) { conn.socket.close(4004, "Revoked"); this.connections.delete(nodeId); + this.syncState.delete(nodeId); } } @@ -359,6 +412,7 @@ class NodeConnectionManager { conn.socket.close(1001, "Server shutting down"); } this.connections.clear(); + this.syncState.clear(); } } diff --git a/node/media-node/internal/control/client.go b/node/media-node/internal/control/client.go index 064f60d..3f8c1c8 100644 --- a/node/media-node/internal/control/client.go +++ b/node/media-node/internal/control/client.go @@ -21,20 +21,34 @@ import ( "github.com/sthmedia/media-node/internal/streaming" ) -const protocolVersion = 1 +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 - 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{} + 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{} + syncAckCh chan syncAck + fullSyncActive bool + pendingEvents []scanner.LibraryEvent } func NewClient( @@ -54,6 +68,7 @@ func NewClient( version: version, startTime: time.Now(), stopCh: make(chan struct{}), + syncAckCh: make(chan syncAck, 16), } } @@ -92,11 +107,18 @@ func (c *Client) Stop() { default: close(c.stopCh) } + c.forceClose() +} + +func (c *Client) forceClose() { c.mu.Lock() - if c.conn != nil { - _ = c.conn.Close() - } + conn := c.conn + c.conn = nil + c.connected = false c.mu.Unlock() + if conn != nil { + _ = conn.Close() + } } func (c *Client) connect() error { @@ -111,12 +133,25 @@ func (c *Client) connect() error { 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() }() @@ -159,6 +194,7 @@ func (c *Client) heartbeatLoop(done <-chan struct{}) { case <-ticker.C: if err := c.sendHeartbeat(); err != nil { log.Printf("heartbeat failed: %v", err) + c.forceClose() return } } @@ -189,7 +225,13 @@ func (c *Client) sendHeartbeat() error { func (c *Client) readLoop() { for { - _, data, err := c.conn.ReadMessage() + 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 @@ -219,7 +261,20 @@ func (c *Client) handleMessage(msgType string, payload json.RawMessage) { return } log.Println("Connected to Master") - c.sendFullSync() + // 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"` @@ -273,7 +328,7 @@ func (c *Client) handleMessage(msgType string, payload json.RawMessage) { case "RESCAN": go c.scanner.FullScan() case "FULL_LIBRARY_SYNC": - c.sendFullSync() + go c.sendFullSync() case "CONFIG_UPDATE": log.Println("CONFIG_UPDATE received (apply on next restart)") case "ERROR": @@ -285,26 +340,117 @@ 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() { - 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) + c.mu.Lock() + if c.fullSyncActive { + c.mu.Unlock() return } - log.Printf("Full library sync sent (%d files)", len(files)) + 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) { @@ -316,6 +462,10 @@ func (c *Client) NotifySessionEnded(sessionID, reason string) { } 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 { @@ -325,7 +475,7 @@ func (c *Client) write(msg interface{}) error { if err != nil { return err } - _ = c.conn.SetWriteDeadline(time.Now().Add(15 * time.Second)) + _ = c.conn.SetWriteDeadline(time.Now().Add(deadline)) return c.conn.WriteMessage(websocket.TextMessage, data) } diff --git a/packages/protocol/src/index.ts b/packages/protocol/src/index.ts index 3772909..4c538c9 100644 --- a/packages/protocol/src/index.ts +++ b/packages/protocol/src/index.ts @@ -58,6 +58,18 @@ export interface LibraryEventPayload { export interface FullLibrarySyncPayload { nodeRevision: number; files: LibraryEvent["file"][]; + /** Present when sync is split across multiple WebSocket messages */ + syncId?: string; + batchIndex?: number; + batchCount?: number; + /** True on the final batch (or when sending a single unbatched sync) */ + isLast?: boolean; +} + +export interface FullLibrarySyncAckPayload { + accepted: boolean; + syncId?: string; + batchIndex?: number; } export interface RescanPayload {