diff --git a/nodejs/__test__/embedding.test.ts b/nodejs/__test__/embedding.test.ts index e56e80631..0133905dc 100644 --- a/nodejs/__test__/embedding.test.ts +++ b/nodejs/__test__/embedding.test.ts @@ -184,6 +184,49 @@ describe("embedding functions", () => { const vector0 = JSON.parse(JSON.stringify(arr[0].vector)); expect(vector0).toEqual([1, 2, 3]); }); + it("should append using embedding metadata created by Python", async () => { + @register("python-mock") + // biome-ignore lint/correctness/noUnusedVariables: the decorator registers this class + class MockEmbeddingFunction extends EmbeddingFunction { + ndims() { + return 3; + } + embeddingDataType(): Float { + return new Float32(); + } + async computeQueryEmbeddings(_data: string) { + return [1, 2, 3]; + } + async computeSourceEmbeddings(data: string[]) { + return data.map(() => [1, 2, 3]); + } + } + + const metadata = new Map([ + [ + "embedding_functions", + '[{"source_column":"text","vector_column":"vector","name":"python-mock","model":{}}]', + ], + ]); + const schema = new Schema( + [ + new Field("text", new Utf8(), true), + new Field( + "vector", + new FixedSizeList(3, new Field("item", new Float32(), true)), + true, + ), + ], + metadata, + ); + + const db = await connect(tmpDir.name); + const table = await db.createEmptyTable("test", schema); + await table.add([{ text: "hello world" }]); + + const rows = await table.query().toArray(); + expect(JSON.parse(JSON.stringify(rows[0].vector))).toEqual([1, 2, 3]); + }); it("should error when appending to a table with an unregistered embedding function", async () => { @register("mock") class MockEmbeddingFunction extends EmbeddingFunction { diff --git a/nodejs/lancedb/arrow.ts b/nodejs/lancedb/arrow.ts index 587d30b19..3cb2068ab 100644 --- a/nodejs/lancedb/arrow.ts +++ b/nodejs/lancedb/arrow.ts @@ -48,6 +48,7 @@ import { } from "apache-arrow"; import { Buffers } from "apache-arrow/data"; import { type EmbeddingFunction } from "./embedding/embedding_function"; +import { parseEmbeddingFunctionMetadata } from "./embedding/metadata"; import { EmbeddingFunctionConfig, getRegistry } from "./embedding/registry"; import { sanitizeField, @@ -1385,11 +1386,13 @@ function validateSchemaEmbeddings( // Check schema metadata for embedding functions if (schema.metadata.has("embedding_functions")) { - const embeddings = JSON.parse( - schema.metadata.get("embedding_functions")!, - ); - // biome-ignore lint/suspicious/noExplicitAny: we don't know the type of `f` - if (embeddings.find((f: any) => f["vectorColumn"] === field.name)) { + const embeddings = parseEmbeddingFunctionMetadata(schema.metadata); + if ( + embeddings.find( + (embedding) => + (embedding.vectorColumn ?? "vector") === field.name, + ) + ) { hasEmbeddingFunction = true; } } diff --git a/nodejs/lancedb/embedding/metadata.ts b/nodejs/lancedb/embedding/metadata.ts new file mode 100644 index 000000000..f9dbac55e --- /dev/null +++ b/nodejs/lancedb/embedding/metadata.ts @@ -0,0 +1,50 @@ +// SPDX-License-Identifier: Apache-2.0 +// SPDX-FileCopyrightText: Copyright The LanceDB Authors + +import type { EmbeddingFunction } from "./embedding_function"; + +interface SerializedEmbeddingFunction { + name: string; + model: EmbeddingFunction["TOptions"]; + sourceColumn?: string; + vectorColumn?: string; + [key: string]: unknown; +} + +export interface EmbeddingFunctionMetadata { + name: string; + model: EmbeddingFunction["TOptions"]; + sourceColumn: string; + vectorColumn?: string; +} + +export function parseEmbeddingFunctionMetadata( + metadata: Map, +): EmbeddingFunctionMetadata[] { + const serialized = metadata.get("embedding_functions"); + if (serialized === undefined) { + return []; + } + + return (JSON.parse(serialized) as SerializedEmbeddingFunction[]).map( + (functionMetadata) => { + const sourceColumn = + functionMetadata.sourceColumn ?? + (functionMetadata["source_column"] as string | undefined); + if (sourceColumn === undefined) { + throw new Error( + "Embedding function metadata is missing a source column", + ); + } + + return { + name: functionMetadata.name, + model: functionMetadata.model, + sourceColumn, + vectorColumn: + functionMetadata.vectorColumn ?? + (functionMetadata["vector_column"] as string | undefined), + }; + }, + ); +} diff --git a/nodejs/lancedb/embedding/registry.ts b/nodejs/lancedb/embedding/registry.ts index 2eae90ed3..d17613805 100644 --- a/nodejs/lancedb/embedding/registry.ts +++ b/nodejs/lancedb/embedding/registry.ts @@ -5,6 +5,7 @@ import { type EmbeddingFunction, type EmbeddingFunctionConstructor, } from "./embedding_function"; +import { parseEmbeddingFunctionMetadata } from "./metadata"; import "reflect-metadata"; export type CreateReturnType = T extends { init: () => Promise } @@ -105,20 +106,10 @@ export class EmbeddingFunctionRegistry { this: EmbeddingFunctionRegistry, metadata: Map, ): Promise> { - if (!metadata.has("embedding_functions")) { + const functions = parseEmbeddingFunctionMetadata(metadata); + if (functions.length === 0) { return new Map(); } else { - type FunctionConfig = { - name: string; - sourceColumn: string; - vectorColumn: string; - model: EmbeddingFunction["TOptions"]; - }; - - const functions = ( - JSON.parse(metadata.get("embedding_functions")!) - ); - const items: [string, EmbeddingFunctionConfig][] = await Promise.all( functions.map(async (f) => { const fn = this.get(f.name);