fix(node): stabilize deferred auto search

This commit is contained in:
Gatefixer
2026-08-08 13:21:59 +00:00
parent 3efc9187da
commit a66b715ccb
10 changed files with 285 additions and 107 deletions
-12
View File
@@ -16,18 +16,6 @@ A builder for LanceDB queries.
- `StandardQueryBase`<`NativeQuery`>
## Properties
### inner
```ts
protected inner: Query | Promise<Query>;
```
#### Inherited from
`StandardQueryBase.inner`
## Methods
### analyzePlan()
-8
View File
@@ -25,14 +25,6 @@ Common methods supported by all query types
- `AsyncIterable`&lt;`RecordBatch`&gt;
## Properties
### inner
```ts
protected inner: NativeQueryType | Promise<NativeQueryType>;
```
## Methods
### analyzePlan()
-12
View File
@@ -12,18 +12,6 @@ A query that returns a subset of the rows in the table.
- [`QueryBase`](QueryBase.md)&lt;`NativeTakeQuery`&gt;
## Properties
### inner
```ts
protected inner: TakeQuery | Promise<TakeQuery>;
```
#### Inherited from
[`QueryBase`](QueryBase.md).[`inner`](QueryBase.md#inner)
## Methods
### analyzePlan()
-12
View File
@@ -18,18 +18,6 @@ This builder can be reused to execute the query many times.
- `StandardQueryBase`&lt;`NativeVectorQuery`&gt;
## Properties
### inner
```ts
protected inner: VectorQuery | Promise<VectorQuery>;
```
#### Inherited from
`StandardQueryBase.inner`
## Methods
### addQueryVector()
+60 -3
View File
@@ -47,6 +47,7 @@ import {
BooleanQuery,
Occur,
Operator,
VectorQuery,
instanceOfFullTextQuery,
} from "../lancedb/query";
@@ -2337,9 +2338,23 @@ describe.each([arrow15, arrow16, arrow17, arrow18])(
expect(results2[0].text).toBe(data[1].text);
});
test("auto search follows embedding metadata across executions", async () => {
test("auto search stays consistent with the active revision", async () => {
let initCalls = 0;
let queryCalls = 0;
let markStarted!: () => void;
const started = new Promise<void>((resolve) => {
markStarted = resolve;
});
let releaseEmbedding!: () => void;
const embeddingReleased = new Promise<void>((resolve) => {
releaseEmbedding = resolve;
});
@register("refresh-test")
class TestEmbedding extends EmbeddingFunction<string> {
async init() {
initCalls += 1;
}
ndims() {
return 1;
}
@@ -2347,6 +2362,11 @@ describe.each([arrow15, arrow16, arrow17, arrow18])(
return new arrow.Float32();
}
async computeQueryEmbeddings(value: string) {
queryCalls += 1;
if (value === "blocked") {
markStarted();
await embeddingReleased;
}
return value === "greetings" ? [0.1] : [0.2];
}
async computeSourceEmbeddings(values: string[]) {
@@ -2362,7 +2382,10 @@ describe.each([arrow15, arrow16, arrow17, arrow18])(
readConsistencyInterval: 0,
});
const tracked = await reader.openTable("test");
const autoQuery = tracked.search("greetings").select(["text"]).limit(1);
const autoQuery = (tracked.search("greetings") as VectorQuery)
.nprobes(1)
.select(["text"])
.limit(1);
const func = new TestEmbedding();
const schema = LanceSchema({
@@ -2371,18 +2394,52 @@ describe.each([arrow15, arrow16, arrow17, arrow18])(
});
const data = [{ text: "hello world" }, { text: "goodbye world" }];
await writer.createTable("test", data, { mode: "overwrite", schema });
const baselineInitCalls = initCalls;
expect(
(await tracked.schema()).metadata.get("embedding_functions"),
).toBeDefined();
const results = await autoQuery.toArray();
expect(results[0].text).toBe(data[0].text);
expect(initCalls).toBe(baselineInitCalls + 1);
expect(queryCalls).toBe(1);
const ftsData = [{ text: "greetings from full text", vector: [0.0] }];
const repeatedResults = await autoQuery.toArray();
expect(repeatedResults[0].text).toBe(data[0].text);
expect(initCalls).toBe(baselineInitCalls + 1);
expect(queryCalls).toBe(1);
const multiVectorResults = await (
tracked.search("greetings") as VectorQuery
)
.addQueryVector(Promise.resolve([0.2]))
.select(["text"])
.limit(1)
.toArray();
expect(multiVectorResults).toHaveLength(2);
expect(multiVectorResults.map((row) => row.text).sort()).toEqual(
data.map((row) => row.text).sort(),
);
const pending = tracked
.search("blocked")
.select(["text"])
.limit(1)
.toArray();
await started;
const ftsData = [
{ text: "greetings from full text", vector: [0.0] },
{ text: "blocked from full text", vector: [0.0] },
];
const ftsTable = await writer.createTable("test", ftsData, {
mode: "overwrite",
});
await ftsTable.createIndex("text", { config: Index.fts() });
releaseEmbedding();
const pendingResults = await pending;
expect(pendingResults[0].text).toBe(ftsData[1].text);
expect(
(await tracked.schema()).metadata.get("embedding_functions"),
+176 -55
View File
@@ -103,7 +103,7 @@ export interface FullTextSearchOptions {
type NativeQueryLike = NativeQuery | NativeVectorQuery | NativeTakeQuery;
class DeferredNativeQuery<NativeQueryType extends NativeQueryLike> {
private readonly calls: Array<(inner: NativeQueryType) => void> = [];
protected readonly calls: Array<(inner: NativeQueryType) => void> = [];
constructor(private readonly factory: () => Promise<NativeQueryType>) {}
@@ -131,6 +131,18 @@ function nearestToNative(
return inner.nearestTo(Float32Array.from(vector as number[]));
}
function addQueryVectorToNative(
inner: NativeVectorQuery,
vector: Awaited<IntoVector>,
) {
const raw = Array.isArray(vector) ? null : extractVectorBuffer(vector);
if (raw) {
inner.addQueryVectorRaw(raw.data, raw.dtype);
} else {
inner.addQueryVector(Float32Array.from(vector as number[]));
}
}
/** Common methods supported by all query types
*
* @see {@link Query}
@@ -142,13 +154,24 @@ export class QueryBase<
NativeQueryType extends NativeQuery | NativeVectorQuery | NativeTakeQuery,
> implements AsyncIterable<RecordBatch>
{
/**
* @hidden
*/
protected inner:
| NativeQueryType
| Promise<NativeQueryType>
| DeferredNativeQuery<NativeQueryType>;
/**
* @hidden
*/
protected constructor(
protected inner: NativeQueryType | Promise<NativeQueryType>,
inner:
| NativeQueryType
| Promise<NativeQueryType>
| DeferredNativeQuery<NativeQueryType>,
) {
// intentionally empty
this.inner = inner;
}
// call a function on the inner (either a promise or the actual object)
@@ -156,10 +179,8 @@ export class QueryBase<
* @hidden
*/
protected doCall(fn: (inner: NativeQueryType) => void) {
if ((this.inner as unknown) instanceof DeferredNativeQuery) {
const deferred = this
.inner as unknown as DeferredNativeQuery<NativeQueryType>;
deferred.doCall(fn);
if (this.inner instanceof DeferredNativeQuery) {
this.inner.doCall(fn);
} else if (this.inner instanceof Promise) {
this.inner = this.inner.then((inner) => {
fn(inner);
@@ -174,10 +195,8 @@ export class QueryBase<
* @hidden
*/
protected resolveInner(): NativeQueryType | Promise<NativeQueryType> {
if ((this.inner as unknown) instanceof DeferredNativeQuery) {
const deferred = this
.inner as unknown as DeferredNativeQuery<NativeQueryType>;
return deferred.resolve();
if (this.inner instanceof DeferredNativeQuery) {
return this.inner.resolve();
}
return this.inner;
}
@@ -400,7 +419,12 @@ export class StandardQueryBase<
extends QueryBase<NativeQueryType>
implements ExecutableQuery
{
constructor(inner: NativeQueryType | Promise<NativeQueryType>) {
constructor(
inner:
| NativeQueryType
| Promise<NativeQueryType>
| DeferredNativeQuery<NativeQueryType>,
) {
super(inner);
}
@@ -550,10 +574,22 @@ export class VectorQuery extends StandardQueryBase<NativeVectorQuery> {
/**
* @hidden
*/
constructor(inner: NativeVectorQuery | Promise<NativeVectorQuery>) {
constructor(
inner:
| NativeVectorQuery
| Promise<NativeVectorQuery>
| DeferredNativeQuery<NativeVectorQuery>,
) {
super(inner);
}
/**
* @hidden
*/
protected doVectorCall(fn: (inner: NativeVectorQuery) => void) {
super.doCall(fn);
}
/**
* Set the number of partitions to search (probe)
*
@@ -581,7 +617,7 @@ export class VectorQuery extends StandardQueryBase<NativeVectorQuery> {
* the minimum and maximum to the same value.
*/
nprobes(nprobes: number): VectorQuery {
super.doCall((inner) => inner.nprobes(nprobes));
this.doVectorCall((inner) => inner.nprobes(nprobes));
return this;
}
@@ -595,7 +631,7 @@ export class VectorQuery extends StandardQueryBase<NativeVectorQuery> {
* but will also increase latency.
*/
minimumNprobes(minimumNprobes: number): VectorQuery {
super.doCall((inner) => inner.minimumNprobes(minimumNprobes));
this.doVectorCall((inner) => inner.minimumNprobes(minimumNprobes));
return this;
}
@@ -609,7 +645,7 @@ export class VectorQuery extends StandardQueryBase<NativeVectorQuery> {
* potential false negatives.
*/
maximumNprobes(maximumNprobes: number): VectorQuery {
super.doCall((inner) => inner.maximumNprobes(maximumNprobes));
this.doVectorCall((inner) => inner.maximumNprobes(maximumNprobes));
return this;
}
@@ -622,7 +658,7 @@ export class VectorQuery extends StandardQueryBase<NativeVectorQuery> {
* `undefined` means no lower or upper bound.
*/
distanceRange(lowerBound?: number, upperBound?: number): VectorQuery {
super.doCall((inner) => inner.distanceRange(lowerBound, upperBound));
this.doVectorCall((inner) => inner.distanceRange(lowerBound, upperBound));
return this;
}
@@ -636,7 +672,7 @@ export class VectorQuery extends StandardQueryBase<NativeVectorQuery> {
* also increase the latency of your query. The default value is 1.5*limit.
*/
ef(ef: number): VectorQuery {
super.doCall((inner) => inner.ef(ef));
this.doVectorCall((inner) => inner.ef(ef));
return this;
}
@@ -650,7 +686,7 @@ export class VectorQuery extends StandardQueryBase<NativeVectorQuery> {
* whose data type is a fixed-size-list of floats.
*/
column(column: string): VectorQuery {
super.doCall((inner) => inner.column(column));
this.doVectorCall((inner) => inner.column(column));
return this;
}
@@ -671,7 +707,7 @@ export class VectorQuery extends StandardQueryBase<NativeVectorQuery> {
distanceType(
distanceType: Required<IvfPqOptions>["distanceType"],
): VectorQuery {
super.doCall((inner) => inner.distanceType(distanceType));
this.doVectorCall((inner) => inner.distanceType(distanceType));
return this;
}
@@ -705,7 +741,7 @@ export class VectorQuery extends StandardQueryBase<NativeVectorQuery> {
* distance between the query vector and the actual uncompressed vector.
*/
refineFactor(refineFactor: number): VectorQuery {
super.doCall((inner) => inner.refineFactor(refineFactor));
this.doVectorCall((inner) => inner.refineFactor(refineFactor));
return this;
}
@@ -730,7 +766,7 @@ export class VectorQuery extends StandardQueryBase<NativeVectorQuery> {
* factor can often help restore some of the results lost by post filtering.
*/
postfilter(): VectorQuery {
super.doCall((inner) => inner.postfilter());
this.doVectorCall((inner) => inner.postfilter());
return this;
}
@@ -744,7 +780,7 @@ export class VectorQuery extends StandardQueryBase<NativeVectorQuery> {
* calculate your recall to select an appropriate value for nprobes.
*/
bypassVectorIndex(): VectorQuery {
super.doCall((inner) => inner.bypassVectorIndex());
this.doVectorCall((inner) => inner.bypassVectorIndex());
return this;
}
@@ -761,34 +797,19 @@ export class VectorQuery extends StandardQueryBase<NativeVectorQuery> {
addQueryVector(vector: IntoVector): VectorQuery {
if (vector instanceof Promise) {
const res = (async () => {
try {
const v = await vector;
// biome-ignore lint/suspicious/noExplicitAny: we need to get the `inner`, but js has no package scoping
const value: any = this.addQueryVector(v);
const inner = value.inner as
| NativeVectorQuery
| Promise<NativeVectorQuery>;
return inner;
} catch (e) {
return Promise.reject(e);
}
const inner = (await this.resolveInner()) as NativeVectorQuery;
addQueryVectorToNative(inner, await vector);
return inner;
})();
return new VectorQuery(res);
} else {
super.doCall((inner) => {
const raw = Array.isArray(vector) ? null : extractVectorBuffer(vector);
if (raw) {
inner.addQueryVectorRaw(raw.data, raw.dtype);
} else {
inner.addQueryVector(Float32Array.from(vector as number[]));
}
});
this.doVectorCall((inner) => addQueryVectorToNative(inner, vector));
return this;
}
}
rerank(reranker: Reranker): VectorQuery {
super.doCall((inner) =>
this.doVectorCall((inner) =>
inner.rerank(async (args) => {
const vecResults = await fromBufferToRecordBatch(args.vecResults);
const ftsResults = await fromBufferToRecordBatch(args.ftsResults);
@@ -807,11 +828,54 @@ export class VectorQuery extends StandardQueryBase<NativeVectorQuery> {
}
}
type AutoQueryResolution = {
inner: NativeQuery | NativeVectorQuery;
route: "fts" | "vector";
};
class DeferredAutoNativeQuery extends DeferredNativeQuery<NativeVectorQuery> {
private readonly vectorCalls: Array<
(inner: NativeVectorQuery) => void | Promise<void>
> = [];
constructor(
private readonly autoFactory: () => Promise<AutoQueryResolution>,
) {
super(async () => (await autoFactory()).inner as NativeVectorQuery);
}
doVectorCall(fn: (inner: NativeVectorQuery) => void | Promise<void>) {
this.vectorCalls.push(fn);
}
async resolve(): Promise<NativeVectorQuery> {
const resolution = await this.autoFactory();
for (const call of this.calls) {
call(resolution.inner as NativeVectorQuery);
}
if (resolution.route === "vector") {
for (const call of this.vectorCalls) {
await call(resolution.inner as NativeVectorQuery);
}
}
return resolution.inner as NativeVectorQuery;
}
}
class DeferredAutoQuery extends VectorQuery {
constructor(factory: () => Promise<NativeQuery | NativeVectorQuery>) {
super(
new DeferredNativeQuery(factory) as unknown as Promise<NativeVectorQuery>,
);
constructor(private readonly deferred: DeferredAutoNativeQuery) {
super(deferred);
}
protected doVectorCall(fn: (inner: NativeVectorQuery) => void) {
this.deferred.doVectorCall(fn);
}
addQueryVector(vector: IntoVector): VectorQuery {
this.deferred.doVectorCall(async (inner) => {
addQueryVectorToNative(inner, await vector);
});
return this;
}
}
@@ -825,17 +889,74 @@ export function createAutoQuery(
table: NativeTable,
query: string,
columns: string[] | null,
getVector: () => Promise<Awaited<IntoVector> | undefined>,
getVector: (metadata: string) => Promise<Awaited<IntoVector>>,
): VectorQuery {
return new DeferredAutoQuery(async () => {
const vector = await getVector();
const inner = table.query();
if (vector === undefined) {
inner.fullTextSearch({ query, columns });
return inner;
type RouteSnapshot = {
table: NativeTable;
embeddingMetadata: string | undefined;
};
type CachedPreparation = {
metadata: string;
vector: Promise<Awaited<IntoVector>>;
};
let cachedPreparation: CachedPreparation | undefined;
const snapshotRoute = async (): Promise<RouteSnapshot> => {
const snapshot = await table.querySnapshot();
const schema = tableFromIPC(await snapshot.schema()).schema;
return {
table: snapshot,
embeddingMetadata: schema.metadata.get("embedding_functions"),
};
};
const deferred = new DeferredAutoNativeQuery(async () => {
while (true) {
const initial = await snapshotRoute();
if (initial.embeddingMetadata === undefined) {
cachedPreparation = undefined;
const inner = initial.table.query();
inner.fullTextSearch({ query, columns });
return { inner, route: "fts" };
}
const metadata = initial.embeddingMetadata;
if (cachedPreparation?.metadata !== metadata) {
cachedPreparation = {
metadata,
vector: getVector(metadata),
};
}
const preparation = cachedPreparation;
let vector: Awaited<IntoVector>;
try {
vector = await preparation.vector;
} catch (error) {
if (cachedPreparation === preparation) {
cachedPreparation = undefined;
}
throw error;
}
// Provider preparation can perform arbitrary asynchronous work. Take a
// fresh pinned snapshot afterwards and only use the prepared vector if
// that snapshot has the same embedding configuration.
const current = await snapshotRoute();
if (current.embeddingMetadata !== metadata) {
cachedPreparation = undefined;
continue;
}
return {
inner: nearestToNative(current.table.query(), vector),
route: "vector",
};
}
return nearestToNative(inner, vector);
});
return new DeferredAutoQuery(deferred);
}
/**
+8 -5
View File
@@ -1044,15 +1044,18 @@ export class LocalTable extends Table {
const columns =
typeof ftsColumns === "string" ? [ftsColumns] : (ftsColumns ?? null);
return createAutoQuery(this.inner, query, columns, async () => {
const functions = await this.getEmbeddingFunctions();
return createAutoQuery(this.inner, query, columns, async (metadata) => {
const functions = await getRegistry().parseFunctions(
new Map([["embedding_functions", metadata]]),
);
// TODO: Support multiple embedding functions
const embeddingFunc: EmbeddingFunctionConfig | undefined = functions
.values()
.next().value;
if (!embeddingFunc) {
return undefined;
}
// The route only calls this callback when embedding metadata exists.
// parseFunctions either yields a provider or reports malformed metadata.
if (!embeddingFunc)
throw new Error("Invalid embedding function metadata");
return await embeddingFunc.function.computeQueryEmbeddings(query);
});
}
+7
View File
@@ -278,6 +278,13 @@ impl Table {
Ok(Query::new(self.inner_ref()?.query()))
}
/// Return a read-only table handle pinned to the current query revision.
#[napi(catch_unwind)]
pub async fn query_snapshot(&self) -> napi::Result<Table> {
let snapshot = self.inner_ref()?.query_snapshot().await.default_error()?;
Ok(Table::new(snapshot))
}
#[napi(catch_unwind)]
pub fn take_offsets(&self, offsets: Vec<i64>) -> napi::Result<TakeQuery> {
Ok(TakeQuery::new(
+8
View File
@@ -1678,6 +1678,14 @@ impl<S: HttpSend> BaseTable for RemoteTable<S> {
fn id(&self) -> &str {
&self.identifier
}
async fn query_snapshot(&self) -> Result<Arc<dyn BaseTable>> {
let description = self.describe().await?;
let schema: arrow_schema::Schema = description.schema.try_into()?;
let snapshot = self.with_branch(self.branch.clone());
*snapshot.version.write().await = Some(description.version);
snapshot.schema_cache.seed(Arc::new(schema));
Ok(Arc::new(snapshot))
}
async fn version(&self) -> Result<u64> {
self.describe().await.map(|desc| desc.version)
}
+26
View File
@@ -581,6 +581,13 @@ pub trait BaseTable: std::fmt::Display + std::fmt::Debug + Send + Sync {
fn id(&self) -> &str;
/// Get the arrow [Schema] of the table.
async fn schema(&self) -> Result<SchemaRef>;
/// Create a read-only handle pinned to the table's current active revision.
///
/// The returned handle is independent from later refreshes or checkouts on
/// this handle. This is used by bindings that must prepare client-side
/// query state from the same revision that the query will execute against.
#[doc(hidden)]
async fn query_snapshot(&self) -> Result<Arc<dyn BaseTable>>;
/// Count the number of rows in this table.
async fn count_rows(&self, filter: Option<Filter>) -> Result<usize>;
/// Create a physical plan for the query.
@@ -1068,6 +1075,16 @@ impl Table {
self.inner.schema().await
}
/// Create a read-only handle pinned to the current active revision.
#[doc(hidden)]
pub async fn query_snapshot(&self) -> Result<Self> {
Ok(Self {
inner: self.inner.query_snapshot().await?,
database: self.database.clone(),
embedding_registry: self.embedding_registry.clone(),
})
}
/// Count the number of rows in this dataset.
///
/// # Arguments
@@ -2835,6 +2852,15 @@ impl BaseTable for NativeTable {
&self.id
}
async fn query_snapshot(&self) -> Result<Arc<dyn BaseTable>> {
let dataset = self.dataset.get().await?;
let snapshot = dataset::DatasetConsistencyWrapper::new_time_travel(
(*dataset).clone(),
self.read_consistency_interval,
);
Ok(Arc::new(self.with_dataset(snapshot)))
}
async fn version(&self) -> Result<u64> {
Ok(self.dataset.get().await?.version().version)
}