Batch full library sync over WebSocket to avoid connection resets.

Large libraries (~30k files) overflowed a single WS message; sync now runs in ACK'd batches with event buffering on the node.

Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
Jos Vooges | STH 2026-08-25 01:58:31 +02:00
parent 7741500e73
commit 7519c9228f
5 changed files with 313 additions and 66 deletions

View file

@ -38,7 +38,12 @@ async function main() {
timeWindow: "1 minute", 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) => { app.addHook("onSend", async (request, reply, payload) => {
if (!request.url.startsWith("/stremio/")) { if (!request.url.startsWith("/stremio/")) {

View file

@ -1,61 +1,87 @@
import type { LibraryEvent } from "@media-cluster/shared-types";
import { prisma } from "../database/client"; import { prisma } from "../database/client";
import type { LibraryEvent } from "@media-cluster/shared-types";
import { MetadataService } from "../metadata/service"; import { MetadataService } from "../metadata/service";
const PRUNE_CHUNK = 500;
export class LibrarySyncService { export class LibrarySyncService {
constructor(private readonly metadata: MetadataService) {} constructor(private metadata: MetadataService) {}
async processEvents(nodeId: string, events: LibraryEvent[]): Promise<number> {
let maxRevision = 0;
async processEvents(nodeId: string, events: LibraryEvent[]): Promise<void> {
for (const event of events) { for (const event of events) {
maxRevision = Math.max(maxRevision, event.nodeRevision); switch (event.type) {
case "ADDED":
if (event.type === "FILE_REMOVED") { case "UPDATED":
await prisma.mediaFile.updateMany({ await this.metadata.matchMediaFile(nodeId, event.file);
where: { nodeId, localFileId: event.file.localFileId }, break;
data: { available: false }, case "REMOVED":
}); await prisma.mediaFile.updateMany({
continue; 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({ await prisma.node.update({
where: { id: nodeId }, 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( async processFullSync(
nodeId: string, nodeId: string,
nodeRevision: number, nodeRevision: number,
files: LibraryEvent["file"][] files: LibraryEvent["file"][],
options: {
isLast: boolean;
fileIdsSoFar: Set<string>;
}
): Promise<void> { ): Promise<void> {
const localFileIds = files.map((f) => f.localFileId);
await prisma.mediaFile.updateMany({
where: {
nodeId,
localFileId: { notIn: localFileIds },
},
data: { available: false },
});
for (const file of files) { for (const file of files) {
options.fileIdsSoFar.add(file.localFileId);
await this.metadata.matchMediaFile(nodeId, file); 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({ await prisma.node.update({
where: { id: nodeId }, where: { id: nodeId },
data: { data: {
lastRevision: nodeRevision, lastRevision: nodeRevision,
libraryFileCount: files.length, libraryFileCount: keepIds.size,
}, },
}); });
} }

View file

@ -27,12 +27,18 @@ interface PendingAck {
timer: ReturnType<typeof setTimeout>; timer: ReturnType<typeof setTimeout>;
} }
interface SyncState {
syncId: string;
fileIds: Set<string>;
}
class NodeConnectionManager { class NodeConnectionManager {
private connections = new Map<string, NodeConnection>(); private connections = new Map<string, NodeConnection>();
private librarySync: LibrarySyncService | null = null; private librarySync: LibrarySyncService | null = null;
private config: Config | null = null; private config: Config | null = null;
private offlineCheckInterval: ReturnType<typeof setInterval> | null = null; private offlineCheckInterval: ReturnType<typeof setInterval> | null = null;
private pendingSessionAcks = new Map<string, PendingAck>(); private pendingSessionAcks = new Map<string, PendingAck>();
private syncState = new Map<string, SyncState>();
init(config: Config): void { init(config: Config): void {
this.config = config; this.config = config;
@ -69,6 +75,7 @@ class NodeConnectionManager {
socket.on("close", () => { socket.on("close", () => {
if (nodeId) { if (nodeId) {
this.connections.delete(nodeId); this.connections.delete(nodeId);
this.syncState.delete(nodeId);
void prisma.node.update({ void prisma.node.update({
where: { id: nodeId }, where: { id: nodeId },
data: { status: "OFFLINE" }, data: { status: "OFFLINE" },
@ -147,6 +154,7 @@ class NodeConnectionManager {
return; return;
} }
this.syncState.delete(payload.nodeId);
this.connections.set(payload.nodeId, { this.connections.set(payload.nodeId, {
nodeId: payload.nodeId, nodeId: payload.nodeId,
socket, socket,
@ -210,6 +218,13 @@ class NodeConnectionManager {
const conn = [...this.connections.values()].find((c) => c.socket === socket); const conn = [...this.connections.values()].find((c) => c.socket === socket);
if (!conn || !this.librarySync) return; 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 } }); const node = await prisma.node.findUnique({ where: { id: conn.nodeId } });
if (!node) return; if (!node) return;
@ -232,12 +247,49 @@ class NodeConnectionManager {
const conn = [...this.connections.values()].find((c) => c.socket === socket); const conn = [...this.connections.values()].find((c) => c.socket === socket);
if (!conn || !this.librarySync) return; if (!conn || !this.librarySync) return;
await this.librarySync.processFullSync( const syncId = payload.syncId ?? `legacy-${Date.now()}`;
conn.nodeId, const isLast = payload.isLast ?? true;
payload.nodeRevision, const batchIndex = payload.batchIndex ?? 0;
payload.files
); let state = this.syncState.get(conn.nodeId);
this.send(socket, createMessage("FULL_LIBRARY_SYNC_ACK", { accepted: true })); 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 { private handlePlaybackSessionAck(payload: CreatePlaybackSessionAckPayload): void {
@ -309,6 +361,7 @@ class NodeConnectionManager {
if (conn) { if (conn) {
conn.socket.close(4004, "Revoked"); conn.socket.close(4004, "Revoked");
this.connections.delete(nodeId); this.connections.delete(nodeId);
this.syncState.delete(nodeId);
} }
} }
@ -359,6 +412,7 @@ class NodeConnectionManager {
conn.socket.close(1001, "Server shutting down"); conn.socket.close(1001, "Server shutting down");
} }
this.connections.clear(); this.connections.clear();
this.syncState.clear();
} }
} }

View file

@ -21,20 +21,34 @@ import (
"github.com/sthmedia/media-node/internal/streaming" "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 { type Client struct {
cfg *config.Config cfg *config.Config
creds *config.Credentials creds *config.Credentials
store *database.Store store *database.Store
scanner *scanner.Scanner scanner *scanner.Scanner
streamer *streaming.Server streamer *streaming.Server
version string version string
conn *websocket.Conn conn *websocket.Conn
mu sync.Mutex mu sync.Mutex
connected bool connected bool
startTime time.Time startTime time.Time
stopCh chan struct{} stopCh chan struct{}
syncAckCh chan syncAck
fullSyncActive bool
pendingEvents []scanner.LibraryEvent
} }
func NewClient( func NewClient(
@ -54,6 +68,7 @@ func NewClient(
version: version, version: version,
startTime: time.Now(), startTime: time.Now(),
stopCh: make(chan struct{}), stopCh: make(chan struct{}),
syncAckCh: make(chan syncAck, 16),
} }
} }
@ -92,11 +107,18 @@ func (c *Client) Stop() {
default: default:
close(c.stopCh) close(c.stopCh)
} }
c.forceClose()
}
func (c *Client) forceClose() {
c.mu.Lock() c.mu.Lock()
if c.conn != nil { conn := c.conn
_ = c.conn.Close() c.conn = nil
} c.connected = false
c.mu.Unlock() c.mu.Unlock()
if conn != nil {
_ = conn.Close()
}
} }
func (c *Client) connect() error { func (c *Client) connect() error {
@ -111,12 +133,25 @@ func (c *Client) connect() error {
c.mu.Lock() c.mu.Lock()
c.conn = conn c.conn = conn
c.connected = true c.connected = true
c.fullSyncActive = false
c.pendingEvents = nil
c.mu.Unlock() c.mu.Unlock()
// Drain stale ACKs from a previous connection
for {
select {
case <-c.syncAckCh:
default:
goto drained
}
}
drained:
defer func() { defer func() {
c.mu.Lock() c.mu.Lock()
c.connected = false c.connected = false
c.conn = nil c.conn = nil
c.fullSyncActive = false
c.mu.Unlock() c.mu.Unlock()
_ = conn.Close() _ = conn.Close()
}() }()
@ -159,6 +194,7 @@ func (c *Client) heartbeatLoop(done <-chan struct{}) {
case <-ticker.C: case <-ticker.C:
if err := c.sendHeartbeat(); err != nil { if err := c.sendHeartbeat(); err != nil {
log.Printf("heartbeat failed: %v", err) log.Printf("heartbeat failed: %v", err)
c.forceClose()
return return
} }
} }
@ -189,7 +225,13 @@ func (c *Client) sendHeartbeat() error {
func (c *Client) readLoop() { func (c *Client) readLoop() {
for { 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 { if err != nil {
log.Printf("WebSocket read error: %v", err) log.Printf("WebSocket read error: %v", err)
return return
@ -219,7 +261,20 @@ func (c *Client) handleMessage(msgType string, payload json.RawMessage) {
return return
} }
log.Println("Connected to Master") 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": case "CREATE_PLAYBACK_SESSION":
var p struct { var p struct {
SessionID string `json:"sessionId"` SessionID string `json:"sessionId"`
@ -273,7 +328,7 @@ func (c *Client) handleMessage(msgType string, payload json.RawMessage) {
case "RESCAN": case "RESCAN":
go c.scanner.FullScan() go c.scanner.FullScan()
case "FULL_LIBRARY_SYNC": case "FULL_LIBRARY_SYNC":
c.sendFullSync() go c.sendFullSync()
case "CONFIG_UPDATE": case "CONFIG_UPDATE":
log.Println("CONFIG_UPDATE received (apply on next restart)") log.Println("CONFIG_UPDATE received (apply on next restart)")
case "ERROR": case "ERROR":
@ -285,26 +340,117 @@ func (c *Client) SendLibraryEvents(events []scanner.LibraryEvent) {
if len(events) == 0 { if len(events) == 0 {
return 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{}{ msg := controlMessage("LIBRARY_EVENT", map[string]interface{}{
"events": events, "events": events,
}) })
if err := c.write(msg); err != nil { if err := c.write(msg); err != nil {
log.Printf("failed to send library events: %v", err) log.Printf("failed to send library events: %v", err)
c.forceClose()
} }
} }
func (c *Client) sendFullSync() { func (c *Client) sendFullSync() {
files := c.scanner.AllMediaFiles() c.mu.Lock()
rev, _ := c.store.MaxRevision() if c.fullSyncActive {
msg := controlMessage("FULL_LIBRARY_SYNC", map[string]interface{}{ c.mu.Unlock()
"nodeRevision": rev,
"files": files,
})
if err := c.write(msg); err != nil {
log.Printf("full library sync failed: %v", err)
return 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) { 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 { 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() c.mu.Lock()
defer c.mu.Unlock() defer c.mu.Unlock()
if c.conn == nil { if c.conn == nil {
@ -325,7 +475,7 @@ func (c *Client) write(msg interface{}) error {
if err != nil { if err != nil {
return err 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) return c.conn.WriteMessage(websocket.TextMessage, data)
} }

View file

@ -58,6 +58,18 @@ export interface LibraryEventPayload {
export interface FullLibrarySyncPayload { export interface FullLibrarySyncPayload {
nodeRevision: number; nodeRevision: number;
files: LibraryEvent["file"][]; 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 { export interface RescanPayload {