diff --git a/src/client.ts b/src/client.ts index 0d3fabe..a60ec79 100644 --- a/src/client.ts +++ b/src/client.ts @@ -7,6 +7,7 @@ import { } from "./auth/login.js"; import { unwrapAuth } from "./auth/unwrap.js"; import { init, fromBase64, toBase64 } from "./crypto/index.js"; +import { fetchMLDataBatch, type MLData } from "./mldata-fetch.js"; import { decryptCollection, decryptFile } from "./model/index.js"; import { downloadFile as dlFile, @@ -273,6 +274,18 @@ export class Client { return files; } + // Fetch machine-learning data (face detections + CLIP embeddings) for up + // to a batch of files, each decrypted with its own key. One request; the + // library batches at `MLDATA_BATCH_SIZE` and schedules each batch through + // its metadata request pool. + async fetchMLData(args: { + fileIDs: number[]; + fileKeys: Map; + }): Promise> { + this.assertLoggedIn(); + return fetchMLDataBatch(this.api, args.fileIDs, args.fileKeys); + } + async downloadFile( file: EnteFile, outPath?: string, diff --git a/src/library/index.ts b/src/library/index.ts index 4d06e87..52e0a0e 100644 --- a/src/library/index.ts +++ b/src/library/index.ts @@ -23,6 +23,8 @@ import { join } from "node:path"; import envPaths from "env-paths"; import { MetadataStore } from "./store.js"; +import { MLDataStore } from "./mldata.js"; +import { RequestPools } from "./pools.js"; import { deriveRecords, snapshotFrom, @@ -32,6 +34,7 @@ import { type LibraryChange, } from "./records.js"; import type { CollectionsPage, FilesPage } from "../client.js"; +import { MLDATA_BATCH_SIZE, type MLData } from "../mldata-fetch.js"; import type { Collection, EnteFile } from "../model/types.js"; export const DEFAULT_REFRESH_INTERVAL_SECONDS = 3; @@ -47,14 +50,24 @@ export interface LibraryClient { collectionKey: Uint8Array; sinceTime: number; }): Promise; + // Fetch ML data (face detections + CLIP embeddings) for up to a batch of + // files. Optional: a client without it simply disables ML fetching, leaving + // the metadata refresh untouched. + fetchMLData?(args: { + fileIDs: number[]; + fileKeys: Map; + }): Promise>; } -// A single refresh cycle's progress. "started" fires before the network work, -// then exactly one of "done" or "failed"; "failed" carries the error message. +// A progress event for one unit of background work. A metadata "refresh" or an +// ML "fetchMLData" pass each fire "started" before their network work and then +// exactly one of "done" or "failed"; "failed" carries the error message and +// an ML "done" reports how many payloads it stored. export interface RefreshEvent { - operation: "refresh"; + operation: "refresh" | "fetchMLData"; status: "started" | "done" | "failed"; error?: string; + fetched?: number; } export type RefreshProgressCallback = (event: RefreshEvent) => void; @@ -69,6 +82,9 @@ export interface LibraryOptions { downloadDirectory?: string; refreshIntervalSeconds?: number; onProgress?: RefreshProgressCallback; + // The bounded request pools (issue #45). ML data is fetched through the + // metadata pool. Defaults to a fresh set at the design's caps. + pools?: RequestPools; } export interface LibraryStatus { @@ -81,6 +97,15 @@ export interface LibraryStatus { // The message from the most recent refresh, set only while that refresh // failed; cleared by the next success. lastError?: string; + // Wall-clock ms of the last ML fetch pass that succeeded, or undefined if + // none has yet (or ML fetching is disabled). + lastMLFetchAt?: number; + // The most recent ML fetch pass's error, set only while it failed. + lastMLError?: string; + // ML payloads stored on disk and CLIP embeddings in the index; undefined + // when ML fetching is disabled. + mlStored?: number; + mlIndexed?: number; closed: boolean; } @@ -93,12 +118,20 @@ export class Library { private readonly userID: number; private readonly intervalMs: number; private readonly onProgress?: RefreshProgressCallback; + private readonly pools: RequestPools; + // The ML-data cache, present only when the client can fetch ML data. + private readonly mldata?: MLDataStore; private timer?: ReturnType; private refreshing = false; + // Guards the ML fetch pass so a slow backfill never runs twice at once; a + // refresh whose pass is still running kicks nothing new. + private mlFetching = false; private closed = false; private lastRefreshAt?: number; private lastError?: string; + private lastMLFetchAt?: number; + private lastMLError?: string; // The plain-record projection as of the last refresh, and the GUI change // subscribers. A refresh that alters the projection notifies each with the // delta; `lastRecords` is kept current every refresh so a subscriber that @@ -118,6 +151,8 @@ export class Library { downloadDirectory?: string; intervalMs: number; onProgress?: RefreshProgressCallback; + pools: RequestPools; + mldata?: MLDataStore; }) { this.client = args.client; this.store = args.store; @@ -126,6 +161,8 @@ export class Library { this.downloadDirectory = args.downloadDirectory; this.intervalMs = args.intervalMs; this.onProgress = args.onProgress; + this.pools = args.pools; + this.mldata = args.mldata; this.lastRecords = this.deriveNow(); } @@ -148,6 +185,12 @@ export class Library { (opts.refreshIntervalSeconds ?? DEFAULT_REFRESH_INTERVAL_SECONDS) * 1000; + // The ML cache only earns its keep when the client can fetch ML data; + // a client without that capability opens no `mldata/` directory. + const mldata = opts.client.fetchMLData + ? await MLDataStore.open(join(cacheDirectory, "mldata")) + : undefined; + const lib = new Library({ client: opts.client, store, @@ -156,6 +199,8 @@ export class Library { downloadDirectory: opts.downloadDirectory, intervalMs, onProgress: opts.onProgress, + pools: opts.pools ?? new RequestPools(), + mldata, }); if (store.loadedFromDisk) { @@ -215,12 +260,17 @@ export class Library { for (const c of collections) { files += this.store.listFiles(c.id).length; } + const ml = this.mldata?.stats(); return { userID: this.store.userID, collections: collections.length, files, lastRefreshAt: this.lastRefreshAt, lastError: this.lastError, + lastMLFetchAt: this.lastMLFetchAt, + lastMLError: this.lastMLError, + mlStored: ml?.stored, + mlIndexed: ml?.indexed, closed: this.closed, }; } @@ -255,6 +305,11 @@ export class Library { this.lastRefreshAt = Date.now(); this.lastError = undefined; this.emit({ operation: "refresh", status: "done" }); + // Backfill ML data for the files this refresh knows about. It runs + // outside the refresh's success/failure so a fetch or disk problem + // there never marks the metadata refresh failed, and it is not + // awaited so it never stalls the refresh interval. + void this.runMLFetch(); } catch (err) { const error = err instanceof Error ? err.message : String(err); this.lastError = error; @@ -355,6 +410,85 @@ export class Library { } } + // One ML fetch pass: fetch, decrypt and store the ML data for every file + // the store knows about that is not cached (or whose `updationTime` has + // advanced), through the metadata pool, and update the CLIP index. Guarded + // so passes never overlap; a failure is reported, not thrown. + private async runMLFetch(): Promise { + const mldata = this.mldata; + // Bind so the call keeps the client as its receiver when invoked + // through the pool below. + const fetchMLData = this.client.fetchMLData?.bind(this.client); + if (!mldata || !fetchMLData || this.closed || this.mlFetching) return; + + const files = this.uniqueFiles(); + const needed = mldata.neededFor(files); + if (needed.length === 0) return; + + this.mlFetching = true; + this.emit({ operation: "fetchMLData", status: "started" }); + try { + const fileKeys = new Map(); + const updation = new Map(); + for (const f of files) { + fileKeys.set(f.id, f.key); + updation.set(f.id, f.updationTime); + } + + let stored = 0; + for (let i = 0; i < needed.length; i += MLDATA_BATCH_SIZE) { + if (this.closed) break; + const batch = needed.slice(i, i + MLDATA_BATCH_SIZE); + const payloads = await this.pools.metadata.run( + () => fetchMLData({ fileIDs: batch, fileKeys }), + { priority: "background" }, + ); + stored += (await mldata.storeFetched(payloads, updation)) + .stored; + } + + this.lastMLFetchAt = Date.now(); + this.lastMLError = undefined; + this.emit({ + operation: "fetchMLData", + status: "done", + fetched: stored, + }); + } catch (err) { + const error = err instanceof Error ? err.message : String(err); + this.lastMLError = error; + this.emit({ operation: "fetchMLData", status: "failed", error }); + } finally { + this.mlFetching = false; + } + } + + // The distinct files the store holds, one entry per fileID (a file in + // several collections shares its ML data), each carrying the key and the + // newest `updationTime` seen across its memberships. + private uniqueFiles(): { + id: number; + key: Uint8Array; + updationTime: number; + }[] { + const byID = new Map< + number, + { id: number; key: Uint8Array; updationTime: number } + >(); + for (const collection of this.store.listCollections()) { + for (const f of this.store.listFiles(collection.id)) { + const seen = byID.get(f.id); + if (seen === undefined || f.updationTime > seen.updationTime) + byID.set(f.id, { + id: f.id, + key: f.key, + updationTime: f.updationTime, + }); + } + } + return [...byID.values()]; + } + // Gather every file membership and project the store into by-id records. private deriveNow(): DerivedRecords { const collections = this.store.listCollections(); diff --git a/src/library/mldata.ts b/src/library/mldata.ts new file mode 100644 index 0000000..4c70129 --- /dev/null +++ b/src/library/mldata.ts @@ -0,0 +1,361 @@ +// The on-disk cache of Ente's per-file machine-learning data and the CLIP +// index derived from it (issue #49). +// +// Under `/mldata/` this keeps: +// +// - `.json` — one decrypted, gunzipped payload per file, written by +// rename. Its presence means it is complete: a torn write never leaves a +// half-file, so the set of these files is the source of truth for what is +// cached. The full payload (face boxes, landmarks, embeddings) is read back +// from here on demand and never held in RAM. +// +// - `clip.f32` + `clip.json` — the derived index the content search runs on. +// `clip.json` lists the indexed fileIDs in order plus the embedding length; +// `clip.f32` is those CLIP embeddings packed as one `Float32Array`, so the +// index loads in a single read with no per-vector parse. The index is +// rebuilt from the payloads whenever it is missing or structurally +// disagrees with the files present, and appended to as new payloads arrive. +// +// - `fetched.json` — a small map of fileID to the `updationTime` it was +// fetched at. This is best-effort bookkeeping for refetch decisions (a file +// whose `updationTime` later advances is refetched); the payloads, not this +// file, remain the record of what is cached, so losing it only forgoes +// update-driven refetch until the next fetch rewrites it. +// +// In RAM this holds only the id list and the packed `Float32Array`. + +import { mkdir, readFile, readdir } from "node:fs/promises"; +import { join } from "node:path"; + +import { writeAtomic } from "../download/index.js"; +import type { MLData } from "../mldata-fetch.js"; + +const CLIP_VECTORS = "clip.f32"; +const CLIP_INDEX = "clip.json"; +const FETCHED = "fetched.json"; +// A payload file is named for its fileID alone; the derived files above are +// not, so this pattern picks out payloads and nothing else. +const PAYLOAD_RE = /^(\d+)\.json$/; +const BYTES_PER_FLOAT = 4; + +// The on-disk form of `clip.json`. +interface ClipIndexFile { + fileIDs: number[]; + embeddingLength: number; +} + +// A file the model knows about, for deciding what to fetch. +export interface MLDataFile { + id: number; + updationTime: number; +} + +// The RAM index the search reads: `fileIDs[i]` owns the `embeddingLength` +// floats of `embeddings` starting at `i * embeddingLength`. +export interface MLIndex { + fileIDs: number[]; + embeddingLength: number; + embeddings: Float32Array; +} + +// Pull the CLIP embedding out of a payload, or undefined when it is absent or +// misshapen. Kept strict so a bad payload is skipped rather than corrupting the +// packed index. +const clipEmbedding = (payload: MLData): number[] | undefined => { + const clip = payload.clip; + if (typeof clip !== "object" || clip === null) return undefined; + const embedding = (clip as { embedding?: unknown }).embedding; + if (!Array.isArray(embedding)) return undefined; + if (embedding.some((v) => typeof v !== "number" || !Number.isFinite(v))) + return undefined; + return embedding as number[]; +}; + +export class MLDataStore { + readonly dir: string; + + // fileIDs whose payload JSON is present on disk (present means complete). + private readonly present = new Set(); + // fileID -> updationTime it was fetched at. + private readonly fetched = new Map(); + + // The packed index and where each id sits in it. + private ids: number[] = []; + private embeddingLength = 0; + private embeddings = new Float32Array(0); + private readonly pos = new Map(); + + private constructor(dir: string) { + this.dir = dir; + } + + // Open (creating the directory) and load the id list and packed index into + // RAM, rebuilding the index from the payloads when it is missing or does + // not match the files present. + static async open(dir: string): Promise { + const store = new MLDataStore(dir); + await mkdir(dir, { recursive: true }); + await store.loadPresent(); + await store.loadFetched(); + if (!(await store.tryLoadIndex())) await store.rebuildIndex(); + return store; + } + + // The fileIDs among `files` that must be fetched: every file with no + // payload yet (first run, then new files), plus any whose `updationTime` + // has advanced past the one its cached payload was fetched at. Returned + // sorted and unique. + neededFor(files: MLDataFile[]): number[] { + const latest = new Map(); + for (const f of files) { + const seen = latest.get(f.id); + if (seen === undefined || f.updationTime > seen) + latest.set(f.id, f.updationTime); + } + const needed: number[] = []; + for (const [id, updationTime] of latest) { + if (!this.present.has(id)) { + needed.push(id); + continue; + } + const at = this.fetched.get(id); + if (at !== undefined && updationTime > at) needed.push(id); + } + return needed.sort((a, b) => a - b); + } + + // Store a batch of fetched payloads: write one file per id, fold their CLIP + // embeddings into the packed index (in place for a refetch, appended for a + // new file), and persist the derived files. Returns how many payloads were + // stored and how many ids the index now holds. + async storeFetched( + payloads: Map, + updation: Map, + ): Promise<{ stored: number; indexed: number }> { + if (payloads.size === 0) return { stored: 0, indexed: this.ids.length }; + + for (const [id, payload] of payloads) { + await this.writePayload(id, payload); + this.present.add(id); + const at = updation.get(id); + if (at !== undefined) this.fetched.set(id, at); + } + + const updates: { at: number; vector: number[] }[] = []; + const appends: { id: number; vector: number[] }[] = []; + for (const [id, payload] of payloads) { + const vector = clipEmbedding(payload); + if (!vector) continue; + if (this.embeddingLength === 0 && this.ids.length === 0) + this.embeddingLength = vector.length; + // The index is fixed-width; a vector of another length (never seen + // from Ente's CLIP model) is stored but left out of the index. + if (vector.length !== this.embeddingLength) continue; + const at = this.pos.get(id); + if (at !== undefined) updates.push({ at, vector }); + else appends.push({ id, vector }); + } + + for (const { at, vector } of updates) + this.embeddings.set(vector, at * this.embeddingLength); + + if (appends.length > 0) { + const length = this.embeddingLength; + const grown = new Float32Array( + this.embeddings.length + appends.length * length, + ); + grown.set(this.embeddings); + let offset = this.embeddings.length; + for (const { id, vector } of appends) { + grown.set(vector, offset); + this.pos.set(id, this.ids.length); + this.ids.push(id); + offset += length; + } + this.embeddings = grown; + } + + await this.persistIndex(); + await this.persistFetched(); + return { stored: payloads.size, indexed: this.ids.length }; + } + + // The packed index the search runs on. The id list is copied so callers + // cannot disturb the store's own order; the embeddings are the live buffer. + getIndex(): MLIndex { + return { + fileIDs: [...this.ids], + embeddingLength: this.embeddingLength, + embeddings: this.embeddings, + }; + } + + // The full payload for a file, read from disk, or undefined when it is not + // cached or does not parse. + async readPayload(fileID: number): Promise { + if (!this.present.has(fileID)) return undefined; + let raw: string; + try { + raw = await readFile(this.payloadPath(fileID), "utf-8"); + } catch { + return undefined; + } + try { + return JSON.parse(raw) as MLData; + } catch { + return undefined; + } + } + + stats(): { stored: number; indexed: number } { + return { stored: this.present.size, indexed: this.ids.length }; + } + + private payloadPath(id: number): string { + return join(this.dir, `${id}.json`); + } + + private async writePayload(id: number, payload: MLData): Promise { + await writeAtomic( + this.payloadPath(id), + new TextEncoder().encode(JSON.stringify(payload)), + ); + } + + private async loadPresent(): Promise { + let names: string[]; + try { + names = await readdir(this.dir); + } catch { + return; + } + for (const name of names) { + const match = PAYLOAD_RE.exec(name); + if (match) this.present.add(Number(match[1])); + } + } + + private async loadFetched(): Promise { + let raw: string; + try { + raw = await readFile(join(this.dir, FETCHED), "utf-8"); + } catch { + return; + } + try { + const parsed = JSON.parse(raw) as Record; + for (const [key, value] of Object.entries(parsed)) { + const id = Number(key); + if ( + Number.isInteger(id) && + typeof value === "number" && + this.present.has(id) + ) + this.fetched.set(id, value); + } + } catch { + // Corrupt bookkeeping degrades refetch decisions, never fails open. + } + } + + // Load the packed index if it is present and consistent with the payloads: + // its ids must all still be present and its vector file must be exactly the + // size the id count and embedding length imply. Returns whether it loaded. + private async tryLoadIndex(): Promise { + let metaRaw: string; + try { + metaRaw = await readFile(join(this.dir, CLIP_INDEX), "utf-8"); + } catch { + return false; + } + let meta: ClipIndexFile; + try { + meta = JSON.parse(metaRaw) as ClipIndexFile; + } catch { + return false; + } + if ( + !Array.isArray(meta.fileIDs) || + typeof meta.embeddingLength !== "number" + ) + return false; + if (meta.fileIDs.some((id) => !this.present.has(id))) return false; + + let bytes: Buffer; + try { + bytes = await readFile(join(this.dir, CLIP_VECTORS)); + } catch { + return false; + } + const expected = + meta.fileIDs.length * meta.embeddingLength * BYTES_PER_FLOAT; + if (bytes.byteLength !== expected) return false; + + // One read, no parse: copy into an aligned buffer and view it as + // floats. The copy is needed because a Buffer from the pool can start + // at an offset a Float32Array cannot be laid over. + const aligned = new Uint8Array(bytes.byteLength); + aligned.set(bytes); + this.embeddings = new Float32Array(aligned.buffer); + this.embeddingLength = meta.embeddingLength; + this.ids = [...meta.fileIDs]; + this.pos.clear(); + this.ids.forEach((id, i) => this.pos.set(id, i)); + return true; + } + + // Rebuild the packed index by reading every payload present, then persist + // it. Payloads without a CLIP embedding (or of an unexpected length) are + // simply not indexed. + private async rebuildIndex(): Promise { + this.ids = []; + this.pos.clear(); + this.embeddingLength = 0; + const vectors: number[][] = []; + for (const id of [...this.present].sort((a, b) => a - b)) { + const payload = await this.readPayload(id); + if (!payload) continue; + const vector = clipEmbedding(payload); + if (!vector) continue; + if (this.embeddingLength === 0) + this.embeddingLength = vector.length; + if (vector.length !== this.embeddingLength) continue; + this.pos.set(id, this.ids.length); + this.ids.push(id); + vectors.push(vector); + } + const length = this.embeddingLength; + const packed = new Float32Array(this.ids.length * length); + vectors.forEach((vector, i) => packed.set(vector, i * length)); + this.embeddings = packed; + await this.persistIndex(); + } + + private async persistIndex(): Promise { + const meta: ClipIndexFile = { + fileIDs: this.ids, + embeddingLength: this.embeddingLength, + }; + await writeAtomic( + join(this.dir, CLIP_INDEX), + new TextEncoder().encode(JSON.stringify(meta)), + ); + await writeAtomic( + join(this.dir, CLIP_VECTORS), + new Uint8Array( + this.embeddings.buffer, + this.embeddings.byteOffset, + this.embeddings.byteLength, + ), + ); + } + + private async persistFetched(): Promise { + const record: Record = {}; + for (const [id, at] of this.fetched) record[id] = at; + await writeAtomic( + join(this.dir, FETCHED), + new TextEncoder().encode(JSON.stringify(record)), + ); + } +} diff --git a/src/metadata-backup.ts b/src/metadata-backup.ts index f227ddf..89bc1a9 100644 --- a/src/metadata-backup.ts +++ b/src/metadata-backup.ts @@ -1,4 +1,3 @@ -import { gunzipSync } from "node:zlib"; import { mkdirSync, mkdtempSync, @@ -11,7 +10,7 @@ import { tmpdir } from "node:os"; import * as jpeg from "jpeg-js"; import exifReader from "exif-reader"; import type { Client } from "./client.js"; -import { decryptBlob, fromBase64 } from "./crypto/index.js"; +import { fetchMLData } from "./mldata-fetch.js"; import type { EnteFile } from "./model/types.js"; export type ProgressCallback = (message: string) => void; @@ -24,50 +23,6 @@ export interface MetadataBackupOptions { const sanitizePath = (name: string): string => name.replace(/[/\\:*?"<>|]/g, "_").replace(/^\.+/, "_"); -interface RawRemoteFileData { - fileID: number; - encryptedData: string; - decryptionHeader: string; - updatedAt?: number; -} - -const fetchMLDataForFiles = async ( - client: Client, - fileIDs: number[], - fileKeys: Map, -): Promise>> => { - const api = client.getApiClient(); - const result = new Map>(); - const batchSize = 200; - - for (let i = 0; i < fileIDs.length; i += batchSize) { - const batch = fileIDs.slice(i, i + batchSize); - const { data } = await api.postJSON<{ data: RawRemoteFileData[] }>( - "/files/data/fetch", - { type: "mldata", fileIDs: batch }, - ); - - for (const entry of data ?? []) { - const key = fileKeys.get(entry.fileID); - if (!key) continue; - try { - const decrypted = decryptBlob( - fromBase64(entry.encryptedData), - fromBase64(entry.decryptionHeader), - key, - ); - const jsonStr = gunzipSync(Buffer.from(decrypted)).toString( - "utf-8", - ); - result.set(entry.fileID, JSON.parse(jsonStr)); - } catch { - // Corrupted ML data for this file; skip it - } - } - } - return result; -}; - // Extract the raw EXIF APP1 segment from JPEG bytes. Returns the EXIF // data buffer (starting after the APP1 length field, at the "Exif\0\0" // header) or undefined if no APP1 marker is found. @@ -228,8 +183,8 @@ export const runMetadataBackup = async ( } log("Fetching ML data (face detections, CLIP embeddings)..."); - const mlDataMap = await fetchMLDataForFiles( - client, + const mlDataMap = await fetchMLData( + client.getApiClient(), [...fileKeys.keys()], fileKeys, ); diff --git a/src/mldata-fetch.ts b/src/mldata-fetch.ts new file mode 100644 index 0000000..d1b58b0 --- /dev/null +++ b/src/mldata-fetch.ts @@ -0,0 +1,93 @@ +// Fetch and decrypt Ente's per-file machine-learning data ("magic" search +// data: face detections + CLIP embeddings). +// +// The data lives behind `/files/data/fetch` with `type: "mldata"`. Each entry +// comes back encrypted under the file's own key and gzipped; decrypting and +// gunzipping yields the JSON payload +// `{ face: { faces: [...] }, clip: { embedding } }`. Ente caps a request at 200 +// ids, so `fetchMLData` batches for callers that want many at once while +// `fetchMLDataBatch` is the single-request unit the library submits to its +// request pool. + +import { gunzipSync } from "node:zlib"; + +import type { ApiClient } from "./api/client.js"; +import { decryptBlob, fromBase64 } from "./crypto/index.js"; + +// The most ids one `/files/data/fetch` request may carry. +export const MLDATA_BATCH_SIZE = 200; + +// The decrypted, gunzipped per-file payload. Its concrete shape is Ente's; the +// store keeps the whole object verbatim and each consumer reads the fields it +// needs, so it stays an open record rather than a fixed interface. +export type MLData = Record; + +interface RawRemoteFileData { + fileID: number; + encryptedData: string; + decryptionHeader: string; + updatedAt?: number; +} + +// Decrypt one entry with its file key and gunzip the JSON payload. Returns +// undefined when the key is unknown or the entry does not decrypt/parse, so one +// corrupt file never fails a whole batch. +const decodeEntry = ( + entry: RawRemoteFileData, + key: Uint8Array | undefined, +): MLData | undefined => { + if (!key) return undefined; + try { + const decrypted = decryptBlob( + fromBase64(entry.encryptedData), + fromBase64(entry.decryptionHeader), + key, + ); + const json = gunzipSync(Buffer.from(decrypted)).toString("utf-8"); + return JSON.parse(json) as MLData; + } catch { + return undefined; + } +}; + +// Fetch ML data for up to `MLDATA_BATCH_SIZE` ids in a single request. This is +// the unit the request pools schedule; callers with more ids split them into +// batches and submit each batch to the pool. +export const fetchMLDataBatch = async ( + api: ApiClient, + fileIDs: number[], + fileKeys: Map, +): Promise> => { + const { data } = await api.postJSON<{ data: RawRemoteFileData[] }>( + "/files/data/fetch", + { type: "mldata", fileIDs }, + ); + const result = new Map(); + for (const entry of data ?? []) { + const payload = decodeEntry(entry, fileKeys.get(entry.fileID)); + if (payload) result.set(entry.fileID, payload); + } + return result; +}; + +// Fetch ML data for arbitrarily many ids, batching at `MLDATA_BATCH_SIZE`. Used +// by the one-shot metadata backup; the library fetches through its request pool +// with `fetchMLDataBatch` instead. +export const fetchMLData = async ( + api: ApiClient, + fileIDs: number[], + fileKeys: Map, +): Promise> => { + const result = new Map(); + for (let i = 0; i < fileIDs.length; i += MLDATA_BATCH_SIZE) { + const batch = fileIDs.slice(i, i + MLDATA_BATCH_SIZE); + for (const [id, payload] of await fetchMLDataBatch( + api, + batch, + fileKeys, + )) { + result.set(id, payload); + } + } + return result; +}; diff --git a/test/library/mldata.test.ts b/test/library/mldata.test.ts new file mode 100644 index 0000000..d95a641 --- /dev/null +++ b/test/library/mldata.test.ts @@ -0,0 +1,410 @@ +/** + * Tests for the ML-data cache and its derived CLIP index (issue #49). + * + * Two layers are exercised: + * + * 1. `MLDataStore` on its own: storing one payload file per fileID (present + * means complete), building a `clip.f32` + `clip.json` index that reloads + * in a single read, rebuilding that index from the payloads when it is + * missing or disagrees with the files present, appending as new payloads + * arrive, overwriting a refetched file in place, and deciding what to + * (re)fetch as `updationTime` advances. + * + * 2. `Library` wiring: after each refresh the library fetches ML data through + * the metadata pool for every known file not yet cached, is incremental on + * later refreshes, and refetches a file whose `updationTime` advanced. + * + * Embedding values are chosen to be exactly representable as float32 so the + * round-trip through `clip.f32` compares equal. + */ + +import { describe, it, expect, beforeEach, afterEach, vi } from "vitest"; +import { existsSync, mkdtempSync, rmSync, writeFileSync } from "node:fs"; +import { tmpdir } from "node:os"; +import { join } from "node:path"; + +import { MLDataStore } from "../../src/library/mldata.js"; +import { Library } from "../../src/library/index.js"; +import type { CollectionsPage, FilesPage } from "../../src/client.js"; +import type { MLData } from "../../src/mldata-fetch.js"; +import type { Collection, EnteFile } from "../../src/model/types.js"; + +// A payload shaped like Ente's: a CLIP embedding plus face data that only the +// on-disk payload carries (never the RAM index). +const payload = (embedding: number[]): MLData => ({ + face: { + faces: [{ faceID: "f", detection: { box: { x: 0.5 } } }], + }, + clip: { embedding }, +}); + +describe("MLDataStore", () => { + let dir: string; + + beforeEach(() => { + dir = mkdtempSync(join(tmpdir(), "quak-mldata-")); + }); + + afterEach(() => { + rmSync(dir, { recursive: true, force: true }); + }); + + it("stores one payload file per fileID and builds a one-read index", async () => { + const store = await MLDataStore.open(dir); + const res = await store.storeFetched( + new Map([ + [100, payload([0.5, 0.25, 0.75])], + [200, payload([1, -2, 0.5])], + ]), + new Map([ + [100, 10], + [200, 20], + ]), + ); + expect(res).toEqual({ stored: 2, indexed: 2 }); + + // One payload file per fileID, and the derived index files. + expect(existsSync(join(dir, "100.json"))).toBe(true); + expect(existsSync(join(dir, "200.json"))).toBe(true); + expect(existsSync(join(dir, "clip.f32"))).toBe(true); + expect(existsSync(join(dir, "clip.json"))).toBe(true); + + // Reopening loads the index from disk in one read. + const reopened = await MLDataStore.open(dir); + const index = reopened.getIndex(); + expect(index.fileIDs).toEqual([100, 200]); + expect(index.embeddingLength).toBe(3); + expect([...index.embeddings]).toEqual([0.5, 0.25, 0.75, 1, -2, 0.5]); + + // The full payload (face boxes) is read back from disk on demand. + const full = await reopened.readPayload(100); + expect(full?.face).toBeDefined(); + expect(await reopened.readPayload(999)).toBeUndefined(); + }); + + it("rebuilds the index from payloads when it is missing", async () => { + const store = await MLDataStore.open(dir); + await store.storeFetched( + new Map([[100, payload([0.5, 0.25, 0.75])]]), + new Map([[100, 10]]), + ); + + // The derived index is lost but the payloads survive. + rmSync(join(dir, "clip.f32")); + rmSync(join(dir, "clip.json")); + + const reopened = await MLDataStore.open(dir); + const index = reopened.getIndex(); + expect(index.fileIDs).toEqual([100]); + expect([...index.embeddings]).toEqual([0.5, 0.25, 0.75]); + expect(existsSync(join(dir, "clip.f32"))).toBe(true); + }); + + it("rebuilds the index when it disagrees with the files present", async () => { + const store = await MLDataStore.open(dir); + await store.storeFetched( + new Map([ + [100, payload([0.5, 0.25, 0.75])], + [200, payload([1, -2, 0.5])], + ]), + new Map([ + [100, 10], + [200, 20], + ]), + ); + + // A payload disappears out from under the index (leaving it referencing + // a file no longer present); the index must be rebuilt from what is + // actually on disk. + rmSync(join(dir, "200.json")); + + const reopened = await MLDataStore.open(dir); + expect(reopened.getIndex().fileIDs).toEqual([100]); + }); + + it("appends new payloads and overwrites a refetched file in place", async () => { + const store = await MLDataStore.open(dir); + await store.storeFetched( + new Map([[100, payload([0.5, 0.25, 0.75])]]), + new Map([[100, 10]]), + ); + // A later batch adds a new file: appended after the first. + await store.storeFetched( + new Map([[200, payload([1, -2, 0.5])]]), + new Map([[200, 20]]), + ); + // Refetching 100 (its embedding changed) updates it in place, not a + // duplicate row. + await store.storeFetched( + new Map([[100, payload([9, 9, 9])]]), + new Map([[100, 30]]), + ); + + const index = store.getIndex(); + expect(index.fileIDs).toEqual([100, 200]); + expect([...index.embeddings]).toEqual([9, 9, 9, 1, -2, 0.5]); + }); + + it("keeps a payload without a CLIP embedding out of the index", async () => { + const store = await MLDataStore.open(dir); + const res = await store.storeFetched( + new Map([[100, { face: { faces: [] } }]]), + new Map([[100, 10]]), + ); + expect(res.stored).toBe(1); + expect(res.indexed).toBe(0); + // The payload is still cached (present means complete). + expect(existsSync(join(dir, "100.json"))).toBe(true); + expect(store.getIndex().fileIDs).toEqual([]); + }); + + it("fetches only what is missing or has a newer updationTime", async () => { + const store = await MLDataStore.open(dir); + await store.storeFetched( + new Map([[100, payload([0.5, 0.25, 0.75])]]), + new Map([[100, 10]]), + ); + + // 100 is cached and current; 200 has never been fetched. + expect( + store.neededFor([ + { id: 100, updationTime: 10 }, + { id: 200, updationTime: 5 }, + ]), + ).toEqual([200]); + + // 100's updationTime advanced past what it was fetched at: refetch. + expect(store.neededFor([{ id: 100, updationTime: 15 }])).toEqual([100]); + + // Nothing advanced: nothing to fetch. + expect(store.neededFor([{ id: 100, updationTime: 10 }])).toEqual([]); + }); + + it("survives a corrupt index without losing the payloads", async () => { + const store = await MLDataStore.open(dir); + await store.storeFetched( + new Map([[100, payload([0.5, 0.25, 0.75])]]), + new Map([[100, 10]]), + ); + writeFileSync(join(dir, "clip.json"), "not json"); + + const reopened = await MLDataStore.open(dir); + expect(reopened.getIndex().fileIDs).toEqual([100]); + }); +}); + +// --- Library wiring --------------------------------------------------------- + +const USER_ID = 42; +const FAST_INTERVAL = 0.02; + +const collection = (id: number, updationTime: number): Collection => ({ + id, + ownerID: USER_ID, + key: new Uint8Array([id & 0xff]), + name: `album-${id}`, + type: "album", + updationTime, + isShared: false, +}); + +const file = ( + id: number, + collectionID: number, + updationTime: number, +): EnteFile => ({ + id, + collectionID, + ownerID: USER_ID, + key: new Uint8Array([id & 0xff]), + metadata: { + title: `file-${id}.jpg`, + fileType: "image", + creationTime: updationTime, + modificationTime: updationTime, + }, + file: { decryptionHeader: "aGVhZGVy" }, + thumbnail: { decryptionHeader: "dGh1bWI=" }, + updationTime, +}); + +// A mock client that serves scripted collection/file pages and per-file ML +// payloads, recording every ML fetch request so incremental behaviour is +// provable. +class MLMockClient { + userID = USER_ID; + collectionsQueue: CollectionsPage[] = []; + filesByCollection = new Map(); + mlByFile = new Map(); + mlFetchCalls: number[][] = []; + + whoami(): { email: string; userID: number } { + return { email: "user@example.com", userID: this.userID }; + } + + async collectionsSince(args: { + sinceTime: number; + }): Promise { + return ( + this.collectionsQueue.shift() ?? { + collections: [], + deleted: [], + cursor: args.sinceTime, + } + ); + } + + async filesSince(args: { + collectionID: number; + collectionKey: Uint8Array; + sinceTime: number; + }): Promise { + const queue = this.filesByCollection.get(args.collectionID); + return ( + queue?.shift() ?? { + files: [], + deleted: [], + cursor: args.sinceTime, + } + ); + } + + async fetchMLData(args: { + fileIDs: number[]; + fileKeys: Map; + }): Promise> { + this.mlFetchCalls.push([...args.fileIDs]); + const result = new Map(); + for (const id of args.fileIDs) { + const p = this.mlByFile.get(id); + if (p) result.set(id, p); + } + return result; + } + + filesFor(collectionID: number, ...pages: FilesPage[]): void { + this.filesByCollection.set(collectionID, pages); + } +} + +describe("Library ML-data fetch on refresh", () => { + let dir: string; + let cacheDirectory: string; + + beforeEach(() => { + dir = mkdtempSync(join(tmpdir(), "quak-lib-mldata-")); + cacheDirectory = join(dir, "cache"); + }); + + afterEach(() => { + rmSync(dir, { recursive: true, force: true }); + }); + + it("fetches, stores and indexes ML data for known files, then is incremental", async () => { + const client = new MLMockClient(); + client.collectionsQueue.push({ + collections: [collection(1, 100)], + deleted: [], + cursor: 100, + }); + client.filesFor(1, { + files: [file(1001, 1, 90), file(1002, 1, 95)], + deleted: [], + cursor: 95, + }); + client.mlByFile.set(1001, payload([0.5, 0.25, 0.75])); + client.mlByFile.set(1002, payload([1, -2, 0.5])); + + const lib = await Library.open({ + client, + cacheDirectory, + refreshIntervalSeconds: FAST_INTERVAL, + }); + try { + await vi.waitFor( + () => { + expect(lib.status().mlIndexed).toBe(2); + expect(lib.status().mlStored).toBe(2); + }, + { timeout: 2000, interval: 5 }, + ); + + // Both files were fetched, in one batch. + expect(client.mlFetchCalls.flat().sort((a, b) => a - b)).toEqual([ + 1001, 1002, + ]); + const callsAfterFirst = client.mlFetchCalls.length; + + // The index is on disk and reloads to the same shape. + const reopened = await MLDataStore.open( + join(cacheDirectory, "mldata"), + ); + expect(reopened.getIndex().fileIDs).toEqual([1001, 1002]); + + // Later refreshes with nothing new must not refetch. + await new Promise((r) => setTimeout(r, FAST_INTERVAL * 1000 * 5)); + expect(client.mlFetchCalls.length).toBe(callsAfterFirst); + } finally { + lib.close(); + } + }); + + it("refetches a file whose updationTime advanced", async () => { + const client = new MLMockClient(); + client.collectionsQueue.push({ + collections: [collection(1, 100)], + deleted: [], + cursor: 100, + }); + client.filesFor(1, { + files: [file(1001, 1, 90)], + deleted: [], + cursor: 90, + }); + client.mlByFile.set(1001, payload([0.5, 0.25, 0.75])); + + const lib = await Library.open({ + client, + cacheDirectory, + refreshIntervalSeconds: FAST_INTERVAL, + }); + try { + await vi.waitFor(() => expect(lib.status().mlIndexed).toBe(1), { + timeout: 2000, + interval: 5, + }); + const callsBefore = client.mlFetchCalls.length; + + // The file changes on the server (updationTime advances) with a new + // embedding; the next refresh must refetch it. + client.mlByFile.set(1001, payload([9, 9, 9])); + client.collectionsQueue.push({ + collections: [collection(1, 200)], + deleted: [], + cursor: 200, + }); + client.filesFor(1, { + files: [file(1001, 1, 190)], + deleted: [], + cursor: 190, + }); + + await vi.waitFor( + () => { + expect(client.mlFetchCalls.length).toBeGreaterThan( + callsBefore, + ); + expect(client.mlFetchCalls.flat()).toContain(1001); + }, + { timeout: 2000, interval: 5 }, + ); + + const reopened = await MLDataStore.open( + join(cacheDirectory, "mldata"), + ); + expect([...reopened.getIndex().embeddings]).toEqual([9, 9, 9]); + } finally { + lib.close(); + } + }); +});