refactor(ts-sdk): lazy-load optional provider SDKs so importing mem0ai/oss never requires them (#6280)

This commit is contained in:
Kartik
2026-07-13 15:23:12 +05:30
committed by GitHub
parent 17836748d7
commit d89793b666
18 changed files with 726 additions and 312 deletions
+48
View File
@@ -168,6 +168,54 @@
},
"@aws-sdk/client-neptune-graph": {
"optional": true
},
"chromadb": {
"optional": true
},
"mongodb": {
"optional": true
},
"weaviate-client": {
"optional": true
},
"cassandra-driver": {
"optional": true
},
"@pinecone-database/pinecone": {
"optional": true
},
"@aws-sdk/client-s3vectors": {
"optional": true
},
"@turbopuffer/turbopuffer": {
"optional": true
},
"@upstash/vector": {
"optional": true
},
"@elastic/elasticsearch": {
"optional": true
},
"@opensearch-project/opensearch": {
"optional": true
},
"cohere-ai": {
"optional": true
},
"zeroentropy": {
"optional": true
},
"fastembed": {
"optional": true
},
"@google-cloud/aiplatform": {
"optional": true
},
"@huggingface/transformers": {
"optional": true
},
"iovalkey": {
"optional": true
}
},
"engines": {
+37 -14
View File
@@ -1,19 +1,27 @@
import { EmbeddingModel, FlagEmbedding } from "fastembed";
import type { FlagEmbedding } from "fastembed";
import { Embedder } from "./base";
import { EmbeddingConfig } from "../types";
const DEFAULT_MODEL = EmbeddingModel.BGESmallENV15;
type FastEmbedModel = Exclude<EmbeddingModel, EmbeddingModel.CUSTOM>;
// FastEmbed only ships a fixed set of ONNX models. Keep the list handy so we can
// reject unknown model names up front with a clear message instead of letting
// FlagEmbedding.init fail later with an opaque download error.
const SUPPORTED_MODELS = Object.values(EmbeddingModel).filter(
(model) => model !== EmbeddingModel.CUSTOM,
) as FastEmbedModel[];
// FastEmbed only ships a fixed set of ONNX models (fastembed's `EmbeddingModel`
// enum, minus CUSTOM). Mirrored here as literals so an invalid model name can
// be rejected synchronously in the constructor — with a clear message instead
// of a `FlagEmbedding.init()` download error — without eagerly importing the
// optional 'fastembed' package just to read its enum. Keep in sync if
// fastembed adds a model.
const SUPPORTED_MODELS = [
"fast-all-MiniLM-L6-v2",
"fast-bge-base-en",
"fast-bge-base-en-v1.5",
"fast-bge-small-en",
"fast-bge-small-en-v1.5",
"fast-bge-small-zh-v1.5",
"fast-multilingual-e5-large",
] as const;
type FastEmbedModel = (typeof SUPPORTED_MODELS)[number];
const DEFAULT_MODEL: FastEmbedModel = "fast-bge-small-en-v1.5";
export class FastEmbedEmbedder implements Embedder {
private modelName: FastEmbedModel;
private readonly modelName: FastEmbedModel;
private embeddingModel?: Promise<FlagEmbedding>;
constructor(config: EmbeddingConfig) {
@@ -32,9 +40,7 @@ export class FastEmbedEmbedder implements Embedder {
private getEmbeddingModel(): Promise<FlagEmbedding> {
if (!this.embeddingModel) {
this.embeddingModel = FlagEmbedding.init({
model: this.modelName,
}).catch((error) => {
this.embeddingModel = this.initEmbeddingModel().catch((error) => {
this.embeddingModel = undefined;
throw error;
});
@@ -43,6 +49,23 @@ export class FastEmbedEmbedder implements Embedder {
return this.embeddingModel;
}
/**
* Lazily import the optional `fastembed` peer and initialize the model, so
* consumers that never touch FastEmbed don't need it installed.
*/
private async initEmbeddingModel(): Promise<FlagEmbedding> {
let sdk: any;
try {
sdk = await import("fastembed");
} catch {
throw new Error(
"The 'fastembed' package is required to use the FastEmbed embedder. Install it with: npm install fastembed",
);
}
return sdk.FlagEmbedding.init({ model: this.modelName });
}
private normalizeInput(text: string): string {
return text.replace(/\n/g, " ");
}
+33 -5
View File
@@ -1,11 +1,12 @@
import { CohereClient } from "cohere-ai";
import { RerankerConfig } from "../types";
import { Reranker, RerankResult } from "./base";
const DEFAULT_MODEL = "rerank-v3.5";
export class CohereReranker implements Reranker {
private client: CohereClient;
private clientInstance?: any;
private clientPromise?: Promise<any>;
private readonly apiKey: string;
private model: string;
private topK?: number;
private returnDocuments: boolean;
@@ -18,13 +19,39 @@ export class CohereReranker implements Reranker {
"Cohere API key is required. Set COHERE_API_KEY environment variable or pass apiKey in config.",
);
}
this.client = new CohereClient({ token: apiKey });
this.apiKey = apiKey;
this.model = config.model || DEFAULT_MODEL;
this.topK = config.topK;
this.returnDocuments = config.returnDocuments ?? false;
this.maxChunksPerDoc = config.maxChunksPerDoc;
}
/**
* Lazily construct (or reuse) the Cohere client, importing the optional
* `cohere-ai` peer only when the reranker is first used so consumers that
* never touch Cohere don't need it installed.
*/
private async getClient(): Promise<any> {
if (this.clientInstance) return this.clientInstance;
if (!this.clientPromise) {
this.clientPromise = this.createClient();
}
this.clientInstance = await this.clientPromise;
return this.clientInstance;
}
private async createClient(): Promise<any> {
let sdk: any;
try {
sdk = await import("cohere-ai");
} catch {
throw new Error(
"The 'cohere-ai' package is required to use the Cohere reranker. Install it with: npm install cohere-ai",
);
}
return new sdk.CohereClient({ token: this.apiKey });
}
async rerank(
query: string,
documents: string[],
@@ -33,7 +60,8 @@ export class CohereReranker implements Reranker {
if (documents.length === 0) return [];
try {
const response = await this.client.rerank({
const client = await this.getClient();
const response = await client.rerank({
model: this.model,
query,
documents,
@@ -42,7 +70,7 @@ export class CohereReranker implements Reranker {
maxChunksPerDoc: this.maxChunksPerDoc,
});
return response.results.map((result) => ({
return response.results.map((result: any) => ({
index: result.index,
rerankScore: result.relevanceScore,
}));
+33 -5
View File
@@ -1,11 +1,12 @@
import { ZeroEntropy } from "zeroentropy";
import { RerankerConfig } from "../types";
import { Reranker, RerankResult } from "./base";
const DEFAULT_MODEL = "zerank-1";
export class ZeroEntropyReranker implements Reranker {
private client: ZeroEntropy;
private clientInstance?: any;
private clientPromise?: Promise<any>;
private readonly apiKey: string;
private model: string;
private topK?: number;
@@ -16,11 +17,37 @@ export class ZeroEntropyReranker implements Reranker {
"Zero Entropy API key is required. Set ZERO_ENTROPY_API_KEY environment variable or pass apiKey in config.",
);
}
this.client = new ZeroEntropy({ apiKey });
this.apiKey = apiKey;
this.model = config.model || DEFAULT_MODEL;
this.topK = config.topK;
}
/**
* Lazily construct (or reuse) the ZeroEntropy client, importing the
* optional `zeroentropy` peer only when the reranker is first used so
* consumers that never touch ZeroEntropy don't need it installed.
*/
private async getClient(): Promise<any> {
if (this.clientInstance) return this.clientInstance;
if (!this.clientPromise) {
this.clientPromise = this.createClient();
}
this.clientInstance = await this.clientPromise;
return this.clientInstance;
}
private async createClient(): Promise<any> {
let sdk: any;
try {
sdk = await import("zeroentropy");
} catch {
throw new Error(
"The 'zeroentropy' package is required to use the ZeroEntropy reranker. Install it with: npm install zeroentropy",
);
}
return new sdk.ZeroEntropy({ apiKey: this.apiKey });
}
async rerank(
query: string,
documents: string[],
@@ -29,13 +56,14 @@ export class ZeroEntropyReranker implements Reranker {
if (documents.length === 0) return [];
try {
const response = await this.client.models.rerank({
const client = await this.getClient();
const response = await client.models.rerank({
model: this.model,
query,
documents,
});
const scored = response.results.map((result) => ({
const scored: RerankResult[] = response.results.map((result: any) => ({
index: result.index,
rerankScore: result.relevance_score,
}));
+47 -61
View File
@@ -1,60 +1,42 @@
// jest.mock is hoisted above variable declarations, so shared mock functions
// are attached to the module-level `__mocks__` object populated inside the
// factory. The hoisted mock reaches them through this stable reference.
// The provider imports `chromadb` lazily (await import) only on first use, so
// the jest.mock factory no longer runs at module-eval time. Create the shared
// mock handles at module top-level (not inside the factory) so `beforeEach` can
// reach them before the lazy import has fired; the factory, which runs on first
// use, just returns references to them.
const add = jest.fn().mockResolvedValue(undefined);
const query = jest
.fn()
.mockResolvedValue({ ids: [[]], distances: [[]], metadatas: [[]] });
const get = jest.fn().mockResolvedValue({ ids: [], metadatas: [] });
const update = jest.fn().mockResolvedValue(undefined);
const upsert = jest.fn().mockResolvedValue(undefined);
const deleteFn = jest.fn().mockResolvedValue(undefined);
const __mocks__: {
add: jest.Mock;
query: jest.Mock;
get: jest.Mock;
update: jest.Mock;
upsert: jest.Mock;
deleteFn: jest.Mock;
getOrCreateCollection: jest.Mock;
deleteCollection: jest.Mock;
ChromaClient: jest.Mock;
CloudClient: jest.Mock;
} = {} as any;
const collectionHandle = { add, query, get, update, upsert, delete: deleteFn };
const getOrCreateCollection = jest.fn().mockResolvedValue(collectionHandle);
const deleteCollection = jest.fn().mockResolvedValue(undefined);
jest.mock("chromadb", () => {
const add = jest.fn().mockResolvedValue(undefined);
const query = jest
.fn()
.mockResolvedValue({ ids: [[]], distances: [[]], metadatas: [[]] });
const get = jest.fn().mockResolvedValue({ ids: [], metadatas: [] });
const update = jest.fn().mockResolvedValue(undefined);
const upsert = jest.fn().mockResolvedValue(undefined);
const deleteFn = jest.fn().mockResolvedValue(undefined);
const clientImpl = () => ({ getOrCreateCollection, deleteCollection });
const ChromaClient = jest.fn().mockImplementation(clientImpl);
const CloudClient = jest.fn().mockImplementation(clientImpl);
const collectionHandle = {
add,
query,
get,
update,
upsert,
delete: deleteFn,
};
const getOrCreateCollection = jest.fn().mockResolvedValue(collectionHandle);
const deleteCollection = jest.fn().mockResolvedValue(undefined);
const __mocks__ = {
add,
query,
get,
update,
upsert,
deleteFn,
getOrCreateCollection,
deleteCollection,
ChromaClient,
CloudClient,
};
const clientImpl = () => ({ getOrCreateCollection, deleteCollection });
const ChromaClient = jest.fn().mockImplementation(clientImpl);
const CloudClient = jest.fn().mockImplementation(clientImpl);
Object.assign(__mocks__, {
add,
query,
get,
update,
upsert,
deleteFn,
getOrCreateCollection,
deleteCollection,
ChromaClient,
CloudClient,
});
return { ChromaClient, CloudClient };
});
jest.mock("chromadb", () => ({
ChromaClient: __mocks__.ChromaClient,
CloudClient: __mocks__.CloudClient,
}));
import { ChromaDB } from "../vector_stores/chroma";
import { VectorStoreFactory } from "../utils/factory";
@@ -127,8 +109,10 @@ describe("VectorStoreFactory", () => {
});
describe("Constructor", () => {
it("builds a local ChromaClient with host and port", () => {
makeDb({ host: "localhost", port: 8000 });
// The client is built lazily on first use, so trigger initialize() before
// asserting how it was constructed.
it("builds a local ChromaClient with host and port", async () => {
await initDb({ host: "localhost", port: 8000 });
expect(__mocks__.ChromaClient).toHaveBeenCalledWith({
host: "localhost",
port: 8000,
@@ -136,8 +120,8 @@ describe("Constructor", () => {
expect(__mocks__.CloudClient).not.toHaveBeenCalled();
});
it("passes ssl and path through to ChromaClient when provided", () => {
makeDb({ host: "example.com", port: 443, ssl: true, path: "/db" });
it("passes ssl and path through to ChromaClient when provided", async () => {
await initDb({ host: "example.com", port: 443, ssl: true, path: "/db" });
expect(__mocks__.ChromaClient).toHaveBeenCalledWith({
host: "example.com",
port: 443,
@@ -146,12 +130,13 @@ describe("Constructor", () => {
});
});
it("builds a CloudClient when apiKey and tenant are set", () => {
new ChromaDB({
it("builds a CloudClient when apiKey and tenant are set", async () => {
const db = new ChromaDB({
collectionName: "test-collection",
apiKey: "key-123",
tenant: "tenant-abc",
} as any);
await db.initialize();
expect(__mocks__.CloudClient).toHaveBeenCalledWith({
apiKey: "key-123",
tenant: "tenant-abc",
@@ -160,13 +145,14 @@ describe("Constructor", () => {
expect(__mocks__.ChromaClient).not.toHaveBeenCalled();
});
it("honors an explicit cloud database name", () => {
new ChromaDB({
it("honors an explicit cloud database name", async () => {
const db = new ChromaDB({
collectionName: "test-collection",
apiKey: "key-123",
tenant: "tenant-abc",
database: "custom-db",
} as any);
await db.initialize();
expect(__mocks__.CloudClient).toHaveBeenCalledWith(
expect.objectContaining({ database: "custom-db" }),
);
+49 -60
View File
@@ -3,73 +3,62 @@
// the module-level `__mocks__` object that is populated inside the factory so
// that the hoisted mock can reach them via a stable reference.
const __mocks__: {
upsert: jest.Mock;
query: jest.Mock;
fetch: jest.Mock;
deleteOne: jest.Mock;
namespace: jest.Mock;
describeIndexStats: jest.Mock;
index: jest.Mock;
listIndexes: jest.Mock;
createIndex: jest.Mock;
deleteIndex: jest.Mock;
Pinecone: jest.Mock;
} = {} as any;
// The provider imports `@pinecone-database/pinecone` lazily (await import) only
// on first use, so the jest.mock factory no longer runs at module-eval time.
// Create the shared mock handles at module top-level (not inside the factory)
// so `beforeEach` can reach them before the lazy import has fired; the factory,
// which runs on first use, just returns a reference to the Pinecone class mock.
const upsert = jest.fn().mockResolvedValue(undefined);
const query = jest.fn().mockResolvedValue({ matches: [] });
const fetch = jest.fn().mockResolvedValue({ records: {} });
const deleteOne = jest.fn().mockResolvedValue(undefined);
jest.mock("@pinecone-database/pinecone", () => {
// These are created fresh inside the factory so hoisting is safe.
const upsert = jest.fn().mockResolvedValue(undefined);
const query = jest.fn().mockResolvedValue({ matches: [] });
const fetch = jest.fn().mockResolvedValue({ records: {} });
const deleteOne = jest.fn().mockResolvedValue(undefined);
const nsHandle = { upsert, query, fetch, deleteOne };
const namespace = jest.fn().mockReturnValue(nsHandle);
const nsHandle = { upsert, query, fetch, deleteOne };
const namespace = jest.fn().mockReturnValue(nsHandle);
const describeIndexStats = jest
.fn()
.mockResolvedValue({ totalRecordCount: 0, namespaces: {} });
const describeIndexStats = jest
.fn()
.mockResolvedValue({ totalRecordCount: 0, namespaces: {} });
const indexHandle = {
namespace,
describeIndexStats,
// expose ops directly for the no-namespace path
upsert,
query,
fetch,
deleteOne,
};
const index = jest.fn().mockReturnValue(indexHandle);
const indexHandle = {
namespace,
describeIndexStats,
// expose ops directly for the no-namespace path
upsert,
query,
fetch,
deleteOne,
};
const index = jest.fn().mockReturnValue(indexHandle);
const listIndexes = jest.fn().mockResolvedValue({ indexes: [] });
const createIndex = jest.fn().mockResolvedValue(undefined);
const deleteIndex = jest.fn().mockResolvedValue(undefined);
const listIndexes = jest.fn().mockResolvedValue({ indexes: [] });
const createIndex = jest.fn().mockResolvedValue(undefined);
const deleteIndex = jest.fn().mockResolvedValue(undefined);
const Pinecone = jest.fn().mockImplementation(() => ({
listIndexes,
createIndex,
deleteIndex,
index,
}));
const Pinecone = jest.fn().mockImplementation(() => ({
listIndexes,
createIndex,
deleteIndex,
index,
}));
const __mocks__ = {
upsert,
query,
fetch,
deleteOne,
namespace,
describeIndexStats,
index,
listIndexes,
createIndex,
deleteIndex,
Pinecone,
};
// Populate the shared reference so tests can reach the mocks.
Object.assign(__mocks__, {
upsert,
query,
fetch,
deleteOne,
namespace,
describeIndexStats,
index,
listIndexes,
createIndex,
deleteIndex,
Pinecone,
});
return { Pinecone };
});
jest.mock("@pinecone-database/pinecone", () => ({
Pinecone: __mocks__.Pinecone,
}));
import { PineconeDB } from "../vector_stores/pinecone";
import { VectorStoreFactory } from "../utils/factory";
+24 -8
View File
@@ -1,4 +1,3 @@
import cassandra from "cassandra-driver";
import { VectorStore } from "./base";
import { SearchFilters, VectorStoreConfig, VectorStoreResult } from "../types";
@@ -18,7 +17,9 @@ interface CassandraConfig extends VectorStoreConfig {
protocolVersion?: number;
loadBalancingPolicy?: any;
client?: CassandraClientLike;
driver?: typeof cassandra;
/** Pre-configured Cassandra driver module (typed as `any` to keep the
* optional driver's types out of the published type declarations). */
driver?: any;
}
interface CassandraClientLike {
@@ -38,7 +39,7 @@ interface CassandraVector {
export class CassandraDB implements VectorStore {
private static readonly PAGE_SIZE = 500;
private readonly driver: typeof cassandra;
private readonly driver?: any;
private readonly contactPoints?: string[];
private readonly port: number;
private readonly username?: string;
@@ -54,7 +55,7 @@ export class CassandraDB implements VectorStore {
private _initPromise?: Promise<void>;
constructor(config: CassandraConfig) {
this.driver = config.driver || cassandra;
this.driver = config.driver;
this.contactPoints = config.contactPoints;
this.port = config.port || 9042;
this.username = config.username;
@@ -85,7 +86,7 @@ export class CassandraDB implements VectorStore {
private async _doInitialize(): Promise<void> {
if (!this.client) {
this.client = this.createClient();
this.client = await this.createClient();
}
if (typeof this.client.connect === "function") {
await this.client.connect();
@@ -326,7 +327,8 @@ export class CassandraDB implements VectorStore {
);
}
private createClient(): CassandraClientLike {
private async createClient(): Promise<CassandraClientLike> {
const driver = this.driver ?? (await this.loadDriver());
const clientConfig: Record<string, any> = {};
if (this.secureConnectBundle) {
@@ -363,13 +365,27 @@ export class CassandraDB implements VectorStore {
};
}
if (this.username && this.password) {
clientConfig.authProvider = new this.driver.auth.PlainTextAuthProvider(
clientConfig.authProvider = new driver.auth.PlainTextAuthProvider(
this.username,
this.password,
);
}
return new this.driver.Client(clientConfig);
return new driver.Client(clientConfig);
}
// Loaded dynamically: cassandra-driver is an optional peer dependency, so a static
// value import would break `import { Memory } from "mem0ai/oss"` for everyone else.
private async loadDriver(): Promise<any> {
let sdk: any;
try {
sdk = await import("cassandra-driver");
} catch {
throw new Error(
"The 'cassandra-driver' package is required to use the Cassandra vector store. Install it with: npm install cassandra-driver",
);
}
return sdk.default ?? sdk;
}
private validateIdentifier(name: string, label: string): string {
+56 -22
View File
@@ -1,4 +1,4 @@
import { ChromaClient, CloudClient } from "chromadb";
import type { ChromaClient, CloudClient } from "chromadb";
import { VectorStore } from "./base";
import { SearchFilters, VectorStoreConfig, VectorStoreResult } from "../types";
@@ -32,36 +32,68 @@ const MIGRATIONS_COLLECTION = "memory_migrations";
* so no embedding function is required on the collection.
*/
export class ChromaDB implements VectorStore {
private client: any;
private clientInstance?: any;
private clientPromise?: Promise<any>;
private readonly config: ChromaConfig;
private readonly collectionName: string;
private collectionPromise?: Promise<any>;
private migrationsPromise?: Promise<any>;
constructor(config: ChromaConfig) {
if (config.client) {
this.client = config.client;
} else if (config.apiKey && config.tenant) {
this.client = new CloudClient({
apiKey: config.apiKey,
tenant: config.tenant,
database: config.database || "mem0",
} as any);
} else {
const params: Record<string, any> = {};
if (config.host) params.host = config.host;
if (config.port) params.port = config.port;
if (config.ssl !== undefined) params.ssl = config.ssl;
if (config.path) params.path = config.path;
this.client = new ChromaClient(params as any);
}
this.config = config;
this.collectionName = config.collectionName;
this.initialize().catch(console.error);
}
/**
* Lazily construct (or reuse) the ChromaDB client, importing the optional
* `chromadb` peer only when the store is first used so consumers that never
* touch Chroma don't need it installed.
*/
private async getClient(): Promise<any> {
if (this.clientInstance) return this.clientInstance;
if (!this.clientPromise) {
this.clientPromise = this.createClient();
}
this.clientInstance = await this.clientPromise;
return this.clientInstance;
}
private async createClient(): Promise<any> {
const config = this.config;
if (config.client) {
return config.client;
}
let sdk: any;
try {
sdk = await import("chromadb");
} catch {
throw new Error(
"The 'chromadb' package is required to use the Chroma vector store. Install it with: npm install chromadb",
);
}
if (config.apiKey && config.tenant) {
return new sdk.CloudClient({
apiKey: config.apiKey,
tenant: config.tenant,
database: config.database || "mem0",
} as any);
}
const params: Record<string, any> = {};
if (config.host) params.host = config.host;
if (config.port) params.port = config.port;
if (config.ssl !== undefined) params.ssl = config.ssl;
if (config.path) params.path = config.path;
return new sdk.ChromaClient(params as any);
}
private async getCollection(): Promise<any> {
if (!this.collectionPromise) {
this.collectionPromise = this.client.getOrCreateCollection({
const client = await this.getClient();
this.collectionPromise = client.getOrCreateCollection({
name: this.collectionName,
embeddingFunction: null,
});
@@ -71,7 +103,8 @@ export class ChromaDB implements VectorStore {
private async getMigrationsCollection(): Promise<any> {
if (!this.migrationsPromise) {
this.migrationsPromise = this.client.getOrCreateCollection({
const client = await this.getClient();
this.migrationsPromise = client.getOrCreateCollection({
name: MIGRATIONS_COLLECTION,
embeddingFunction: null,
});
@@ -170,7 +203,8 @@ export class ChromaDB implements VectorStore {
}
async deleteCol(): Promise<void> {
await this.client.deleteCollection({ name: this.collectionName });
const client = await this.getClient();
await client.deleteCollection({ name: this.collectionName });
this.collectionPromise = undefined;
}
@@ -1,9 +1,11 @@
import { Client } from "@elastic/elasticsearch";
import type { Client } from "@elastic/elasticsearch";
import { VectorStore } from "./base";
import { SearchFilters, VectorStoreConfig, VectorStoreResult } from "../types";
interface ElasticsearchConfig extends VectorStoreConfig {
client?: Client;
/** Pre-configured Elasticsearch client instance (typed as `any` to keep the
* optional driver's types out of the published type declarations). */
client?: any;
host?: string;
port?: number;
cloudId?: string;
@@ -41,17 +43,28 @@ function validateFilter(key: string, value: unknown): void {
}
export class ElasticsearchDB implements VectorStore {
private client: Client;
private client!: Client;
private readonly config: ElasticsearchConfig;
private readonly collectionName: string;
private readonly dimension: number;
private readonly autoCreateIndex: boolean;
private _initPromise?: Promise<void>;
constructor(config: ElasticsearchConfig) {
this.config = config;
this.collectionName = config.collectionName;
this.dimension = config.dimension || config.embeddingModelDims || 1536;
this.autoCreateIndex = config.autoCreateIndex !== false;
this.initialize().catch(console.error);
}
// The client is created lazily on first initialize() so the optional
// `@elastic/elasticsearch` peer is only loaded when the store is actually used.
private async ensureClient(): Promise<void> {
if (this.client) return;
const config = this.config;
if (config.client) {
this.client = config.client;
} else {
@@ -87,10 +100,17 @@ export class ElasticsearchDB implements VectorStore {
params.headers = config.headers;
}
this.client = new Client(params);
}
let sdk: any;
try {
sdk = await import("@elastic/elasticsearch");
} catch {
throw new Error(
"The '@elastic/elasticsearch' package is required to use the Elasticsearch vector store. Install it with: npm install @elastic/elasticsearch",
);
}
this.initialize().catch(console.error);
this.client = new sdk.Client(params);
}
}
async initialize(): Promise<void> {
@@ -101,6 +121,7 @@ export class ElasticsearchDB implements VectorStore {
}
private async _doInitialize(): Promise<void> {
await this.ensureClient();
try {
if (this.autoCreateIndex) {
await this.ensureIndex(this.collectionName, this.dimension);
@@ -149,6 +170,7 @@ export class ElasticsearchDB implements VectorStore {
ids: string[],
payloads: Record<string, any>[],
): Promise<void> {
await this.initialize();
const operations: any[] = [];
for (let i = 0; i < vectors.length; i++) {
operations.push(
@@ -165,6 +187,7 @@ export class ElasticsearchDB implements VectorStore {
topK: number = 5,
filters?: SearchFilters,
): Promise<VectorStoreResult[]> {
await this.initialize();
const searchBody: Record<string, any> = {
knn: {
field: "vector",
@@ -195,6 +218,7 @@ export class ElasticsearchDB implements VectorStore {
}
async get(vectorId: string): Promise<VectorStoreResult | null> {
await this.initialize();
try {
const response = await this.client.get({
index: this.collectionName,
@@ -215,6 +239,7 @@ export class ElasticsearchDB implements VectorStore {
vector: number[],
payload: Record<string, any>,
): Promise<void> {
await this.initialize();
const doc: Record<string, any> = {};
if (vector) doc.vector = vector;
if (payload) doc.metadata = payload;
@@ -227,6 +252,7 @@ export class ElasticsearchDB implements VectorStore {
}
async delete(vectorId: string): Promise<void> {
await this.initialize();
await this.client.delete({
index: this.collectionName,
id: vectorId,
@@ -234,6 +260,7 @@ export class ElasticsearchDB implements VectorStore {
}
async deleteCol(): Promise<void> {
await this.initialize();
await this.client.indices.delete({ index: this.collectionName });
}
@@ -241,6 +268,7 @@ export class ElasticsearchDB implements VectorStore {
filters?: SearchFilters,
topK: number = 100,
): Promise<[VectorStoreResult[], number]> {
await this.initialize();
const query: Record<string, any> = { query: { match_all: {} } };
if (filters && Object.keys(filters).length > 0) {
@@ -275,6 +303,7 @@ export class ElasticsearchDB implements VectorStore {
}
async getUserId(): Promise<string> {
await this.initialize();
try {
const response = await this.client.search({
index: "memory_migrations",
@@ -308,6 +337,7 @@ export class ElasticsearchDB implements VectorStore {
}
async setUserId(userId: string): Promise<void> {
await this.initialize();
try {
const response = await this.client.search({
index: "memory_migrations",
+25 -5
View File
@@ -1,4 +1,4 @@
import { MongoClient, Collection, Db } from "mongodb";
import type { MongoClient, Collection, Db } from "mongodb";
import { VectorStore } from "./base";
import { SearchFilters, VectorStoreConfig, VectorStoreResult } from "../types";
@@ -8,13 +8,16 @@ export interface MongoDBConfig extends VectorStoreConfig {
collectionName?: string;
embeddingModelDims?: number;
dimension?: number;
client?: MongoClient;
/** Pre-configured MongoDB client instance (typed as `any` to keep the
* optional driver's types out of the published type declarations). */
client?: any;
}
export class MongoDB implements VectorStore {
private client: MongoClient;
private db: Db;
private client!: MongoClient;
private db!: Db;
private collection!: Collection;
private readonly config: MongoDBConfig;
private readonly collectionName: string;
private readonly dbName: string;
private readonly embeddingModelDims: number;
@@ -22,17 +25,33 @@ export class MongoDB implements VectorStore {
private _initPromise?: Promise<void>;
constructor(config: MongoDBConfig) {
this.config = config;
this.collectionName = config.collectionName || "mem0";
this.dbName = config.dbName || "mem0_db";
this.embeddingModelDims =
config.embeddingModelDims || config.dimension || 1536;
this.indexName = `${this.collectionName}_vector_index`;
// The client/db are created lazily on first initialize() so the optional
// `mongodb` peer is only loaded when the store is actually used.
}
private async ensureClient(): Promise<void> {
if (this.client) return;
const config = this.config;
if (config.client) {
this.client = config.client;
} else {
let sdk: any;
try {
sdk = await import("mongodb");
} catch {
throw new Error(
"The 'mongodb' package is required to use the MongoDB vector store. Install it with: npm install mongodb",
);
}
const url = config.url || "mongodb://localhost:27017";
this.client = new MongoClient(url, { appName: "Mem0" });
this.client = new sdk.MongoClient(url, { appName: "Mem0" });
}
this.db = this.client.db(this.dbName);
@@ -46,6 +65,7 @@ export class MongoDB implements VectorStore {
}
private async _doInitialize(): Promise<void> {
await this.ensureClient();
try {
const collections = await this.db
.listCollections({ name: this.collectionName })
@@ -1,4 +1,4 @@
import { Client } from "@opensearch-project/opensearch";
import type { Client } from "@opensearch-project/opensearch";
import { VectorStore } from "./base";
import { SearchFilters, VectorStoreConfig, VectorStoreResult } from "../types";
@@ -14,7 +14,9 @@ type OpenSearchAuth =
| Record<string, any>;
interface OpenSearchConfig extends VectorStoreConfig {
client?: Client;
/** Pre-configured OpenSearch client instance (typed as `any` to keep the
* optional driver's types out of the published type declarations). */
client?: any;
host?: string;
port?: number;
httpAuth?: OpenSearchAuth | [string, string];
@@ -61,17 +63,29 @@ function escapeWildcard(value: string): string {
}
export class OpenSearchDB implements VectorStore {
private client: Client;
private client!: Client;
private readonly config: OpenSearchConfig;
private readonly collectionName: string;
private readonly embeddingModelDims: number;
private readonly autoRefresh: boolean;
private _initPromise?: Promise<void>;
constructor(config: OpenSearchConfig) {
this.config = config;
this.collectionName = config.collectionName;
this.embeddingModelDims = config.embeddingModelDims;
this.autoRefresh = config.autoRefresh ?? false;
this.initialize().catch(console.error);
}
// The client is created lazily on first initialize() so the optional
// `@opensearch-project/opensearch` peer is only loaded when the store is
// actually used.
private async ensureClient(): Promise<void> {
if (this.client) return;
const config = this.config;
if (config.client) {
this.client = config.client;
} else {
@@ -84,7 +98,16 @@ export class OpenSearchDB implements VectorStore {
? { username: config.user, password: config.password }
: undefined);
this.client = new Client({
let sdk: any;
try {
sdk = await import("@opensearch-project/opensearch");
} catch {
throw new Error(
"The '@opensearch-project/opensearch' package is required to use the OpenSearch vector store. Install it with: npm install @opensearch-project/opensearch",
);
}
this.client = new sdk.Client({
node: `${useSSL ? "https" : "http"}://${host}:${port}`,
auth: this.normalizeAuth(auth),
ssl: {
@@ -94,8 +117,6 @@ export class OpenSearchDB implements VectorStore {
},
});
}
this.initialize().catch(console.error);
}
async initialize(): Promise<void> {
@@ -107,6 +128,7 @@ export class OpenSearchDB implements VectorStore {
}
private async _doInitialize(): Promise<void> {
await this.ensureClient();
await this.createCol(this.collectionName, this.embeddingModelDims);
await this.ensureMigrationIndex();
}
@@ -210,6 +232,7 @@ export class OpenSearchDB implements VectorStore {
ids: string[],
payloads: Record<string, any>[],
): Promise<void> {
await this.initialize();
vectors.forEach((vector, index) => this.validateVector(vector, index));
const operations = vectors.flatMap((vector, index) => {
@@ -250,6 +273,7 @@ export class OpenSearchDB implements VectorStore {
topK: number = 5,
filters?: SearchFilters,
): Promise<VectorStoreResult[] | null> {
await this.initialize();
const boolQuery: Record<string, any> = {
should: [
{ match: { "payload.data": query } },
@@ -281,6 +305,7 @@ export class OpenSearchDB implements VectorStore {
topK: number = 5,
filters?: SearchFilters,
): Promise<VectorStoreResult[]> {
await this.initialize();
const knnQuery = {
knn: {
vector_field: {
@@ -316,6 +341,7 @@ export class OpenSearchDB implements VectorStore {
}
async get(vectorId: string): Promise<VectorStoreResult | null> {
await this.initialize();
try {
const response = responseBody<{ _source?: OpenSearchHit["_source"] }>(
await this.client.get({
@@ -343,6 +369,7 @@ export class OpenSearchDB implements VectorStore {
vector: number[],
payload: Record<string, any>,
): Promise<void> {
await this.initialize();
if (vector) {
this.validateVector(vector, 0);
}
@@ -362,6 +389,7 @@ export class OpenSearchDB implements VectorStore {
}
async delete(vectorId: string): Promise<void> {
await this.initialize();
try {
await this.client.delete({
index: this.collectionName,
@@ -377,6 +405,7 @@ export class OpenSearchDB implements VectorStore {
}
async deleteCol(): Promise<void> {
await this.initialize();
if (!(await this.indexExists(this.collectionName))) return;
await this.client.indices.delete({ index: this.collectionName });
}
@@ -385,6 +414,7 @@ export class OpenSearchDB implements VectorStore {
filters?: SearchFilters,
topK: number = 100,
): Promise<[VectorStoreResult[], number]> {
await this.initialize();
const filter = this.buildFilterClauses(filters);
const query = filter.length ? { bool: { filter } } : { match_all: {} };
@@ -413,11 +443,13 @@ export class OpenSearchDB implements VectorStore {
}
async reset(): Promise<void> {
await this.initialize();
await this.deleteCol();
await this.createCol(this.collectionName, this.embeddingModelDims);
}
async getUserId(): Promise<string> {
await this.initialize();
await this.ensureMigrationIndex();
const response = responseBody<{ hits: { hits: OpenSearchHit[] } }>(
@@ -442,6 +474,7 @@ export class OpenSearchDB implements VectorStore {
}
async setUserId(userId: string): Promise<void> {
await this.initialize();
await this.ensureMigrationIndex();
await this.client.index({
index: "memory_migrations",
+34 -8
View File
@@ -1,5 +1,4 @@
import { Pinecone } from "@pinecone-database/pinecone";
import type { Index } from "@pinecone-database/pinecone";
import type { Pinecone, Index } from "@pinecone-database/pinecone";
import { VectorStore } from "./base";
import { SearchFilters, VectorStoreConfig, VectorStoreResult } from "../types";
@@ -9,7 +8,9 @@ const MIGRATIONS_RECORD_ID = "mem0-user-id";
interface PineconeDBConfig extends VectorStoreConfig {
collectionName: string;
embeddingModelDims: number;
client?: Pinecone;
/** Pre-configured Pinecone client instance (typed as `any` to keep the
* optional driver's types out of the published type declarations). */
client?: any;
apiKey?: string;
serverlessConfig?: { cloud: string; region: string };
podConfig?: {
@@ -26,7 +27,8 @@ interface PineconeDBConfig extends VectorStoreConfig {
}
export class PineconeDB implements VectorStore {
private client: Pinecone;
private client!: Pinecone;
private readonly config: PineconeDBConfig;
private readonly collectionName: string;
private readonly dimension: number;
private readonly metric: "cosine" | "dotproduct" | "euclidean";
@@ -45,18 +47,16 @@ export class PineconeDB implements VectorStore {
private _initPromise?: Promise<void>;
constructor(config: PineconeDBConfig) {
if (config.client) {
this.client = config.client;
} else {
if (!config.client) {
const apiKey = config.apiKey || process.env.PINECONE_API_KEY;
if (!apiKey) {
throw new Error(
"Pinecone API key required: pass apiKey or set PINECONE_API_KEY env var",
);
}
this.client = new Pinecone({ apiKey });
}
this.config = config;
this.collectionName = config.collectionName;
this.dimension = config.embeddingModelDims || config.dimension || 1536;
this.metric = config.metric || "cosine";
@@ -69,6 +69,31 @@ export class PineconeDB implements VectorStore {
this.initialize().catch(console.error);
}
/**
* Lazily construct (or reuse) the Pinecone client, importing the optional
* `@pinecone-database/pinecone` peer only when the store is first used so
* consumers that never touch Pinecone don't need it installed.
*/
private async ensureClient(): Promise<void> {
if (this.client) return;
const config = this.config;
if (config.client) {
this.client = config.client;
} else {
const apiKey = config.apiKey || process.env.PINECONE_API_KEY;
let sdk: any;
try {
sdk = await import("@pinecone-database/pinecone");
} catch {
throw new Error(
"The '@pinecone-database/pinecone' package is required to use the Pinecone vector store. Install it with: npm install @pinecone-database/pinecone",
);
}
this.client = new sdk.Pinecone({ apiKey });
}
}
async initialize(): Promise<void> {
if (!this._initPromise) {
this._initPromise = this._doInitialize();
@@ -77,6 +102,7 @@ export class PineconeDB implements VectorStore {
}
private async _doInitialize(): Promise<void> {
await this.ensureClient();
await this._ensureIndex();
this._index = this.client.index({ name: this.collectionName });
}
+110 -62
View File
@@ -1,21 +1,8 @@
import {
CreateIndexCommand,
CreateVectorBucketCommand,
DeleteIndexCommand,
DeleteVectorsCommand,
GetIndexCommand,
GetVectorBucketCommand,
GetVectorsCommand,
ListVectorsCommand,
PutVectorsCommand,
QueryVectorsCommand,
S3VectorsClient,
type DistanceMetric,
type GetOutputVector,
type ListOutputVector,
type QueryOutputVector,
type S3VectorsClientConfig,
type VectorData,
import type {
GetOutputVector,
ListOutputVector,
QueryOutputVector,
VectorData,
} from "@aws-sdk/client-s3vectors";
import { VectorStore } from "./base";
import { SearchFilters, VectorStoreConfig, VectorStoreResult } from "../types";
@@ -32,11 +19,13 @@ interface S3VectorsConfig extends VectorStoreConfig {
collectionName: string;
embeddingModelDims?: number;
dimension?: number;
distanceMetric?: DistanceMetric | "cosine" | "euclidean";
distanceMetric?: "cosine" | "euclidean";
region?: string;
regionName?: string;
client?: S3VectorsClientLike;
clientConfig?: S3VectorsClientConfig;
/** Pre-configured S3 Vectors client options (typed as `any` to keep the
* optional SDK's types out of the published type declarations). */
clientConfig?: any;
}
interface S3VectorsClientLike {
@@ -44,11 +33,14 @@ interface S3VectorsClientLike {
}
export class S3Vectors implements VectorStore {
private readonly client: S3VectorsClientLike;
private readonly config: S3VectorsConfig;
private readonly vectorBucketName: string;
private readonly collectionName: string;
private readonly dimension: number;
private readonly distanceMetric: DistanceMetric | "cosine" | "euclidean";
private readonly distanceMetric: "cosine" | "euclidean";
private client?: S3VectorsClientLike;
private clientPromise?: Promise<S3VectorsClientLike>;
private sdkPromise?: Promise<any>;
private _initPromise?: Promise<void>;
private cachedUserId?: string;
@@ -65,22 +57,55 @@ export class S3Vectors implements VectorStore {
throw new Error("embeddingModelDims or dimension is required");
}
this.config = config;
this.vectorBucketName = config.vectorBucketName;
this.collectionName = config.collectionName;
this.dimension = dimension;
this.distanceMetric = config.distanceMetric || "cosine";
this.client =
config.client ||
new S3VectorsClient({
...(config.clientConfig || {}),
...(config.region || config.regionName
? { region: config.region || config.regionName }
: {}),
});
void this.initialize().catch(console.error);
}
/**
* Lazily import the optional `@aws-sdk/client-s3vectors` peer so consumers
* who never use the S3 Vectors store don't need it installed.
*/
private getSdk(): Promise<any> {
if (!this.sdkPromise) {
this.sdkPromise = import("@aws-sdk/client-s3vectors").catch(() => {
throw new Error(
"The '@aws-sdk/client-s3vectors' package is required to use the S3 Vectors store. Install it with: npm install @aws-sdk/client-s3vectors",
);
});
}
return this.sdkPromise;
}
/** Lazily construct (or reuse) the S3 Vectors client. */
private async getClient(): Promise<S3VectorsClientLike> {
if (this.client) return this.client;
if (!this.clientPromise) {
this.clientPromise = this.createClient();
}
this.client = await this.clientPromise;
return this.client;
}
private async createClient(): Promise<S3VectorsClientLike> {
const config = this.config;
if (config.client) {
return config.client;
}
const sdk = await this.getSdk();
return new sdk.S3VectorsClient({
...(config.clientConfig || {}),
...(config.region || config.regionName
? { region: config.region || config.regionName }
: {}),
});
}
async initialize(): Promise<void> {
if (!this._initPromise) {
this._initPromise = this._doInitialize();
@@ -105,8 +130,10 @@ export class S3Vectors implements VectorStore {
await this.initialize();
this.assertBatchDimensions(vectors, "Insert");
await this.client.send(
new PutVectorsCommand({
const sdk = await this.getSdk();
const client = await this.getClient();
await client.send(
new sdk.PutVectorsCommand({
vectorBucketName: this.vectorBucketName,
indexName: this.collectionName,
vectors: vectors.map((vector, index) => ({
@@ -134,8 +161,10 @@ export class S3Vectors implements VectorStore {
if (filter && this.isAlwaysFalseFilter(filter)) {
return [];
}
const response = await this.client.send(
new QueryVectorsCommand({
const sdk = await this.getSdk();
const client = await this.getClient();
const response = await client.send(
new sdk.QueryVectorsCommand({
vectorBucketName: this.vectorBucketName,
indexName: this.collectionName,
queryVector: this.toVectorData(query),
@@ -158,9 +187,11 @@ export class S3Vectors implements VectorStore {
async get(vectorId: string): Promise<VectorStoreResult | null> {
await this.initialize();
const sdk = await this.getSdk();
const client = await this.getClient();
try {
const response = await this.client.send(
new GetVectorsCommand({
const response = await client.send(
new sdk.GetVectorsCommand({
vectorBucketName: this.vectorBucketName,
indexName: this.collectionName,
keys: [vectorId],
@@ -203,8 +234,10 @@ export class S3Vectors implements VectorStore {
this.assertVectorDimension(nextVector, "Vector");
await this.client.send(
new PutVectorsCommand({
const sdk = await this.getSdk();
const client = await this.getClient();
await client.send(
new sdk.PutVectorsCommand({
vectorBucketName: this.vectorBucketName,
indexName: this.collectionName,
vectors: [
@@ -221,8 +254,10 @@ export class S3Vectors implements VectorStore {
async delete(vectorId: string): Promise<void> {
await this.initialize();
await this.client.send(
new DeleteVectorsCommand({
const sdk = await this.getSdk();
const client = await this.getClient();
await client.send(
new sdk.DeleteVectorsCommand({
vectorBucketName: this.vectorBucketName,
indexName: this.collectionName,
keys: [vectorId],
@@ -233,9 +268,11 @@ export class S3Vectors implements VectorStore {
async deleteCol(): Promise<void> {
await this.initialize();
const sdk = await this.getSdk();
const client = await this.getClient();
try {
await this.client.send(
new DeleteIndexCommand({
await client.send(
new sdk.DeleteIndexCommand({
vectorBucketName: this.vectorBucketName,
indexName: this.collectionName,
}),
@@ -257,14 +294,16 @@ export class S3Vectors implements VectorStore {
const filter = this.convertFilters(filters);
const results: VectorStoreResult[] = [];
let nextToken: string | undefined;
const sdk = await this.getSdk();
const client = await this.getClient();
// Stop paginating once we have topK matches. Both callers of list()
// discard the count, so scanning the whole index just to total every
// match is wasted round-trips (O(index size) on getAll/deleteAll).
// Return the page length like qdrant does.
do {
const response = await this.client.send(
new ListVectorsCommand({
const response = await client.send(
new sdk.ListVectorsCommand({
vectorBucketName: this.vectorBucketName,
indexName: this.collectionName,
maxResults: DEFAULT_PAGE_SIZE,
@@ -298,8 +337,10 @@ export class S3Vectors implements VectorStore {
await this.ensureMigrationIndex();
const response = await this.client.send(
new GetVectorsCommand({
const sdk = await this.getSdk();
const client = await this.getClient();
const response = await client.send(
new sdk.GetVectorsCommand({
vectorBucketName: this.vectorBucketName,
indexName: MIGRATION_INDEX_NAME,
keys: [MIGRATION_VECTOR_KEY],
@@ -325,8 +366,10 @@ export class S3Vectors implements VectorStore {
await this.initialize();
await this.ensureMigrationIndex();
await this.client.send(
new PutVectorsCommand({
const sdk = await this.getSdk();
const client = await this.getClient();
await client.send(
new sdk.PutVectorsCommand({
vectorBucketName: this.vectorBucketName,
indexName: MIGRATION_INDEX_NAME,
vectors: [
@@ -343,9 +386,11 @@ export class S3Vectors implements VectorStore {
}
private async ensureBucketExists(): Promise<void> {
const sdk = await this.getSdk();
const client = await this.getClient();
try {
await this.client.send(
new GetVectorBucketCommand({
await client.send(
new sdk.GetVectorBucketCommand({
vectorBucketName: this.vectorBucketName,
}),
);
@@ -355,8 +400,8 @@ export class S3Vectors implements VectorStore {
}
try {
await this.client.send(
new CreateVectorBucketCommand({
await client.send(
new sdk.CreateVectorBucketCommand({
vectorBucketName: this.vectorBucketName,
}),
);
@@ -371,11 +416,13 @@ export class S3Vectors implements VectorStore {
private async ensureIndexExists(
indexName: string,
dimension: number,
distanceMetric: DistanceMetric | "cosine" | "euclidean",
distanceMetric: "cosine" | "euclidean",
): Promise<void> {
const sdk = await this.getSdk();
const client = await this.getClient();
try {
await this.client.send(
new GetIndexCommand({
await client.send(
new sdk.GetIndexCommand({
vectorBucketName: this.vectorBucketName,
indexName,
}),
@@ -386,8 +433,8 @@ export class S3Vectors implements VectorStore {
}
try {
await this.client.send(
new CreateIndexCommand({
await client.send(
new sdk.CreateIndexCommand({
vectorBucketName: this.vectorBucketName,
indexName,
dataType: "float32",
@@ -410,8 +457,10 @@ export class S3Vectors implements VectorStore {
private async fetchStoredVector(
vectorId: string,
): Promise<{ vector: number[]; payload: Record<string, any> } | null> {
const response = await this.client.send(
new GetVectorsCommand({
const sdk = await this.getSdk();
const client = await this.getClient();
const response = await client.send(
new sdk.GetVectorsCommand({
vectorBucketName: this.vectorBucketName,
indexName: this.collectionName,
keys: [vectorId],
@@ -770,7 +819,7 @@ export class S3Vectors implements VectorStore {
private normalizeQueryVector(
vector: QueryOutputVector,
distanceMetric: DistanceMetric | "cosine" | "euclidean",
distanceMetric: "cosine" | "euclidean",
): VectorStoreResult {
return {
id: String(vector.key),
@@ -794,8 +843,7 @@ export class S3Vectors implements VectorStore {
private normalizeScore(
distance?: number,
distanceMetric: DistanceMetric | "cosine" | "euclidean" = this
.distanceMetric,
distanceMetric: "cosine" | "euclidean" = this.distanceMetric,
): number | undefined {
if (distance === undefined || distance === null) {
return undefined;
@@ -1,4 +1,3 @@
import Turbopuffer from "@turbopuffer/turbopuffer";
import { VectorStore } from "./base";
import { SearchFilters, VectorStoreConfig, VectorStoreResult } from "../types";
@@ -11,11 +10,10 @@ interface TurbopufferConfig extends VectorStoreConfig {
}
export class TurbopufferDB implements VectorStore {
private client: Turbopuffer;
private ns: ReturnType<InstanceType<typeof Turbopuffer>["namespace"]>;
private migrationsNs: ReturnType<
InstanceType<typeof Turbopuffer>["namespace"]
>;
private clientInstance?: any;
private clientPromise?: Promise<any>;
private readonly apiKey: string;
private readonly region: string;
private readonly collectionName: string;
private readonly distanceMetric: string;
private readonly batchSize: number;
@@ -28,17 +26,55 @@ export class TurbopufferDB implements VectorStore {
);
}
this.client = new Turbopuffer({
apiKey,
region: config.region ?? "gcp-us-central1",
});
this.apiKey = apiKey;
this.region = config.region ?? "gcp-us-central1";
this.collectionName = config.collectionName;
this.distanceMetric = config.distanceMetric ?? "cosine_distance";
this.batchSize = config.batchSize ?? 100;
this.ns = this.client.namespace(this.collectionName);
this.migrationsNs = this.client.namespace(
this.collectionName + "_migrations",
);
}
/**
* Lazily construct (or reuse) the Turbopuffer client, importing the optional
* `@turbopuffer/turbopuffer` peer only when the store is first used so
* consumers that never touch Turbopuffer don't need it installed.
*/
private async getClient(): Promise<any> {
if (this.clientInstance) return this.clientInstance;
if (!this.clientPromise) {
this.clientPromise = this.createClient();
}
this.clientInstance = await this.clientPromise;
return this.clientInstance;
}
private async createClient(): Promise<any> {
let sdk: any;
try {
sdk = await import("@turbopuffer/turbopuffer");
} catch {
throw new Error(
"The '@turbopuffer/turbopuffer' package is required to use the Turbopuffer vector store. Install it with: npm install @turbopuffer/turbopuffer",
);
}
// @turbopuffer/turbopuffer ships `Turbopuffer` as both the default export
// and a named export pointing at the same class. Use `.default` since
// that's what a plain `import Turbopuffer from "..."` resolves to (and
// what test doubles for this module mock).
return new sdk.default({
apiKey: this.apiKey,
region: this.region,
});
}
private async getNs(): Promise<any> {
const client = await this.getClient();
return client.namespace(this.collectionName);
}
private async getMigrationsNs(): Promise<any> {
const client = await this.getClient();
return client.namespace(this.collectionName + "_migrations");
}
async initialize(): Promise<void> {
@@ -50,6 +86,7 @@ export class TurbopufferDB implements VectorStore {
ids: string[],
payloads: Record<string, any>[],
): Promise<void> {
const ns = await this.getNs();
for (let i = 0; i < vectors.length; i += this.batchSize) {
const batchVectors = vectors.slice(i, i + this.batchSize);
const batchIds = ids.slice(i, i + this.batchSize);
@@ -61,7 +98,7 @@ export class TurbopufferDB implements VectorStore {
vector,
}));
await this.ns.write({
await ns.write({
upsert_rows,
distance_metric: this.distanceMetric as any,
});
@@ -82,8 +119,9 @@ export class TurbopufferDB implements VectorStore {
const tpufFilters = this.convertFilters(filters);
if (tpufFilters !== null) queryParams.filters = tpufFilters;
const ns = await this.getNs();
try {
const result = await this.ns.query(queryParams);
const result = await ns.query(queryParams);
return this.parseRows(result.rows ?? []);
} catch (err) {
console.error("Turbopuffer search error:", err);
@@ -96,8 +134,9 @@ export class TurbopufferDB implements VectorStore {
}
async get(vectorId: string): Promise<VectorStoreResult | null> {
const ns = await this.getNs();
try {
const result = await this.ns.query({
const result = await ns.query({
rank_by: ["id", "asc"] as any,
top_k: 1,
include_attributes: true,
@@ -116,24 +155,27 @@ export class TurbopufferDB implements VectorStore {
vector: number[],
payload: Record<string, any>,
): Promise<void> {
const ns = await this.getNs();
if (vector && vector.length > 0) {
await this.ns.write({
await ns.write({
upsert_rows: [{ ...payload, id: vectorId, vector }],
distance_metric: this.distanceMetric as any,
});
} else {
await this.ns.write({
await ns.write({
patch_rows: [{ ...payload, id: vectorId }],
});
}
}
async delete(vectorId: string): Promise<void> {
await this.ns.write({ deletes: [vectorId] });
const ns = await this.getNs();
await ns.write({ deletes: [vectorId] });
}
async deleteCol(): Promise<void> {
await this.ns.deleteAll();
const ns = await this.getNs();
await ns.deleteAll();
}
async list(
@@ -149,8 +191,9 @@ export class TurbopufferDB implements VectorStore {
const tpufFilters = this.convertFilters(filters);
if (tpufFilters !== null) queryParams.filters = tpufFilters;
const ns = await this.getNs();
try {
const result = await this.ns.query(queryParams);
const result = await ns.query(queryParams);
const rows = this.parseRows(result.rows ?? []);
return [rows, rows.length];
} catch (err) {
@@ -161,9 +204,10 @@ export class TurbopufferDB implements VectorStore {
async getUserId(): Promise<string> {
try {
const migrationsNs = await this.getMigrationsNs();
let rows: any[] = [];
try {
const result = await this.migrationsNs.query({
const result = await migrationsNs.query({
rank_by: ["id", "asc"] as any,
top_k: 1,
include_attributes: true,
@@ -181,7 +225,7 @@ export class TurbopufferDB implements VectorStore {
const randomId =
Math.random().toString(36).slice(2, 15) +
Math.random().toString(36).slice(2, 15);
await this.migrationsNs.write({
await migrationsNs.write({
upsert_rows: [{ id: "1", vector: [0.0], user_id: randomId }],
distance_metric: "cosine_distance" as any,
});
@@ -194,7 +238,8 @@ export class TurbopufferDB implements VectorStore {
async setUserId(userId: string): Promise<void> {
try {
await this.migrationsNs.write({
const migrationsNs = await this.getMigrationsNs();
await migrationsNs.write({
upsert_rows: [{ id: "1", vector: [0.0], user_id: userId }],
distance_metric: "cosine_distance" as any,
});
@@ -1,4 +1,4 @@
import { Index, QueryResult, Vector } from "@upstash/vector";
import type { Index, QueryResult, Vector } from "@upstash/vector";
import { VectorStore } from "./base";
import { SearchFilters, VectorStoreConfig, VectorStoreResult } from "../types";
@@ -6,36 +6,59 @@ interface UpstashVectorConfig extends VectorStoreConfig {
collectionName: string;
url?: string;
token?: string;
client?: Index<Record<string, unknown>>;
/** Pre-configured Upstash Vector client instance (typed as `any` to keep
* the optional driver's types out of the published type declarations). */
client?: any;
}
type UpstashMetadata = Record<string, unknown>;
export class UpstashVector implements VectorStore {
private readonly client: Index<UpstashMetadata>;
private client!: Index<UpstashMetadata>;
private readonly config: UpstashVectorConfig;
private readonly collectionName: string;
constructor(config: UpstashVectorConfig) {
if (!config.collectionName) {
throw new Error("collectionName is required for Upstash Vector.");
}
if (config.client) {
this.client = config.client;
} else if (config.url && config.token) {
this.client = new Index({
url: config.url,
token: config.token,
});
} else {
if (!config.client && !(config.url && config.token)) {
throw new Error("Either a client or url and token must be provided.");
}
this.config = config;
this.collectionName = config.collectionName;
}
/**
* Lazily construct (or reuse) the Upstash Vector client, importing the
* optional `@upstash/vector` peer only when the store is first used so
* consumers that never touch Upstash Vector don't need it installed.
*/
private async ensureClient(): Promise<void> {
if (this.client) return;
const config = this.config;
if (config.client) {
this.client = config.client;
} else {
let sdk: any;
try {
sdk = await import("@upstash/vector");
} catch {
throw new Error(
"The '@upstash/vector' package is required to use the Upstash Vector store. Install it with: npm install @upstash/vector",
);
}
this.client = new sdk.Index({
url: config.url,
token: config.token,
});
}
}
async initialize(): Promise<void> {
return;
await this.ensureClient();
}
async insert(
@@ -43,6 +66,7 @@ export class UpstashVector implements VectorStore {
ids: string[],
payloads: Record<string, any>[],
): Promise<void> {
await this.initialize();
const upsertData = vectors.map((vector, idx) => {
return {
id: ids[idx],
@@ -59,6 +83,7 @@ export class UpstashVector implements VectorStore {
topK: number = 5,
filters?: SearchFilters,
): Promise<VectorStoreResult[]> {
await this.initialize();
const response = await this.client.query<UpstashMetadata>(
{
vector: query,
@@ -77,6 +102,7 @@ export class UpstashVector implements VectorStore {
topK: number = 5,
filters?: SearchFilters,
): Promise<VectorStoreResult[] | null> {
await this.initialize();
try {
const response = await this.client.query<UpstashMetadata>(
{
@@ -96,6 +122,7 @@ export class UpstashVector implements VectorStore {
}
async get(vectorId: string): Promise<VectorStoreResult | null> {
await this.initialize();
const response = await this.client.fetch<UpstashMetadata>([vectorId], {
includeMetadata: true,
namespace: this.collectionName,
@@ -117,6 +144,7 @@ export class UpstashVector implements VectorStore {
vector: number[],
payload: Record<string, any>,
): Promise<void> {
await this.initialize();
// Upstash's `update` can't set the vector and metadata in one call (its
// payload is a discriminated union of vector | data | metadata), so a
// single `upsert` replaces both atomically, the same way insert() writes.
@@ -131,10 +159,12 @@ export class UpstashVector implements VectorStore {
}
async delete(vectorId: string): Promise<void> {
await this.initialize();
await this.client.delete(vectorId, { namespace: this.collectionName });
}
async deleteCol(): Promise<void> {
await this.initialize();
await this.client.reset({ namespace: this.collectionName });
}
@@ -142,6 +172,7 @@ export class UpstashVector implements VectorStore {
filters?: SearchFilters,
topK: number = 100,
): Promise<[VectorStoreResult[], number]> {
await this.initialize();
const results: VectorStoreResult[] = [];
let cursor = "0";
+29 -6
View File
@@ -1,10 +1,12 @@
import weaviate, { Filters, type WeaviateClient } from "weaviate-client";
import type { WeaviateClient } from "weaviate-client";
import { v4 as uuidv4 } from "uuid";
import { VectorStore } from "./base";
import { SearchFilters, VectorStoreConfig, VectorStoreResult } from "../types";
interface WeaviateConfig extends VectorStoreConfig {
client?: WeaviateClient;
/** Pre-configured Weaviate client instance (typed as `any` to keep the
* optional driver's types out of the published type declarations). */
client?: any;
clusterUrl?: string;
apiKey?: string;
additionalHeaders?: Record<string, string>;
@@ -28,6 +30,7 @@ const RETURN_PROPERTIES = [
export class WeaviateDB implements VectorStore {
private _config: WeaviateConfig;
private _client!: WeaviateClient;
private _sdk: any;
private _col!: any;
private _userId: string;
private _initPromise?: Promise<void>;
@@ -42,9 +45,23 @@ export class WeaviateDB implements VectorStore {
return (this._initPromise ??= this._doInitialize());
}
private async _doInitialize(): Promise<void> {
const { client, clusterUrl, apiKey, additionalHeaders, collectionName } =
this._config;
// Loaded dynamically: weaviate-client is an optional peer dependency, so a static
// value import would break `import { Memory } from "mem0ai/oss"` for everyone else.
private async ensureClient(): Promise<void> {
if (this._client) return;
let sdk: any;
try {
sdk = await import("weaviate-client");
} catch {
throw new Error(
"The 'weaviate-client' package is required to use the Weaviate vector store. Install it with: npm install weaviate-client",
);
}
this._sdk = sdk;
const { client, clusterUrl, apiKey, additionalHeaders } = this._config;
const weaviate = sdk.default;
if (client) {
this._client = client;
@@ -79,6 +96,12 @@ export class WeaviateDB implements VectorStore {
headers: additionalHeaders,
});
}
}
private async _doInitialize(): Promise<void> {
await this.ensureClient();
const { collectionName } = this._config;
const weaviate = this._sdk.default;
const exists = await this._client.collections.exists(collectionName);
if (!exists) {
@@ -101,7 +124,7 @@ export class WeaviateDB implements VectorStore {
const conditions = (["user_id", "agent_id", "run_id"] as const)
.filter((key) => filters[key] != null)
.map((key) => this._col.filter.byProperty(key).equal(filters[key]));
return conditions.length ? Filters.and(...conditions) : undefined;
return conditions.length ? this._sdk.Filters.and(...conditions) : undefined;
}
async insert(
@@ -50,8 +50,10 @@ beforeEach(() => {
describe("ElasticsearchDB", () => {
describe("constructor", () => {
it("creates client with self-hosted host:port", () => {
new ElasticsearchDB({
// The client is built lazily on first use, so await initialize() before
// asserting how it was constructed.
it("creates client with self-hosted host:port", async () => {
const store = new ElasticsearchDB({
collectionName: "mem0",
embeddingModelDims: 768,
host: "localhost",
@@ -59,6 +61,7 @@ describe("ElasticsearchDB", () => {
username: "user",
password: "pass",
});
await store.initialize();
expect(mockClient).toHaveBeenCalledWith(
expect.objectContaining({
@@ -68,14 +71,15 @@ describe("ElasticsearchDB", () => {
);
});
it("creates client with cloud config", () => {
new ElasticsearchDB({
it("creates client with cloud config", async () => {
const store = new ElasticsearchDB({
collectionName: "mem0",
embeddingModelDims: 1536,
cloudId:
"my-cloud:dXMtZWFzdDQuZ2NwLmVsYXN0aWMtY2xvdWQuY29tOjQ0MyQxMjM0NTY3ODkw",
apiKey: "base64-key",
});
await store.initialize();
expect(mockClient).toHaveBeenCalledWith(
expect.objectContaining({
@@ -98,12 +102,13 @@ describe("ElasticsearchDB", () => {
expect(mockClient).toHaveBeenCalledTimes(1);
});
it("defaults to port 9200 and https", () => {
new ElasticsearchDB({
it("defaults to port 9200 and https", async () => {
const store = new ElasticsearchDB({
collectionName: "mem0",
embeddingModelDims: 384,
host: "es.example.com",
});
await store.initialize();
expect(mockClient).toHaveBeenCalledWith(
expect.objectContaining({
@@ -3606,6 +3606,7 @@ describe("Neptune Analytics – backward compat with mocked client", () => {
profile: "dev-profile",
maxAttempts: 3,
});
await store.initialize();
// The client is now constructed lazily on first use, not in the constructor.
await store.search([1, 2, 3], 1);