From 51b45c2ad5ccc86f5777b8d8f1d7bc366aa68abb Mon Sep 17 00:00:00 2001 From: Peter Steinberger Date: Sun, 13 Sep 2026 22:04:51 -0700 Subject: [PATCH] improve(memory): bound vector recall decision counts (#147900) Co-authored-by: Peter Steinberger <58493+steipete@users.noreply.github.com> --- .../src/memory/manager-search-knn.test.ts | 80 +++++++++++++++++++ .../src/memory/manager-search-knn.ts | 11 ++- 2 files changed, 87 insertions(+), 4 deletions(-) create mode 100644 extensions/memory-core/src/memory/manager-search-knn.test.ts diff --git a/extensions/memory-core/src/memory/manager-search-knn.test.ts b/extensions/memory-core/src/memory/manager-search-knn.test.ts new file mode 100644 index 000000000000..0b3df57738e3 --- /dev/null +++ b/extensions/memory-core/src/memory/manager-search-knn.test.ts @@ -0,0 +1,80 @@ +import { + ensureSqliteLibrarySelected, + openNodeSqliteDatabase, +} from "openclaw/plugin-sdk/memory-core-host-engine-knn"; +import { + ensureMemoryIndexSchema, + loadSqliteVecExtension, +} from "openclaw/plugin-sdk/memory-core-host-engine-storage"; +import { describe, expect, it } from "vitest"; +import { runVectorKnnQuery } from "./manager-search-knn.js"; +import { vectorToBlob } from "./vector-blob.js"; + +describe("memory vector KNN decision counts", () => { + it("stops counting after fallback is decided without enumerating the remaining rows", async () => { + ensureSqliteLibrarySelected(); + const db = openNodeSqliteDatabase(":memory:", { allowExtension: true }); + try { + const loaded = await loadSqliteVecExtension({ db }); + expect(loaded.ok, loaded.error).toBe(true); + ensureMemoryIndexSchema({ db, cacheEnabled: false, ftsEnabled: false }); + db.exec(`CREATE VIRTUAL TABLE memory_index_chunks_vec USING vec0( + id TEXT PRIMARY KEY, embedding FLOAT[2] + )`); + const insertChunk = db.prepare( + `INSERT INTO memory_index_chunks + (id, path, source, start_line, end_line, hash, model, text, embedding, updated_at) + VALUES (?, ?, 'memory', 1, 1, ?, ?, ?, ?, 1)`, + ); + const insertVector = db.prepare( + "INSERT INTO memory_index_chunks_vec (id, embedding) VALUES (?, ?)", + ); + db.exec("BEGIN"); + for (let index = 0; index < 6000; index += 1) { + const id = `chunk-${index}`; + const model = index < 1000 ? "target" : "other"; + const vector = index < 1000 ? [0, 1] : [1, 0]; + insertChunk.run(id, `memory/${id}.md`, id, model, `text ${id}`, JSON.stringify(vector)); + insertVector.run(id, vectorToBlob(vector)); + } + db.exec("COMMIT"); + let chunkVisits = 0; + let vectorVisits = 0; + db.function("visit_chunk", { varargs: true }, () => { + chunkVisits += 1; + return 1; + }); + db.function("visit_vector", { varargs: true }, () => { + vectorVisits += 1; + return 1; + }); + // Transparent views observe native row visits, including KNN, without + // changing stored rows or replacing the real sqlite-vec query engine. + db.exec(` + CREATE TEMP VIEW memory_index_chunks AS + SELECT * FROM main.memory_index_chunks WHERE visit_chunk(id); + CREATE TEMP VIEW observed_vectors AS + SELECT id, embedding, distance, k FROM main.memory_index_chunks_vec + WHERE visit_vector(id); + `); + expect( + runVectorKnnQuery(db, { + vectorTable: "observed_vectors", + providerModels: ["target"], + queryVec: [1, 0], + limit: 2, + snippetMaxChars: 20, + sourceFilter: { sql: "", params: [] }, + }), + ).toEqual({ rows: [], fallbackScanRequired: true }); + // Two KNN attempts visit at most 16 + 4,096 rows. Decision counts need + // only two matching chunks and 4,097 vectors, regardless of the tail. + expect(chunkVisits).toBeGreaterThan(0); + expect(vectorVisits).toBeGreaterThan(0); + expect(chunkVisits).toBeLessThanOrEqual(4114); + expect(vectorVisits).toBeLessThanOrEqual(8209); + } finally { + db.close(); + } + }); +}); diff --git a/extensions/memory-core/src/memory/manager-search-knn.ts b/extensions/memory-core/src/memory/manager-search-knn.ts index 21e9027d1fca..967a1131a283 100644 --- a/extensions/memory-core/src/memory/manager-search-knn.ts +++ b/extensions/memory-core/src/memory/manager-search-knn.ts @@ -158,16 +158,19 @@ export function runVectorKnnQuery( const candidateLimit = Math.min(request.limit * VECTOR_KNN_OVERSAMPLE_FACTOR, MAX_VECTOR_KNN_K); let rows = runVectorQuery(candidateLimit); if (rows.length < request.limit) { + // Only the widening/fallback thresholds matter; stop counting once they are known. const matchingChunkCountRow = db .prepare( - `SELECT COUNT(*) AS count FROM memory_index_chunks c WHERE ${vectorModelFilter}${request.sourceFilter.sql}`, + `SELECT COUNT(*) AS count FROM (\n` + + ` SELECT 1 FROM memory_index_chunks c WHERE ${vectorModelFilter}${request.sourceFilter.sql} LIMIT ?\n` + + `)`, ) - .get(...request.providerModels, ...request.sourceFilter.params); + .get(...request.providerModels, ...request.sourceFilter.params, request.limit); const matchingChunkCount = readCount(matchingChunkCountRow); if (matchingChunkCount > rows.length) { const vectorCountRow = db - .prepare(`SELECT COUNT(*) AS count FROM ${request.vectorTable}`) - .get(); + .prepare(`SELECT COUNT(*) AS count FROM (SELECT 1 FROM ${request.vectorTable} LIMIT ?)`) + .get(MAX_VECTOR_KNN_K + 1); const vectorCount = readCount(vectorCountRow); const widenedLimit = Math.min(vectorCount, MAX_VECTOR_KNN_K); if (widenedLimit > candidateLimit) {