fix(ts-sdk): lazy-load optional provider SDKs in mem0ai/oss (#6389)
This commit is contained in:
@@ -148,7 +148,7 @@ See the full catalog in <Link href="/components/llms/overview">Components</Link>
|
||||
- Qdrant connection errors: confirm port `6333` is exposed and the API key (if set) matches.
|
||||
- Empty search results: verify the embedder model name. A mismatch causes dimension errors.
|
||||
- `Unknown reranker` (Python): upgrade the SDK with `pip install --upgrade mem0ai` to load the latest provider registry.
|
||||
- `Cannot find module` (Node): import from the OSS entry point, `import { Memory } from "mem0ai/oss"`, not `"mem0ai"`.
|
||||
- `Cannot find module` (Node): two common causes. First, import from the OSS entry point, `import { Memory } from "mem0ai/oss"`, not `"mem0ai"`. Second, provider SDKs are optional peer dependencies loaded on demand, so install the one for the provider you configured (for example `npm install @qdrant/js-client-rest` for Qdrant). Installing `mem0ai` alone only pulls in the providers used by default; you do not need SDKs for providers you never select.
|
||||
|
||||
<CardGroup cols={2}>
|
||||
<Card
|
||||
|
||||
@@ -151,6 +151,42 @@
|
||||
"@aws-sdk/client-bedrock-runtime": ">=3.0.0 <3.968.0"
|
||||
},
|
||||
"peerDependenciesMeta": {
|
||||
"@qdrant/js-client-rest": {
|
||||
"optional": true
|
||||
},
|
||||
"redis": {
|
||||
"optional": true
|
||||
},
|
||||
"@supabase/supabase-js": {
|
||||
"optional": true
|
||||
},
|
||||
"cloudflare": {
|
||||
"optional": true
|
||||
},
|
||||
"@azure/search-documents": {
|
||||
"optional": true
|
||||
},
|
||||
"@azure/identity": {
|
||||
"optional": true
|
||||
},
|
||||
"@langchain/core": {
|
||||
"optional": true
|
||||
},
|
||||
"@anthropic-ai/sdk": {
|
||||
"optional": true
|
||||
},
|
||||
"@google/genai": {
|
||||
"optional": true
|
||||
},
|
||||
"groq-sdk": {
|
||||
"optional": true
|
||||
},
|
||||
"@mistralai/mistralai": {
|
||||
"optional": true
|
||||
},
|
||||
"ollama": {
|
||||
"optional": true
|
||||
},
|
||||
"mysql2": {
|
||||
"optional": true
|
||||
},
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
import type { FlagEmbedding } from "fastembed";
|
||||
import { Embedder } from "./base";
|
||||
import { EmbeddingConfig } from "../types";
|
||||
import { loadPeer } from "../utils/load_peer";
|
||||
|
||||
// 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
|
||||
@@ -54,14 +55,11 @@ export class FastEmbedEmbedder implements Embedder {
|
||||
* 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",
|
||||
const sdk = await loadPeer(
|
||||
"fastembed",
|
||||
"FastEmbed embedder",
|
||||
() => import("fastembed"),
|
||||
);
|
||||
}
|
||||
|
||||
return sdk.FlagEmbedding.init({ model: this.modelName });
|
||||
}
|
||||
|
||||
@@ -1,21 +1,32 @@
|
||||
import { GoogleGenAI } from "@google/genai";
|
||||
import type { GoogleGenAI } from "@google/genai";
|
||||
import { Embedder } from "./base";
|
||||
import { EmbeddingConfig } from "../types";
|
||||
import { loadPeer } from "../utils/load_peer";
|
||||
|
||||
export class GoogleEmbedder implements Embedder {
|
||||
private google: GoogleGenAI;
|
||||
private google!: GoogleGenAI;
|
||||
private model: string;
|
||||
private embeddingDims: number | undefined;
|
||||
private readonly apiKey: string | undefined;
|
||||
|
||||
constructor(config: EmbeddingConfig) {
|
||||
this.google = new GoogleGenAI({
|
||||
apiKey: config.apiKey || process.env.GOOGLE_API_KEY,
|
||||
});
|
||||
this.apiKey = config.apiKey || process.env.GOOGLE_API_KEY;
|
||||
this.model = config.model || "gemini-embedding-001";
|
||||
this.embeddingDims = config.embeddingDims;
|
||||
}
|
||||
|
||||
private async ensureClient(): Promise<void> {
|
||||
if (this.google) return;
|
||||
const sdk = await loadPeer(
|
||||
"@google/genai",
|
||||
"Google embedder",
|
||||
() => import("@google/genai"),
|
||||
);
|
||||
this.google = new sdk.GoogleGenAI({ apiKey: this.apiKey });
|
||||
}
|
||||
|
||||
async embed(text: string): Promise<number[]> {
|
||||
await this.ensureClient();
|
||||
const response = await this.google.models.embedContent({
|
||||
model: this.model,
|
||||
contents: text,
|
||||
@@ -27,6 +38,7 @@ export class GoogleEmbedder implements Embedder {
|
||||
}
|
||||
|
||||
async embedBatch(texts: string[]): Promise<number[][]> {
|
||||
await this.ensureClient();
|
||||
const response = await this.google.models.embedContent({
|
||||
model: this.model,
|
||||
contents: texts,
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
import { Embeddings } from "@langchain/core/embeddings";
|
||||
import type { Embeddings } from "@langchain/core/embeddings";
|
||||
import { Embedder } from "./base";
|
||||
import { EmbeddingConfig } from "../types";
|
||||
|
||||
|
||||
@@ -1,19 +1,19 @@
|
||||
import { Ollama } from "ollama";
|
||||
import type { Ollama } from "ollama";
|
||||
import { Embedder } from "./base";
|
||||
import { EmbeddingConfig } from "../types";
|
||||
import { logger } from "../utils/logger";
|
||||
import { loadPeer } from "../utils/load_peer";
|
||||
|
||||
export class OllamaEmbedder implements Embedder {
|
||||
private ollama: Ollama;
|
||||
private ollama!: Ollama;
|
||||
private model: string;
|
||||
private embeddingDims?: number;
|
||||
private readonly host: string;
|
||||
// Using this variable to avoid calling the Ollama server multiple times
|
||||
private initialized: boolean = false;
|
||||
|
||||
constructor(config: EmbeddingConfig) {
|
||||
this.ollama = new Ollama({
|
||||
host: config.url || config.baseURL || "http://localhost:11434",
|
||||
});
|
||||
this.host = config.url || config.baseURL || "http://localhost:11434";
|
||||
this.model = config.model || "nomic-embed-text:latest";
|
||||
this.embeddingDims = config.embeddingDims || 768;
|
||||
this.ensureModelExists().catch((err) => {
|
||||
@@ -21,7 +21,18 @@ export class OllamaEmbedder implements Embedder {
|
||||
});
|
||||
}
|
||||
|
||||
private async ensureClient(): Promise<void> {
|
||||
if (this.ollama) return;
|
||||
const sdk = await loadPeer(
|
||||
"ollama",
|
||||
"Ollama embedder",
|
||||
() => import("ollama"),
|
||||
);
|
||||
this.ollama = new sdk.Ollama({ host: this.host });
|
||||
}
|
||||
|
||||
async embed(text: string): Promise<number[]> {
|
||||
await this.ensureClient();
|
||||
try {
|
||||
await this.ensureModelExists();
|
||||
} catch (err) {
|
||||
@@ -54,6 +65,7 @@ export class OllamaEmbedder implements Embedder {
|
||||
if (this.initialized) {
|
||||
return true;
|
||||
}
|
||||
await this.ensureClient();
|
||||
const local_models = await this.ollama.list();
|
||||
const target = OllamaEmbedder.normalizeModelName(this.model);
|
||||
if (
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
import type { PredictionServiceClient } from "@google-cloud/aiplatform";
|
||||
import { Embedder } from "./base";
|
||||
import { VertexAIConfig } from "../types";
|
||||
import { loadPeer } from "../utils/load_peer";
|
||||
|
||||
type AIPlatform = typeof import("@google-cloud/aiplatform");
|
||||
type ClientOptions = NonNullable<
|
||||
@@ -105,15 +106,11 @@ export class VertexAIEmbedder implements Embedder {
|
||||
}
|
||||
|
||||
private async createClient(): Promise<void> {
|
||||
let aiplatform: AIPlatform;
|
||||
try {
|
||||
aiplatform = await import("@google-cloud/aiplatform");
|
||||
} catch (err) {
|
||||
throw new Error(
|
||||
"Failed to import '@google-cloud/aiplatform'. Please install it to use the Vertex AI embedding provider: " +
|
||||
(err as Error).message,
|
||||
const aiplatform: AIPlatform = await loadPeer(
|
||||
"@google-cloud/aiplatform",
|
||||
"Vertex AI embedding provider",
|
||||
() => import("@google-cloud/aiplatform"),
|
||||
);
|
||||
}
|
||||
|
||||
const client = new aiplatform.PredictionServiceClient(this.clientOptions);
|
||||
|
||||
|
||||
@@ -1,9 +1,11 @@
|
||||
import Anthropic from "@anthropic-ai/sdk";
|
||||
import type Anthropic from "@anthropic-ai/sdk";
|
||||
import { LLM, LLMResponse } from "./base";
|
||||
import { LLMConfig, Message } from "../types";
|
||||
import { loadPeer } from "../utils/load_peer";
|
||||
|
||||
export class AnthropicLLM implements LLM {
|
||||
private client: Anthropic;
|
||||
private client!: Anthropic;
|
||||
private readonly clientArgs: { apiKey: string; baseURL?: string };
|
||||
private model: string;
|
||||
private maxTokens: number;
|
||||
private temperature?: number;
|
||||
@@ -20,7 +22,7 @@ export class AnthropicLLM implements LLM {
|
||||
if (config.baseURL) {
|
||||
clientArgs.baseURL = config.baseURL;
|
||||
}
|
||||
this.client = new Anthropic(clientArgs);
|
||||
this.clientArgs = clientArgs;
|
||||
this.model = config.model || "claude-sonnet-4-6";
|
||||
// Defaults mirror the Python provider's AnthropicConfig
|
||||
// (max_tokens=2000, temperature=0.1, top_p omitted).
|
||||
@@ -29,11 +31,22 @@ export class AnthropicLLM implements LLM {
|
||||
this.topP = config.topP;
|
||||
}
|
||||
|
||||
private async ensureClient(): Promise<void> {
|
||||
if (this.client) return;
|
||||
const sdk = await loadPeer(
|
||||
"@anthropic-ai/sdk",
|
||||
"Anthropic LLM",
|
||||
() => import("@anthropic-ai/sdk"),
|
||||
);
|
||||
this.client = new sdk.default(this.clientArgs);
|
||||
}
|
||||
|
||||
async generateResponse(
|
||||
messages: Message[],
|
||||
responseFormat?: { type: string },
|
||||
tools?: any[],
|
||||
): Promise<string | LLMResponse> {
|
||||
await this.ensureClient();
|
||||
// Extract system message if present
|
||||
const systemMessage = messages.find((msg) => msg.role === "system");
|
||||
const otherMessages = messages.filter((msg) => msg.role !== "system");
|
||||
|
||||
@@ -1,16 +1,28 @@
|
||||
import { GoogleGenAI } from "@google/genai";
|
||||
import type { GoogleGenAI } from "@google/genai";
|
||||
import { LLM, LLMResponse } from "./base";
|
||||
import { LLMConfig, Message } from "../types";
|
||||
import { loadPeer } from "../utils/load_peer";
|
||||
|
||||
export class GoogleLLM implements LLM {
|
||||
private google: GoogleGenAI;
|
||||
private google!: GoogleGenAI;
|
||||
private model: string;
|
||||
private readonly apiKey: string | undefined;
|
||||
|
||||
constructor(config: LLMConfig) {
|
||||
this.google = new GoogleGenAI({ apiKey: config.apiKey });
|
||||
this.apiKey = config.apiKey;
|
||||
this.model = config.model || "gemini-2.0-flash";
|
||||
}
|
||||
|
||||
private async ensureClient(): Promise<void> {
|
||||
if (this.google) return;
|
||||
const sdk = await loadPeer(
|
||||
"@google/genai",
|
||||
"Google LLM",
|
||||
() => import("@google/genai"),
|
||||
);
|
||||
this.google = new sdk.GoogleGenAI({ apiKey: this.apiKey });
|
||||
}
|
||||
|
||||
private formatContents(messages: Message[]) {
|
||||
return messages.map((msg) => ({
|
||||
parts: [
|
||||
@@ -30,6 +42,7 @@ export class GoogleLLM implements LLM {
|
||||
responseFormat?: { type: string },
|
||||
tools?: any[],
|
||||
): Promise<string | LLMResponse> {
|
||||
await this.ensureClient();
|
||||
const contents = this.formatContents(messages);
|
||||
|
||||
// Build config with tools if provided
|
||||
@@ -72,6 +85,7 @@ export class GoogleLLM implements LLM {
|
||||
}
|
||||
|
||||
async generateChat(messages: Message[]): Promise<LLMResponse> {
|
||||
await this.ensureClient();
|
||||
const completion = await this.google.models.generateContent({
|
||||
contents: this.formatContents(messages),
|
||||
model: this.model,
|
||||
|
||||
@@ -1,24 +1,37 @@
|
||||
import { Groq } from "groq-sdk";
|
||||
import type { Groq } from "groq-sdk";
|
||||
import { LLM, LLMResponse } from "./base";
|
||||
import { LLMConfig, Message } from "../types";
|
||||
import { loadPeer } from "../utils/load_peer";
|
||||
|
||||
export class GroqLLM implements LLM {
|
||||
private client: Groq;
|
||||
private client!: Groq;
|
||||
private model: string;
|
||||
private readonly apiKey: string;
|
||||
|
||||
constructor(config: LLMConfig) {
|
||||
const apiKey = config.apiKey || process.env.GROQ_API_KEY;
|
||||
if (!apiKey) {
|
||||
throw new Error("Groq API key is required");
|
||||
}
|
||||
this.client = new Groq({ apiKey });
|
||||
this.apiKey = apiKey;
|
||||
this.model = config.model || "llama3-70b-8192";
|
||||
}
|
||||
|
||||
private async ensureClient(): Promise<void> {
|
||||
if (this.client) return;
|
||||
const sdk = await loadPeer(
|
||||
"groq-sdk",
|
||||
"Groq LLM",
|
||||
() => import("groq-sdk"),
|
||||
);
|
||||
this.client = new sdk.Groq({ apiKey: this.apiKey });
|
||||
}
|
||||
|
||||
async generateResponse(
|
||||
messages: Message[],
|
||||
responseFormat?: { type: string },
|
||||
): Promise<string> {
|
||||
await this.ensureClient();
|
||||
const response = await this.client.chat.completions.create({
|
||||
model: this.model,
|
||||
messages: messages.map((msg) => ({
|
||||
@@ -35,6 +48,7 @@ export class GroqLLM implements LLM {
|
||||
}
|
||||
|
||||
async generateChat(messages: Message[]): Promise<LLMResponse> {
|
||||
await this.ensureClient();
|
||||
const response = await this.client.chat.completions.create({
|
||||
model: this.model,
|
||||
messages: messages.map((msg) => ({
|
||||
|
||||
@@ -1,17 +1,16 @@
|
||||
import { BaseLanguageModel } from "@langchain/core/language_models/base";
|
||||
import {
|
||||
AIMessage,
|
||||
HumanMessage,
|
||||
SystemMessage,
|
||||
BaseMessage,
|
||||
} from "@langchain/core/messages";
|
||||
import type { BaseLanguageModel } from "@langchain/core/language_models/base";
|
||||
import type { BaseMessage } from "@langchain/core/messages";
|
||||
import { z } from "zod";
|
||||
import { LLM, LLMResponse } from "./base";
|
||||
import { LLMConfig, Message } from "../types/index";
|
||||
// Import the schemas directly into LangchainLLM
|
||||
import { FactRetrievalSchema, MemoryUpdateSchema } from "../prompts";
|
||||
|
||||
const convertToLangchainMessages = (messages: Message[]): BaseMessage[] => {
|
||||
const convertToLangchainMessages = async (
|
||||
messages: Message[],
|
||||
): Promise<BaseMessage[]> => {
|
||||
const { AIMessage, HumanMessage, SystemMessage } =
|
||||
await import("@langchain/core/messages");
|
||||
return messages.map((msg) => {
|
||||
const content =
|
||||
typeof msg.content === "string"
|
||||
@@ -62,7 +61,7 @@ export class LangchainLLM implements LLM {
|
||||
response_format?: { type: string },
|
||||
tools?: any[],
|
||||
): Promise<string | LLMResponse> {
|
||||
const langchainMessages = convertToLangchainMessages(messages);
|
||||
const langchainMessages = await convertToLangchainMessages(messages);
|
||||
let runnable: any = this.llmInstance;
|
||||
const invokeOptions: Record<string, any> = {};
|
||||
let isStructuredOutput = false;
|
||||
@@ -170,7 +169,7 @@ export class LangchainLLM implements LLM {
|
||||
}
|
||||
|
||||
async generateChat(messages: Message[]): Promise<LLMResponse> {
|
||||
const langchainMessages = convertToLangchainMessages(messages);
|
||||
const langchainMessages = await convertToLangchainMessages(messages);
|
||||
try {
|
||||
const response = await this.llmInstance.invoke(langchainMessages);
|
||||
if (response && typeof response.content === "string") {
|
||||
|
||||
@@ -1,21 +1,31 @@
|
||||
import { Mistral } from "@mistralai/mistralai";
|
||||
import type { Mistral } from "@mistralai/mistralai";
|
||||
import { LLM, LLMResponse } from "./base";
|
||||
import { LLMConfig, Message } from "../types";
|
||||
import { loadPeer } from "../utils/load_peer";
|
||||
|
||||
export class MistralLLM implements LLM {
|
||||
private client: Mistral;
|
||||
private client!: Mistral;
|
||||
private model: string;
|
||||
private readonly apiKey: string;
|
||||
|
||||
constructor(config: LLMConfig) {
|
||||
if (!config.apiKey) {
|
||||
throw new Error("Mistral API key is required");
|
||||
}
|
||||
this.client = new Mistral({
|
||||
apiKey: config.apiKey,
|
||||
});
|
||||
this.apiKey = config.apiKey;
|
||||
this.model = config.model || "mistral-tiny-latest";
|
||||
}
|
||||
|
||||
private async ensureClient(): Promise<void> {
|
||||
if (this.client) return;
|
||||
const sdk = await loadPeer(
|
||||
"@mistralai/mistralai",
|
||||
"Mistral LLM",
|
||||
() => import("@mistralai/mistralai"),
|
||||
);
|
||||
this.client = new sdk.Mistral({ apiKey: this.apiKey });
|
||||
}
|
||||
|
||||
// Helper function to convert content to string
|
||||
private contentToString(content: any): string {
|
||||
if (typeof content === "string") {
|
||||
@@ -41,6 +51,7 @@ export class MistralLLM implements LLM {
|
||||
responseFormat?: { type: string },
|
||||
tools?: any[],
|
||||
): Promise<string | LLMResponse> {
|
||||
await this.ensureClient();
|
||||
const response = await this.client.chat.complete({
|
||||
model: this.model,
|
||||
messages: messages.map((msg) => ({
|
||||
@@ -82,6 +93,7 @@ export class MistralLLM implements LLM {
|
||||
}
|
||||
|
||||
async generateChat(messages: Message[]): Promise<LLMResponse> {
|
||||
await this.ensureClient();
|
||||
const formattedMessages = messages.map((msg) => ({
|
||||
role: msg.role as "system" | "user" | "assistant",
|
||||
content:
|
||||
|
||||
@@ -1,29 +1,36 @@
|
||||
import { Ollama } from "ollama";
|
||||
import type { Ollama } from "ollama";
|
||||
import { LLM, LLMResponse } from "./base";
|
||||
import { LLMConfig, Message } from "../types";
|
||||
import { logger } from "../utils/logger";
|
||||
import { loadPeer } from "../utils/load_peer";
|
||||
|
||||
export class OllamaLLM implements LLM {
|
||||
private ollama: Ollama;
|
||||
private ollama!: Ollama;
|
||||
private model: string;
|
||||
private readonly host: string;
|
||||
// Using this variable to avoid calling the Ollama server multiple times
|
||||
private initialized: boolean = false;
|
||||
|
||||
constructor(config: LLMConfig) {
|
||||
this.ollama = new Ollama({
|
||||
host: config.url || config.baseURL || "http://localhost:11434",
|
||||
});
|
||||
this.host = config.url || config.baseURL || "http://localhost:11434";
|
||||
this.model = config.model || "llama3.1:8b";
|
||||
this.ensureModelExists().catch((err) => {
|
||||
logger.error(`Error ensuring model exists: ${err}`);
|
||||
});
|
||||
}
|
||||
|
||||
private async ensureClient(): Promise<void> {
|
||||
if (this.ollama) return;
|
||||
const sdk = await loadPeer("ollama", "Ollama LLM", () => import("ollama"));
|
||||
this.ollama = new sdk.Ollama({ host: this.host });
|
||||
}
|
||||
|
||||
async generateResponse(
|
||||
messages: Message[],
|
||||
responseFormat?: { type: string },
|
||||
tools?: any[],
|
||||
): Promise<string | LLMResponse> {
|
||||
await this.ensureClient();
|
||||
try {
|
||||
await this.ensureModelExists();
|
||||
} catch (err) {
|
||||
@@ -63,6 +70,7 @@ export class OllamaLLM implements LLM {
|
||||
}
|
||||
|
||||
async generateChat(messages: Message[]): Promise<LLMResponse> {
|
||||
await this.ensureClient();
|
||||
try {
|
||||
await this.ensureModelExists();
|
||||
} catch (err) {
|
||||
@@ -93,6 +101,7 @@ export class OllamaLLM implements LLM {
|
||||
if (this.initialized) {
|
||||
return true;
|
||||
}
|
||||
await this.ensureClient();
|
||||
const local_models = await this.ollama.list();
|
||||
if (!local_models.models.find((m: any) => m.name === this.model)) {
|
||||
logger.info(`Pulling model ${this.model}...`);
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
import { RerankerConfig } from "../types";
|
||||
import { Reranker, RerankResult } from "./base";
|
||||
import { loadPeer } from "../utils/load_peer";
|
||||
|
||||
const DEFAULT_MODEL = "rerank-v3.5";
|
||||
|
||||
@@ -41,14 +42,11 @@ export class CohereReranker implements Reranker {
|
||||
}
|
||||
|
||||
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",
|
||||
const sdk = await loadPeer(
|
||||
"cohere-ai",
|
||||
"Cohere reranker",
|
||||
() => import("cohere-ai"),
|
||||
);
|
||||
}
|
||||
return new sdk.CohereClient({ token: this.apiKey });
|
||||
}
|
||||
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
import { RerankerConfig } from "../types";
|
||||
import { Reranker, RerankResult } from "./base";
|
||||
import { loadPeer } from "../utils/load_peer";
|
||||
|
||||
const DEFAULT_MODEL = "zerank-1";
|
||||
|
||||
@@ -37,14 +38,11 @@ export class ZeroEntropyReranker implements Reranker {
|
||||
}
|
||||
|
||||
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",
|
||||
const sdk = await loadPeer(
|
||||
"zeroentropy",
|
||||
"ZeroEntropy reranker",
|
||||
() => import("zeroentropy"),
|
||||
);
|
||||
}
|
||||
return new sdk.ZeroEntropy({ apiKey: this.apiKey });
|
||||
}
|
||||
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
import { createClient, SupabaseClient } from "@supabase/supabase-js";
|
||||
import type { SupabaseClient } from "@supabase/supabase-js";
|
||||
import { v4 as uuidv4 } from "uuid";
|
||||
import { HistoryManager } from "./base";
|
||||
import { loadPeer } from "../utils/load_peer";
|
||||
|
||||
interface HistoryEntry {
|
||||
id: string;
|
||||
@@ -20,16 +21,32 @@ interface SupabaseHistoryConfig {
|
||||
}
|
||||
|
||||
export class SupabaseHistoryManager implements HistoryManager {
|
||||
private supabase: SupabaseClient;
|
||||
// ponytail: benign double-construct race — two concurrent first-calls may each
|
||||
// build a client; createClient opens no connection, so last-write-wins is fine.
|
||||
private supabase!: SupabaseClient;
|
||||
private readonly supabaseUrl: string;
|
||||
private readonly supabaseKey: string;
|
||||
private readonly tableName: string;
|
||||
|
||||
constructor(config: SupabaseHistoryConfig) {
|
||||
this.tableName = config.tableName || "memory_history";
|
||||
this.supabase = createClient(config.supabaseUrl, config.supabaseKey);
|
||||
this.supabaseUrl = config.supabaseUrl;
|
||||
this.supabaseKey = config.supabaseKey;
|
||||
this.initializeSupabase().catch(console.error);
|
||||
}
|
||||
|
||||
private async ensureClient(): Promise<void> {
|
||||
if (this.supabase) return;
|
||||
const sdk = await loadPeer(
|
||||
"@supabase/supabase-js",
|
||||
"Supabase history manager",
|
||||
() => import("@supabase/supabase-js"),
|
||||
);
|
||||
this.supabase = sdk.createClient(this.supabaseUrl, this.supabaseKey);
|
||||
}
|
||||
|
||||
private async initializeSupabase(): Promise<void> {
|
||||
await this.ensureClient();
|
||||
// Check if table exists
|
||||
const { error } = await this.supabase
|
||||
.from(this.tableName)
|
||||
@@ -65,6 +82,7 @@ create table ${this.tableName} (
|
||||
updatedAt?: string,
|
||||
isDeleted: number = 0,
|
||||
): Promise<void> {
|
||||
await this.ensureClient();
|
||||
const historyEntry: HistoryEntry = {
|
||||
id: uuidv4(),
|
||||
memory_id: memoryId,
|
||||
@@ -87,6 +105,7 @@ create table ${this.tableName} (
|
||||
}
|
||||
|
||||
async getHistory(memoryId: string): Promise<any[]> {
|
||||
await this.ensureClient();
|
||||
const { data, error } = await this.supabase
|
||||
.from(this.tableName)
|
||||
.select("*")
|
||||
@@ -103,6 +122,7 @@ create table ${this.tableName} (
|
||||
}
|
||||
|
||||
async reset(): Promise<void> {
|
||||
await this.ensureClient();
|
||||
const { error } = await this.supabase
|
||||
.from(this.tableName)
|
||||
.delete()
|
||||
|
||||
@@ -0,0 +1,13 @@
|
||||
export async function loadPeer(
|
||||
pkg: string,
|
||||
label: string,
|
||||
load: () => Promise<any>,
|
||||
): Promise<any> {
|
||||
try {
|
||||
return await load();
|
||||
} catch {
|
||||
throw new Error(
|
||||
`The '${pkg}' package is required to use the ${label}. Install it with: npm install ${pkg}`,
|
||||
);
|
||||
}
|
||||
}
|
||||
@@ -1,7 +1,6 @@
|
||||
import {
|
||||
import type {
|
||||
SearchClient,
|
||||
SearchIndexClient,
|
||||
AzureKeyCredential,
|
||||
SearchIndex,
|
||||
SearchField,
|
||||
SearchFieldDataType,
|
||||
@@ -13,9 +12,9 @@ import {
|
||||
BinaryQuantizationCompression,
|
||||
VectorizedQuery,
|
||||
} from "@azure/search-documents";
|
||||
import { DefaultAzureCredential } from "@azure/identity";
|
||||
import { VectorStore } from "./base";
|
||||
import { SearchFilters, VectorStoreConfig, VectorStoreResult } from "../types";
|
||||
import { loadPeer } from "../utils/load_peer";
|
||||
|
||||
/**
|
||||
* Configuration interface for Azure AI Search vector store
|
||||
@@ -71,8 +70,8 @@ interface AzureAISearchConfig extends VectorStoreConfig {
|
||||
* Supports vector search with hybrid search, compression, and filtering
|
||||
*/
|
||||
export class AzureAISearch implements VectorStore {
|
||||
private searchClient: SearchClient<any>;
|
||||
private indexClient: SearchIndexClient;
|
||||
private searchClient!: SearchClient<any>;
|
||||
private indexClient!: SearchIndexClient;
|
||||
private readonly serviceName: string;
|
||||
private readonly indexName: string;
|
||||
private readonly embeddingModelDims: number;
|
||||
@@ -93,25 +92,42 @@ export class AzureAISearch implements VectorStore {
|
||||
this.vectorFilterMode = config.vectorFilterMode || "preFilter";
|
||||
this.apiKey = config.apiKey;
|
||||
|
||||
// Initialize the index
|
||||
this.initialize().catch(console.error);
|
||||
}
|
||||
|
||||
private async ensureClient(): Promise<void> {
|
||||
if (this.searchClient) return;
|
||||
const searchSdk = await loadPeer(
|
||||
"@azure/search-documents",
|
||||
"Azure AI Search vector store",
|
||||
() => import("@azure/search-documents"),
|
||||
);
|
||||
|
||||
const serviceEndpoint = `https://${this.serviceName}.search.windows.net`;
|
||||
|
||||
// Determine authentication: API key or DefaultAzureCredential
|
||||
const credential =
|
||||
this.apiKey && this.apiKey !== "" && this.apiKey !== "your-api-key"
|
||||
? new AzureKeyCredential(this.apiKey)
|
||||
: new DefaultAzureCredential();
|
||||
let credential: any;
|
||||
if (this.apiKey && this.apiKey !== "" && this.apiKey !== "your-api-key") {
|
||||
credential = new searchSdk.AzureKeyCredential(this.apiKey);
|
||||
} else {
|
||||
const identitySdk = await loadPeer(
|
||||
"@azure/identity",
|
||||
"Azure AI Search without an apiKey",
|
||||
() => import("@azure/identity"),
|
||||
);
|
||||
credential = new identitySdk.DefaultAzureCredential();
|
||||
}
|
||||
|
||||
// Initialize clients
|
||||
this.searchClient = new SearchClient(
|
||||
this.searchClient = new searchSdk.SearchClient(
|
||||
serviceEndpoint,
|
||||
this.indexName,
|
||||
credential,
|
||||
);
|
||||
|
||||
this.indexClient = new SearchIndexClient(serviceEndpoint, credential);
|
||||
|
||||
// Initialize the index
|
||||
this.initialize().catch(console.error);
|
||||
this.indexClient = new searchSdk.SearchIndexClient(
|
||||
serviceEndpoint,
|
||||
credential,
|
||||
);
|
||||
}
|
||||
|
||||
/**
|
||||
@@ -125,6 +141,7 @@ export class AzureAISearch implements VectorStore {
|
||||
}
|
||||
|
||||
private async _doInitialize(): Promise<void> {
|
||||
await this.ensureClient();
|
||||
try {
|
||||
const collections = await this.listCols();
|
||||
if (!collections.includes(this.indexName)) {
|
||||
@@ -262,6 +279,7 @@ export class AzureAISearch implements VectorStore {
|
||||
ids: string[],
|
||||
payloads: Record<string, any>[],
|
||||
): Promise<void> {
|
||||
await this.initialize();
|
||||
console.log(
|
||||
`Inserting ${vectors.length} vectors into index ${this.indexName}`,
|
||||
);
|
||||
@@ -334,6 +352,7 @@ export class AzureAISearch implements VectorStore {
|
||||
topK: number = 5,
|
||||
filters?: SearchFilters,
|
||||
): Promise<VectorStoreResult[] | null> {
|
||||
await this.initialize();
|
||||
try {
|
||||
const filterExpression = filters
|
||||
? this.buildFilterExpression(filters)
|
||||
@@ -373,6 +392,7 @@ export class AzureAISearch implements VectorStore {
|
||||
topK: number = 5,
|
||||
filters?: SearchFilters,
|
||||
): Promise<VectorStoreResult[]> {
|
||||
await this.initialize();
|
||||
const filterExpression = filters
|
||||
? this.buildFilterExpression(filters)
|
||||
: undefined;
|
||||
@@ -429,6 +449,7 @@ export class AzureAISearch implements VectorStore {
|
||||
* Delete a vector by ID
|
||||
*/
|
||||
async delete(vectorId: string): Promise<void> {
|
||||
await this.initialize();
|
||||
const response = await this.searchClient.deleteDocuments([
|
||||
{ id: vectorId },
|
||||
]);
|
||||
@@ -454,6 +475,7 @@ export class AzureAISearch implements VectorStore {
|
||||
vector: number[],
|
||||
payload: Record<string, any>,
|
||||
): Promise<void> {
|
||||
await this.initialize();
|
||||
const document: Record<string, any> = { id: vectorId };
|
||||
|
||||
if (vector) {
|
||||
@@ -486,6 +508,7 @@ export class AzureAISearch implements VectorStore {
|
||||
* Retrieve a vector by ID
|
||||
*/
|
||||
async get(vectorId: string): Promise<VectorStoreResult | null> {
|
||||
await this.initialize();
|
||||
try {
|
||||
const result = await this.searchClient.getDocument(vectorId);
|
||||
const payloadStr = result.payload as string;
|
||||
@@ -521,6 +544,7 @@ export class AzureAISearch implements VectorStore {
|
||||
* Delete the index
|
||||
*/
|
||||
async deleteCol(): Promise<void> {
|
||||
await this.initialize();
|
||||
await this.indexClient.deleteIndex(this.indexName);
|
||||
}
|
||||
|
||||
@@ -542,6 +566,7 @@ export class AzureAISearch implements VectorStore {
|
||||
filters?: SearchFilters,
|
||||
topK: number = 100,
|
||||
): Promise<[VectorStoreResult[], number]> {
|
||||
await this.initialize();
|
||||
const filterExpression = filters
|
||||
? this.buildFilterExpression(filters)
|
||||
: undefined;
|
||||
@@ -586,6 +611,7 @@ export class AzureAISearch implements VectorStore {
|
||||
* Required by VectorStore interface
|
||||
*/
|
||||
async getUserId(): Promise<string> {
|
||||
await this.initialize();
|
||||
try {
|
||||
// Check if memory_migrations index exists
|
||||
const collections = await this.listCols();
|
||||
@@ -648,6 +674,7 @@ export class AzureAISearch implements VectorStore {
|
||||
* Required by VectorStore interface
|
||||
*/
|
||||
async setUserId(userId: string): Promise<void> {
|
||||
await this.initialize();
|
||||
try {
|
||||
// Get existing point ID or generate new one
|
||||
const searchResults = await this.searchClient.search("*", {
|
||||
@@ -677,6 +704,7 @@ export class AzureAISearch implements VectorStore {
|
||||
* Reset the index by deleting and recreating it
|
||||
*/
|
||||
async reset(): Promise<void> {
|
||||
await this.initialize();
|
||||
console.log(`Resetting index ${this.indexName}...`);
|
||||
|
||||
try {
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
import type { Pool, RowDataPacket } from "mysql2/promise";
|
||||
import { VectorStore } from "./base";
|
||||
import { SearchFilters, VectorStoreConfig, VectorStoreResult } from "../types";
|
||||
import { loadPeer } from "../utils/load_peer";
|
||||
|
||||
const SAFE_IDENTIFIER_RE = /^[a-zA-Z_][a-zA-Z0-9_]{0,127}$/;
|
||||
|
||||
@@ -94,14 +95,11 @@ export class AzureMySQLDB implements VectorStore {
|
||||
|
||||
// Loaded dynamically: mysql2 is an optional peer dependency, so a static value import
|
||||
// would break `import { Memory } from "mem0ai/oss"` for everyone else.
|
||||
let createPool: typeof import("mysql2/promise").createPool;
|
||||
try {
|
||||
({ createPool } = await import("mysql2/promise"));
|
||||
} catch {
|
||||
throw new Error(
|
||||
"The Azure MySQL vector store requires the 'mysql2' package. Install it with: npm install mysql2",
|
||||
const { createPool }: typeof import("mysql2/promise") = await loadPeer(
|
||||
"mysql2",
|
||||
"Azure MySQL vector store",
|
||||
() => import("mysql2/promise"),
|
||||
);
|
||||
}
|
||||
|
||||
this.pool = createPool({
|
||||
host: this.config.host,
|
||||
|
||||
@@ -12,6 +12,7 @@ import type {
|
||||
} from "@mochow/mochow-sdk-node";
|
||||
import { VectorStore } from "./base";
|
||||
import { SearchFilters, VectorStoreConfig, VectorStoreResult } from "../types";
|
||||
import { loadPeer } from "../utils/load_peer";
|
||||
|
||||
type MochowSdk = typeof import("@mochow/mochow-sdk-node");
|
||||
|
||||
@@ -156,14 +157,11 @@ export class BaiduDB implements VectorStore {
|
||||
// value import would break `import { Memory } from "mem0ai/oss"` for everyone else.
|
||||
private async loadSdk(): Promise<MochowSdk> {
|
||||
if (!this.sdk) {
|
||||
let module: MochowSdk & { default?: MochowSdk };
|
||||
try {
|
||||
module = await import("@mochow/mochow-sdk-node");
|
||||
} catch {
|
||||
throw new Error(
|
||||
"The Baidu vector store requires the '@mochow/mochow-sdk-node' package. Install it with: npm install @mochow/mochow-sdk-node",
|
||||
const module: MochowSdk & { default?: MochowSdk } = await loadPeer(
|
||||
"@mochow/mochow-sdk-node",
|
||||
"Baidu vector store",
|
||||
() => import("@mochow/mochow-sdk-node"),
|
||||
);
|
||||
}
|
||||
this.sdk = module.default ?? module;
|
||||
}
|
||||
return this.sdk;
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
import { VectorStore } from "./base";
|
||||
import { SearchFilters, VectorStoreConfig, VectorStoreResult } from "../types";
|
||||
import { loadPeer } from "../utils/load_peer";
|
||||
|
||||
const MIGRATION_ROW_ID = "mem0-user";
|
||||
const SAFE_IDENTIFIER_RE = /^[A-Za-z_][A-Za-z0-9_]{0,127}$/;
|
||||
@@ -377,14 +378,11 @@ export class CassandraDB implements VectorStore {
|
||||
// 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",
|
||||
const sdk = await loadPeer(
|
||||
"cassandra-driver",
|
||||
"Cassandra vector store",
|
||||
() => import("cassandra-driver"),
|
||||
);
|
||||
}
|
||||
return sdk.default ?? sdk;
|
||||
}
|
||||
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
import type { ChromaClient, CloudClient } from "chromadb";
|
||||
import { VectorStore } from "./base";
|
||||
import { SearchFilters, VectorStoreConfig, VectorStoreResult } from "../types";
|
||||
import { loadPeer } from "../utils/load_peer";
|
||||
|
||||
interface ChromaConfig extends VectorStoreConfig {
|
||||
/** Pre-configured ChromaDB client instance. */
|
||||
@@ -65,14 +66,11 @@ export class ChromaDB implements VectorStore {
|
||||
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",
|
||||
const sdk = await loadPeer(
|
||||
"chromadb",
|
||||
"Chroma vector store",
|
||||
() => import("chromadb"),
|
||||
);
|
||||
}
|
||||
|
||||
if (config.apiKey && config.tenant) {
|
||||
return new sdk.CloudClient({
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
import type { Client } from "@elastic/elasticsearch";
|
||||
import { VectorStore } from "./base";
|
||||
import { SearchFilters, VectorStoreConfig, VectorStoreResult } from "../types";
|
||||
import { loadPeer } from "../utils/load_peer";
|
||||
|
||||
interface ElasticsearchConfig extends VectorStoreConfig {
|
||||
/** Pre-configured Elasticsearch client instance (typed as `any` to keep the
|
||||
@@ -100,14 +101,11 @@ export class ElasticsearchDB implements VectorStore {
|
||||
params.headers = config.headers;
|
||||
}
|
||||
|
||||
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",
|
||||
const sdk = await loadPeer(
|
||||
"@elastic/elasticsearch",
|
||||
"Elasticsearch vector store",
|
||||
() => import("@elastic/elasticsearch"),
|
||||
);
|
||||
}
|
||||
|
||||
this.client = new sdk.Client(params);
|
||||
}
|
||||
|
||||
@@ -1,5 +1,4 @@
|
||||
import { VectorStore as LangchainVectorStoreInterface } from "@langchain/core/vectorstores";
|
||||
import { Document } from "@langchain/core/documents";
|
||||
import type { VectorStore as LangchainVectorStoreInterface } from "@langchain/core/vectorstores";
|
||||
import { VectorStore } from "./base"; // mem0's VectorStore interface
|
||||
import { SearchFilters, VectorStoreConfig, VectorStoreResult } from "../types";
|
||||
|
||||
@@ -77,6 +76,7 @@ export class LangchainVectorStore implements VectorStore {
|
||||
}
|
||||
|
||||
// Convert payloads to Langchain Document metadata format
|
||||
const { Document } = await import("@langchain/core/documents");
|
||||
const documents = payloads.map((payload, i) => {
|
||||
// Provide empty pageContent, store mem0 id and other data in metadata
|
||||
return new Document({
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
import type { MongoClient, Collection, Db } from "mongodb";
|
||||
import { VectorStore } from "./base";
|
||||
import { SearchFilters, VectorStoreConfig, VectorStoreResult } from "../types";
|
||||
import { loadPeer } from "../utils/load_peer";
|
||||
|
||||
export interface MongoDBConfig extends VectorStoreConfig {
|
||||
url?: string;
|
||||
@@ -42,14 +43,11 @@ export class MongoDB implements VectorStore {
|
||||
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 sdk = await loadPeer(
|
||||
"mongodb",
|
||||
"MongoDB vector store",
|
||||
() => import("mongodb"),
|
||||
);
|
||||
}
|
||||
const url = config.url || "mongodb://localhost:27017";
|
||||
this.client = new sdk.MongoClient(url, { appName: "Mem0" });
|
||||
}
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
import type { Client } from "@opensearch-project/opensearch";
|
||||
import { VectorStore } from "./base";
|
||||
import { SearchFilters, VectorStoreConfig, VectorStoreResult } from "../types";
|
||||
import { loadPeer } from "../utils/load_peer";
|
||||
|
||||
type OpenSearchAuth =
|
||||
| {
|
||||
@@ -98,14 +99,11 @@ export class OpenSearchDB implements VectorStore {
|
||||
? { username: config.user, password: config.password }
|
||||
: undefined);
|
||||
|
||||
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",
|
||||
const sdk = await loadPeer(
|
||||
"@opensearch-project/opensearch",
|
||||
"OpenSearch vector store",
|
||||
() => import("@opensearch-project/opensearch"),
|
||||
);
|
||||
}
|
||||
|
||||
this.client = new sdk.Client({
|
||||
node: `${useSSL ? "https" : "http"}://${host}:${port}`,
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
import type { Pinecone, Index } from "@pinecone-database/pinecone";
|
||||
import { VectorStore } from "./base";
|
||||
import { SearchFilters, VectorStoreConfig, VectorStoreResult } from "../types";
|
||||
import { loadPeer } from "../utils/load_peer";
|
||||
|
||||
const MIGRATIONS_NAMESPACE = "__mem0_migrations__";
|
||||
const MIGRATIONS_RECORD_ID = "mem0-user-id";
|
||||
@@ -82,14 +83,11 @@ export class PineconeDB implements VectorStore {
|
||||
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",
|
||||
const sdk = await loadPeer(
|
||||
"@pinecone-database/pinecone",
|
||||
"Pinecone vector store",
|
||||
() => import("@pinecone-database/pinecone"),
|
||||
);
|
||||
}
|
||||
this.client = new sdk.Pinecone({ apiKey });
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
import { QdrantClient } from "@qdrant/js-client-rest";
|
||||
import type { QdrantClient } from "@qdrant/js-client-rest";
|
||||
import { VectorStore } from "./base";
|
||||
import { SearchFilters, VectorStoreConfig, VectorStoreResult } from "../types";
|
||||
import { loadPeer } from "../utils/load_peer";
|
||||
import * as fs from "fs";
|
||||
|
||||
interface QdrantConfig extends VectorStoreConfig {
|
||||
@@ -55,15 +56,26 @@ const KEY_MAP: Record<string, string> = {
|
||||
};
|
||||
|
||||
export class Qdrant implements VectorStore {
|
||||
private client: QdrantClient;
|
||||
private client!: QdrantClient;
|
||||
private readonly config: QdrantConfig;
|
||||
private readonly collectionName: string;
|
||||
private dimension: number;
|
||||
private _initPromise?: Promise<void>;
|
||||
|
||||
constructor(config: QdrantConfig) {
|
||||
this.config = config;
|
||||
this.collectionName = config.collectionName;
|
||||
this.dimension = config.dimension || 1536; // Default OpenAI dimension
|
||||
this.initialize().catch(console.error);
|
||||
}
|
||||
|
||||
private async ensureClient(): Promise<void> {
|
||||
if (this.client) return;
|
||||
const config = this.config;
|
||||
if (config.client) {
|
||||
this.client = config.client;
|
||||
} else {
|
||||
return;
|
||||
}
|
||||
const params: Record<string, any> = {};
|
||||
if (config.apiKey) {
|
||||
params.apiKey = config.apiKey;
|
||||
@@ -94,12 +106,12 @@ export class Qdrant implements VectorStore {
|
||||
}
|
||||
}
|
||||
|
||||
this.client = new QdrantClient(params);
|
||||
}
|
||||
|
||||
this.collectionName = config.collectionName;
|
||||
this.dimension = config.dimension || 1536; // Default OpenAI dimension
|
||||
this.initialize().catch(console.error);
|
||||
const sdk = await loadPeer(
|
||||
"@qdrant/js-client-rest",
|
||||
"Qdrant vector store",
|
||||
() => import("@qdrant/js-client-rest"),
|
||||
);
|
||||
this.client = new sdk.QdrantClient(params);
|
||||
}
|
||||
|
||||
/**
|
||||
@@ -270,6 +282,7 @@ export class Qdrant implements VectorStore {
|
||||
ids: string[],
|
||||
payloads: Record<string, any>[],
|
||||
): Promise<void> {
|
||||
await this.initialize();
|
||||
const points = vectors.map((vector, idx) => ({
|
||||
id: ids[idx],
|
||||
vector: vector,
|
||||
@@ -290,6 +303,7 @@ export class Qdrant implements VectorStore {
|
||||
topK: number = 5,
|
||||
filters?: SearchFilters,
|
||||
): Promise<VectorStoreResult[]> {
|
||||
await this.initialize();
|
||||
const queryFilter = this.createFilter(filters);
|
||||
const results = await this.client.search(this.collectionName, {
|
||||
vector: query,
|
||||
@@ -305,6 +319,7 @@ export class Qdrant implements VectorStore {
|
||||
}
|
||||
|
||||
async get(vectorId: string): Promise<VectorStoreResult | null> {
|
||||
await this.initialize();
|
||||
const results = await this.client.retrieve(this.collectionName, {
|
||||
ids: [vectorId],
|
||||
with_payload: true,
|
||||
@@ -323,6 +338,7 @@ export class Qdrant implements VectorStore {
|
||||
vector: number[],
|
||||
payload: Record<string, any>,
|
||||
): Promise<void> {
|
||||
await this.initialize();
|
||||
const point = {
|
||||
id: vectorId,
|
||||
vector: vector,
|
||||
@@ -335,12 +351,14 @@ export class Qdrant implements VectorStore {
|
||||
}
|
||||
|
||||
async delete(vectorId: string): Promise<void> {
|
||||
await this.initialize();
|
||||
await this.client.delete(this.collectionName, {
|
||||
points: [vectorId],
|
||||
});
|
||||
}
|
||||
|
||||
async deleteCol(): Promise<void> {
|
||||
await this.initialize();
|
||||
await this.client.deleteCollection(this.collectionName);
|
||||
}
|
||||
|
||||
@@ -348,6 +366,7 @@ export class Qdrant implements VectorStore {
|
||||
filters?: SearchFilters,
|
||||
topK: number = 100,
|
||||
): Promise<[VectorStoreResult[], number]> {
|
||||
await this.initialize();
|
||||
const scrollRequest = {
|
||||
limit: topK,
|
||||
filter: this.createFilter(filters),
|
||||
@@ -380,6 +399,7 @@ export class Qdrant implements VectorStore {
|
||||
}
|
||||
|
||||
async getUserId(): Promise<string> {
|
||||
await this.initialize();
|
||||
try {
|
||||
// Ensure collection exists (idempotent — handles race conditions)
|
||||
await this.ensureCollection("memory_migrations", 1);
|
||||
@@ -417,6 +437,7 @@ export class Qdrant implements VectorStore {
|
||||
}
|
||||
|
||||
async setUserId(userId: string): Promise<void> {
|
||||
await this.initialize();
|
||||
try {
|
||||
// Get existing point ID
|
||||
const result = await this.client.scroll("memory_migrations", {
|
||||
@@ -496,6 +517,7 @@ export class Qdrant implements VectorStore {
|
||||
|
||||
private async _doInitialize(): Promise<void> {
|
||||
try {
|
||||
await this.ensureClient();
|
||||
await this.ensureCollection(this.collectionName, this.dimension);
|
||||
await this.ensureCollection("memory_migrations", 1);
|
||||
} catch (error) {
|
||||
|
||||
@@ -1,4 +1,3 @@
|
||||
import { createClient } from "redis";
|
||||
import type {
|
||||
RedisClientType,
|
||||
RedisDefaultModules,
|
||||
@@ -8,6 +7,7 @@ import type {
|
||||
} from "redis";
|
||||
import { VectorStore } from "./base";
|
||||
import { SearchFilters, VectorStoreConfig, VectorStoreResult } from "../types";
|
||||
import { loadPeer } from "../utils/load_peer";
|
||||
|
||||
/**
|
||||
* Escape RediSearch TAG filter special characters. Any punctuation in the
|
||||
@@ -146,15 +146,21 @@ function toCamelCase(obj: Record<string, any>): Record<string, any> {
|
||||
}
|
||||
|
||||
export class RedisDB implements VectorStore {
|
||||
private client: RedisClientType<
|
||||
private client!: RedisClientType<
|
||||
RedisDefaultModules & RedisModules & RedisFunctions & RedisScripts
|
||||
>;
|
||||
private readonly redisUrl: string;
|
||||
private readonly username?: string;
|
||||
private readonly password?: string;
|
||||
private readonly indexName: string;
|
||||
private readonly indexPrefix: string;
|
||||
private readonly schema: RedisSchema;
|
||||
private _initPromise?: Promise<void>;
|
||||
|
||||
constructor(config: RedisConfig) {
|
||||
this.redisUrl = config.redisUrl;
|
||||
this.username = config.username;
|
||||
this.password = config.password;
|
||||
this.indexName = config.collectionName;
|
||||
this.indexPrefix = `mem0:${config.collectionName}`;
|
||||
|
||||
@@ -177,12 +183,24 @@ export class RedisDB implements VectorStore {
|
||||
}),
|
||||
};
|
||||
|
||||
this.client = createClient({
|
||||
url: config.redisUrl,
|
||||
username: config.username,
|
||||
password: config.password,
|
||||
this.initialize().catch((err) => {
|
||||
console.error("Failed to initialize Redis:", err);
|
||||
});
|
||||
}
|
||||
|
||||
private async ensureClient(): Promise<void> {
|
||||
if (this.client) return;
|
||||
const sdk = await loadPeer(
|
||||
"redis",
|
||||
"Redis vector store",
|
||||
() => import("redis"),
|
||||
);
|
||||
this.client = sdk.createClient({
|
||||
url: this.redisUrl,
|
||||
username: this.username,
|
||||
password: this.password,
|
||||
socket: {
|
||||
reconnectStrategy: (retries) => {
|
||||
reconnectStrategy: (retries: number) => {
|
||||
if (retries > 10) {
|
||||
console.error("Max reconnection attempts reached");
|
||||
return new Error("Max reconnection attempts reached");
|
||||
@@ -194,10 +212,6 @@ export class RedisDB implements VectorStore {
|
||||
|
||||
this.client.on("error", (err) => console.error("Redis Client Error:", err));
|
||||
this.client.on("connect", () => console.log("Redis Client Connected"));
|
||||
|
||||
this.initialize().catch((err) => {
|
||||
console.error("Failed to initialize Redis:", err);
|
||||
});
|
||||
}
|
||||
|
||||
private async createIndex(): Promise<void> {
|
||||
@@ -260,6 +274,7 @@ export class RedisDB implements VectorStore {
|
||||
}
|
||||
|
||||
private async _doInitialize(): Promise<void> {
|
||||
await this.ensureClient();
|
||||
try {
|
||||
await this.client.connect();
|
||||
console.log("Connected to Redis");
|
||||
@@ -331,6 +346,7 @@ export class RedisDB implements VectorStore {
|
||||
ids: string[],
|
||||
payloads: Record<string, any>[],
|
||||
): Promise<void> {
|
||||
await this.initialize();
|
||||
const data = vectors.map((vector, idx) => {
|
||||
const payload = toSnakeCase(payloads[idx]);
|
||||
const id = ids[idx];
|
||||
@@ -389,6 +405,7 @@ export class RedisDB implements VectorStore {
|
||||
topK: number = 5,
|
||||
filters?: SearchFilters,
|
||||
): Promise<VectorStoreResult[]> {
|
||||
await this.initialize();
|
||||
const snakeFilters = filters ? toSnakeCase(filters) : undefined;
|
||||
const filterExpr = snakeFilters
|
||||
? Object.entries(snakeFilters)
|
||||
@@ -456,6 +473,7 @@ export class RedisDB implements VectorStore {
|
||||
}
|
||||
|
||||
async get(vectorId: string): Promise<VectorStoreResult | null> {
|
||||
await this.initialize();
|
||||
try {
|
||||
// Check if the memory exists first
|
||||
const exists = await this.client.exists(
|
||||
@@ -562,6 +580,7 @@ export class RedisDB implements VectorStore {
|
||||
vector: number[],
|
||||
payload: Record<string, any>,
|
||||
): Promise<void> {
|
||||
await this.initialize();
|
||||
const snakePayload = toSnakeCase(payload);
|
||||
const createdAt = snakePayload.created_at
|
||||
? new Date(snakePayload.created_at).getTime()
|
||||
@@ -601,6 +620,7 @@ export class RedisDB implements VectorStore {
|
||||
}
|
||||
|
||||
async delete(vectorId: string): Promise<void> {
|
||||
await this.initialize();
|
||||
try {
|
||||
// Check if memory exists first
|
||||
const key = `${this.indexPrefix}:${vectorId}`;
|
||||
@@ -626,6 +646,7 @@ export class RedisDB implements VectorStore {
|
||||
}
|
||||
|
||||
async deleteCol(): Promise<void> {
|
||||
await this.initialize();
|
||||
await this.client.ft.dropIndex(this.indexName);
|
||||
}
|
||||
|
||||
@@ -633,6 +654,7 @@ export class RedisDB implements VectorStore {
|
||||
filters?: SearchFilters,
|
||||
topK: number = 100,
|
||||
): Promise<[VectorStoreResult[], number]> {
|
||||
await this.initialize();
|
||||
const snakeFilters = filters ? toSnakeCase(filters) : undefined;
|
||||
const filterExpr = snakeFilters
|
||||
? Object.entries(snakeFilters)
|
||||
@@ -676,10 +698,11 @@ export class RedisDB implements VectorStore {
|
||||
}
|
||||
|
||||
async close(): Promise<void> {
|
||||
await this.client.quit();
|
||||
if (this.client) await this.client.quit();
|
||||
}
|
||||
|
||||
async getUserId(): Promise<string> {
|
||||
await this.initialize();
|
||||
try {
|
||||
// Check if the user ID exists in Redis
|
||||
const userId = await this.client.get("memory_migrations:1");
|
||||
@@ -702,6 +725,7 @@ export class RedisDB implements VectorStore {
|
||||
}
|
||||
|
||||
async setUserId(userId: string): Promise<void> {
|
||||
await this.initialize();
|
||||
try {
|
||||
await this.client.set("memory_migrations:1", userId);
|
||||
} catch (error) {
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
import { createClient, SupabaseClient } from "@supabase/supabase-js";
|
||||
import type { SupabaseClient } from "@supabase/supabase-js";
|
||||
import { VectorStore } from "./base";
|
||||
import { SearchFilters, VectorStoreConfig, VectorStoreResult } from "../types";
|
||||
import { loadPeer } from "../utils/load_peer";
|
||||
|
||||
interface VectorData {
|
||||
id: string;
|
||||
@@ -82,14 +83,17 @@ $$;
|
||||
*/
|
||||
|
||||
export class SupabaseDB implements VectorStore {
|
||||
private client: SupabaseClient;
|
||||
private client!: SupabaseClient;
|
||||
private readonly supabaseUrl: string;
|
||||
private readonly supabaseKey: string;
|
||||
private readonly tableName: string;
|
||||
private readonly embeddingColumnName: string;
|
||||
private readonly metadataColumnName: string;
|
||||
private _initPromise?: Promise<void>;
|
||||
|
||||
constructor(config: SupabaseConfig) {
|
||||
this.client = createClient(config.supabaseUrl, config.supabaseKey);
|
||||
this.supabaseUrl = config.supabaseUrl;
|
||||
this.supabaseKey = config.supabaseKey;
|
||||
this.tableName = config.tableName;
|
||||
this.embeddingColumnName = config.embeddingColumnName || "embedding";
|
||||
this.metadataColumnName = config.metadataColumnName || "metadata";
|
||||
@@ -99,6 +103,16 @@ export class SupabaseDB implements VectorStore {
|
||||
});
|
||||
}
|
||||
|
||||
private async ensureClient(): Promise<void> {
|
||||
if (this.client) return;
|
||||
const sdk = await loadPeer(
|
||||
"@supabase/supabase-js",
|
||||
"Supabase vector store",
|
||||
() => import("@supabase/supabase-js"),
|
||||
);
|
||||
this.client = sdk.createClient(this.supabaseUrl, this.supabaseKey);
|
||||
}
|
||||
|
||||
async initialize(): Promise<void> {
|
||||
if (!this._initPromise) {
|
||||
this._initPromise = this._doInitialize();
|
||||
@@ -107,6 +121,7 @@ export class SupabaseDB implements VectorStore {
|
||||
}
|
||||
|
||||
private async _doInitialize(): Promise<void> {
|
||||
await this.ensureClient();
|
||||
try {
|
||||
// Verify table exists and vector operations work by attempting a test insert
|
||||
const testVector = Array(1536).fill(0);
|
||||
@@ -209,6 +224,7 @@ See the SQL migration instructions in the code comments.`,
|
||||
ids: string[],
|
||||
payloads: Record<string, any>[],
|
||||
): Promise<void> {
|
||||
await this.initialize();
|
||||
try {
|
||||
const data = vectors.map((vector, idx) => ({
|
||||
id: ids[idx],
|
||||
@@ -237,6 +253,7 @@ See the SQL migration instructions in the code comments.`,
|
||||
topK: number = 5,
|
||||
filters?: SearchFilters,
|
||||
): Promise<VectorStoreResult[]> {
|
||||
await this.initialize();
|
||||
try {
|
||||
const rpcQuery: VectorQueryParams = {
|
||||
query_embedding: query,
|
||||
@@ -265,6 +282,7 @@ See the SQL migration instructions in the code comments.`,
|
||||
}
|
||||
|
||||
async get(vectorId: string): Promise<VectorStoreResult | null> {
|
||||
await this.initialize();
|
||||
try {
|
||||
const { data, error } = await this.client
|
||||
.from(this.tableName)
|
||||
@@ -290,6 +308,7 @@ See the SQL migration instructions in the code comments.`,
|
||||
vector: number[],
|
||||
payload: Record<string, any>,
|
||||
): Promise<void> {
|
||||
await this.initialize();
|
||||
try {
|
||||
const { error } = await this.client
|
||||
.from(this.tableName)
|
||||
@@ -310,6 +329,7 @@ See the SQL migration instructions in the code comments.`,
|
||||
}
|
||||
|
||||
async delete(vectorId: string): Promise<void> {
|
||||
await this.initialize();
|
||||
try {
|
||||
const { error } = await this.client
|
||||
.from(this.tableName)
|
||||
@@ -324,6 +344,7 @@ See the SQL migration instructions in the code comments.`,
|
||||
}
|
||||
|
||||
async deleteCol(): Promise<void> {
|
||||
await this.initialize();
|
||||
try {
|
||||
const { error } = await this.client
|
||||
.from(this.tableName)
|
||||
@@ -341,6 +362,7 @@ See the SQL migration instructions in the code comments.`,
|
||||
filters?: SearchFilters,
|
||||
topK: number = 100,
|
||||
): Promise<[VectorStoreResult[], number]> {
|
||||
await this.initialize();
|
||||
try {
|
||||
let query = this.client
|
||||
.from(this.tableName)
|
||||
@@ -370,6 +392,7 @@ See the SQL migration instructions in the code comments.`,
|
||||
}
|
||||
|
||||
async getUserId(): Promise<string> {
|
||||
await this.initialize();
|
||||
try {
|
||||
// First check if the table exists
|
||||
const { data: tableExists } = await this.client
|
||||
@@ -421,6 +444,7 @@ See the SQL migration instructions in the code comments.`,
|
||||
}
|
||||
|
||||
async setUserId(userId: string): Promise<void> {
|
||||
await this.initialize();
|
||||
try {
|
||||
const { error: deleteError } = await this.client
|
||||
.from("memory_migrations")
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
import { VectorStore } from "./base";
|
||||
import { SearchFilters, VectorStoreConfig, VectorStoreResult } from "../types";
|
||||
import { loadPeer } from "../utils/load_peer";
|
||||
|
||||
interface TurbopufferConfig extends VectorStoreConfig {
|
||||
apiKey?: string;
|
||||
@@ -48,14 +49,11 @@ export class TurbopufferDB implements VectorStore {
|
||||
}
|
||||
|
||||
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",
|
||||
const sdk = await loadPeer(
|
||||
"@turbopuffer/turbopuffer",
|
||||
"Turbopuffer vector store",
|
||||
() => import("@turbopuffer/turbopuffer"),
|
||||
);
|
||||
}
|
||||
|
||||
// @turbopuffer/turbopuffer ships `Turbopuffer` as both the default export
|
||||
// and a named export pointing at the same class. Use `.default` since
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
import type { Index, QueryResult, Vector } from "@upstash/vector";
|
||||
import { VectorStore } from "./base";
|
||||
import { SearchFilters, VectorStoreConfig, VectorStoreResult } from "../types";
|
||||
import { loadPeer } from "../utils/load_peer";
|
||||
|
||||
interface UpstashVectorConfig extends VectorStoreConfig {
|
||||
collectionName: string;
|
||||
@@ -42,14 +43,11 @@ export class UpstashVector implements VectorStore {
|
||||
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",
|
||||
const sdk = await loadPeer(
|
||||
"@upstash/vector",
|
||||
"Upstash Vector store",
|
||||
() => import("@upstash/vector"),
|
||||
);
|
||||
}
|
||||
this.client = new sdk.Index({
|
||||
url: config.url,
|
||||
token: config.token,
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
import { VectorStore } from "./base";
|
||||
import { SearchFilters, VectorStoreResult } from "../types";
|
||||
import { ValkeyConfig } from "../types/valkey";
|
||||
import { loadPeer } from "../utils/load_peer";
|
||||
|
||||
interface ValkeyClient {
|
||||
call: (...args: (string | number | Buffer)[]) => Promise<unknown>;
|
||||
@@ -142,14 +143,8 @@ function formatTimestamp(timestamp: number, timezone: string = "UTC"): string {
|
||||
return `${yyyy}-${MM}-${dd}T${HH}:${mm}:${ss}${sign}${offHH}:${offMM}`;
|
||||
}
|
||||
|
||||
async function loadIovalkey(): Promise<typeof import("iovalkey")> {
|
||||
try {
|
||||
return await import("iovalkey");
|
||||
} catch {
|
||||
throw new Error(
|
||||
"iovalkey is required for the Valkey vector store. Install it with: npm install iovalkey",
|
||||
);
|
||||
}
|
||||
function loadIovalkey(): Promise<typeof import("iovalkey")> {
|
||||
return loadPeer("iovalkey", "Valkey vector store", () => import("iovalkey"));
|
||||
}
|
||||
|
||||
export class ValkeyDB implements VectorStore {
|
||||
|
||||
@@ -1,7 +1,8 @@
|
||||
import Cloudflare from "cloudflare";
|
||||
import type Cloudflare from "cloudflare";
|
||||
import type { Vectorize, VectorizeVector } from "@cloudflare/workers-types";
|
||||
import { VectorStore } from "./base";
|
||||
import { SearchFilters, VectorStoreConfig, VectorStoreResult } from "../types";
|
||||
import { loadPeer } from "../utils/load_peer";
|
||||
|
||||
interface VectorizeConfig extends VectorStoreConfig {
|
||||
apiKey?: string;
|
||||
@@ -17,24 +18,36 @@ interface CloudflareVector {
|
||||
|
||||
export class VectorizeDB implements VectorStore {
|
||||
private client: Cloudflare | null = null;
|
||||
private apiKey?: string;
|
||||
private dimensions: number;
|
||||
private indexName: string;
|
||||
private accountId: string;
|
||||
private _initPromise?: Promise<void>;
|
||||
|
||||
constructor(config: VectorizeConfig) {
|
||||
this.client = new Cloudflare({ apiToken: config.apiKey });
|
||||
this.apiKey = config.apiKey;
|
||||
this.dimensions = config.dimension || 1536;
|
||||
this.indexName = config.indexName;
|
||||
this.accountId = config.accountId;
|
||||
this.initialize().catch(console.error);
|
||||
}
|
||||
|
||||
private async ensureClient(): Promise<void> {
|
||||
if (this.client) return;
|
||||
const sdk = await loadPeer(
|
||||
"cloudflare",
|
||||
"Vectorize vector store",
|
||||
() => import("cloudflare"),
|
||||
);
|
||||
this.client = new sdk.default({ apiToken: this.apiKey });
|
||||
}
|
||||
|
||||
async insert(
|
||||
vectors: number[][],
|
||||
ids: string[],
|
||||
payloads: Record<string, any>[],
|
||||
): Promise<void> {
|
||||
await this.initialize();
|
||||
try {
|
||||
const vectorObjects: CloudflareVector[] = vectors.map(
|
||||
(vector, index) => ({
|
||||
@@ -83,6 +96,7 @@ export class VectorizeDB implements VectorStore {
|
||||
topK: number = 5,
|
||||
filters?: SearchFilters,
|
||||
): Promise<VectorStoreResult[]> {
|
||||
await this.initialize();
|
||||
try {
|
||||
const result = await this.client?.vectorize.indexes.query(
|
||||
this.indexName,
|
||||
@@ -111,6 +125,7 @@ export class VectorizeDB implements VectorStore {
|
||||
}
|
||||
|
||||
async get(vectorId: string): Promise<VectorStoreResult | null> {
|
||||
await this.initialize();
|
||||
try {
|
||||
const result = (await this.client?.vectorize.indexes.getByIds(
|
||||
this.indexName,
|
||||
@@ -139,6 +154,7 @@ export class VectorizeDB implements VectorStore {
|
||||
vector: number[],
|
||||
payload: Record<string, any>,
|
||||
): Promise<void> {
|
||||
await this.initialize();
|
||||
try {
|
||||
const data: VectorizeVector = {
|
||||
id: vectorId,
|
||||
@@ -173,6 +189,7 @@ export class VectorizeDB implements VectorStore {
|
||||
}
|
||||
|
||||
async delete(vectorId: string): Promise<void> {
|
||||
await this.initialize();
|
||||
try {
|
||||
await this.client?.vectorize.indexes.deleteByIds(this.indexName, {
|
||||
account_id: this.accountId,
|
||||
@@ -187,6 +204,7 @@ export class VectorizeDB implements VectorStore {
|
||||
}
|
||||
|
||||
async deleteCol(): Promise<void> {
|
||||
await this.initialize();
|
||||
try {
|
||||
await this.client?.vectorize.indexes.delete(this.indexName, {
|
||||
account_id: this.accountId,
|
||||
@@ -203,6 +221,7 @@ export class VectorizeDB implements VectorStore {
|
||||
filters?: SearchFilters,
|
||||
topK: number = 20,
|
||||
): Promise<[VectorStoreResult[], number]> {
|
||||
await this.initialize();
|
||||
try {
|
||||
const result = await this.client?.vectorize.indexes.query(
|
||||
this.indexName,
|
||||
@@ -243,6 +262,7 @@ export class VectorizeDB implements VectorStore {
|
||||
}
|
||||
|
||||
async getUserId(): Promise<string> {
|
||||
await this.initialize();
|
||||
try {
|
||||
let found = false;
|
||||
for await (const index of this.client!.vectorize.indexes.list({
|
||||
@@ -309,6 +329,7 @@ export class VectorizeDB implements VectorStore {
|
||||
}
|
||||
|
||||
async setUserId(userId: string): Promise<void> {
|
||||
await this.initialize();
|
||||
try {
|
||||
// Get existing point ID
|
||||
const result: any = await this.client?.vectorize.indexes.query(
|
||||
@@ -355,6 +376,7 @@ export class VectorizeDB implements VectorStore {
|
||||
}
|
||||
|
||||
private async _doInitialize(): Promise<void> {
|
||||
await this.ensureClient();
|
||||
try {
|
||||
// Check if the index already exists
|
||||
let indexFound = false;
|
||||
|
||||
@@ -2,6 +2,7 @@ import type { WeaviateClient } from "weaviate-client";
|
||||
import { v4 as uuidv4 } from "uuid";
|
||||
import { VectorStore } from "./base";
|
||||
import { SearchFilters, VectorStoreConfig, VectorStoreResult } from "../types";
|
||||
import { loadPeer } from "../utils/load_peer";
|
||||
|
||||
interface WeaviateConfig extends VectorStoreConfig {
|
||||
/** Pre-configured Weaviate client instance (typed as `any` to keep the
|
||||
@@ -50,14 +51,11 @@ export class WeaviateDB implements VectorStore {
|
||||
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",
|
||||
const sdk = await loadPeer(
|
||||
"weaviate-client",
|
||||
"Weaviate vector store",
|
||||
() => import("weaviate-client"),
|
||||
);
|
||||
}
|
||||
this._sdk = sdk;
|
||||
|
||||
const { client, clusterUrl, apiKey, additionalHeaders } = this._config;
|
||||
|
||||
@@ -23,11 +23,16 @@ describe("AnthropicLLM (unit)", () => {
|
||||
|
||||
// Regression #5665: a configured baseURL must reach the Anthropic client so
|
||||
// proxy/gateway users are not silently bypassed (TS parity with #5626).
|
||||
it("forwards baseURL to the Anthropic client when set", () => {
|
||||
new AnthropicLLM({
|
||||
// The client is constructed lazily on first use, so drive generateResponse.
|
||||
it("forwards baseURL to the Anthropic client when set", async () => {
|
||||
mockCreate.mockResolvedValueOnce({
|
||||
content: [{ type: "text", text: "ok" }],
|
||||
});
|
||||
const llm = new AnthropicLLM({
|
||||
apiKey: "test-key",
|
||||
baseURL: "https://proxy.example/v1",
|
||||
});
|
||||
await llm.generateResponse([{ role: "user", content: "Hi" }]);
|
||||
|
||||
expect(mockConstructor).toHaveBeenCalledTimes(1);
|
||||
const ctorArgs = mockConstructor.mock.calls[0][0];
|
||||
@@ -37,8 +42,12 @@ describe("AnthropicLLM (unit)", () => {
|
||||
|
||||
// When no baseURL is configured the client must not receive a baseURL key
|
||||
// (so the SDK default endpoint is used).
|
||||
it("does NOT set baseURL when none is configured", () => {
|
||||
new AnthropicLLM({ apiKey: "test-key" });
|
||||
it("does NOT set baseURL when none is configured", async () => {
|
||||
mockCreate.mockResolvedValueOnce({
|
||||
content: [{ type: "text", text: "ok" }],
|
||||
});
|
||||
const llm = new AnthropicLLM({ apiKey: "test-key" });
|
||||
await llm.generateResponse([{ role: "user", content: "Hi" }]);
|
||||
|
||||
expect(mockConstructor).toHaveBeenCalledTimes(1);
|
||||
const ctorArgs = mockConstructor.mock.calls[0][0];
|
||||
|
||||
@@ -22,7 +22,16 @@ function sourceFiles(dir: string, acc: string[] = []): string[] {
|
||||
for (const entry of readdirSync(dir)) {
|
||||
const full = join(dir, entry);
|
||||
if (statSync(full).isDirectory()) {
|
||||
if (entry !== "tests" && entry !== "__tests__") sourceFiles(full, acc);
|
||||
// examples/ are dev scripts and community/ is the separate @mem0/community
|
||||
// package — neither ships in files:["dist"] nor is reachable via the mem0ai/oss
|
||||
// barrel, so a static import there cannot crash a Memory() consumer.
|
||||
if (
|
||||
entry !== "tests" &&
|
||||
entry !== "__tests__" &&
|
||||
entry !== "examples" &&
|
||||
entry !== "community"
|
||||
)
|
||||
sourceFiles(full, acc);
|
||||
} else if (entry.endsWith(".ts") && !entry.endsWith(".test.ts")) {
|
||||
acc.push(full);
|
||||
}
|
||||
|
||||
@@ -40,14 +40,15 @@ beforeEach(() => {
|
||||
});
|
||||
|
||||
describe("Qdrant URL port extraction (qdrant/qdrant-js#59 workaround)", () => {
|
||||
it("extracts port from HTTPS URL with explicit port", () => {
|
||||
new Qdrant({
|
||||
it("extracts port from HTTPS URL with explicit port", async () => {
|
||||
const store = new Qdrant({
|
||||
url: "https://my-cluster.us-west-1-0.aws.cloud.qdrant.io:6333",
|
||||
apiKey: "test-key",
|
||||
collectionName: "test",
|
||||
embeddingModelDims: 768,
|
||||
dimension: 768,
|
||||
});
|
||||
await store.initialize();
|
||||
|
||||
expect(capturedParams).toBeDefined();
|
||||
expect(capturedParams!.url).toBe(
|
||||
@@ -57,47 +58,50 @@ describe("Qdrant URL port extraction (qdrant/qdrant-js#59 workaround)", () => {
|
||||
expect(capturedParams!.apiKey).toBe("test-key");
|
||||
});
|
||||
|
||||
it("extracts port from HTTP URL with explicit port", () => {
|
||||
new Qdrant({
|
||||
it("extracts port from HTTP URL with explicit port", async () => {
|
||||
const store = new Qdrant({
|
||||
url: "http://localhost:6333",
|
||||
collectionName: "test",
|
||||
embeddingModelDims: 768,
|
||||
dimension: 768,
|
||||
});
|
||||
await store.initialize();
|
||||
|
||||
expect(capturedParams).toBeDefined();
|
||||
expect(capturedParams!.url).toBe("http://localhost:6333");
|
||||
expect(capturedParams!.port).toBe(6333);
|
||||
});
|
||||
|
||||
it("defaults to port 6333 when HTTPS URL has no explicit port", () => {
|
||||
new Qdrant({
|
||||
it("defaults to port 6333 when HTTPS URL has no explicit port", async () => {
|
||||
const store = new Qdrant({
|
||||
url: "https://my-cluster.cloud.qdrant.io",
|
||||
apiKey: "test-key",
|
||||
collectionName: "test",
|
||||
embeddingModelDims: 768,
|
||||
dimension: 768,
|
||||
});
|
||||
await store.initialize();
|
||||
|
||||
expect(capturedParams).toBeDefined();
|
||||
expect(capturedParams!.url).toBe("https://my-cluster.cloud.qdrant.io");
|
||||
expect(capturedParams!.port).toBe(6333);
|
||||
});
|
||||
|
||||
it("defaults to port 6333 when HTTP URL has no explicit port", () => {
|
||||
new Qdrant({
|
||||
it("defaults to port 6333 when HTTP URL has no explicit port", async () => {
|
||||
const store = new Qdrant({
|
||||
url: "http://localhost",
|
||||
collectionName: "test",
|
||||
embeddingModelDims: 768,
|
||||
dimension: 768,
|
||||
});
|
||||
await store.initialize();
|
||||
|
||||
expect(capturedParams).toBeDefined();
|
||||
expect(capturedParams!.port).toBe(6333);
|
||||
});
|
||||
|
||||
it("host+port config overrides URL-extracted port", () => {
|
||||
new Qdrant({
|
||||
it("host+port config overrides URL-extracted port", async () => {
|
||||
const store = new Qdrant({
|
||||
url: "https://my-cluster.cloud.qdrant.io:6333",
|
||||
host: "custom-host",
|
||||
port: 9999,
|
||||
@@ -106,28 +110,28 @@ describe("Qdrant URL port extraction (qdrant/qdrant-js#59 workaround)", () => {
|
||||
embeddingModelDims: 768,
|
||||
dimension: 768,
|
||||
});
|
||||
await store.initialize();
|
||||
|
||||
expect(capturedParams).toBeDefined();
|
||||
expect(capturedParams!.host).toBe("custom-host");
|
||||
expect(capturedParams!.port).toBe(9999);
|
||||
});
|
||||
|
||||
it("handles invalid URL gracefully without crashing", () => {
|
||||
expect(() => {
|
||||
new Qdrant({
|
||||
it("handles invalid URL gracefully without crashing", async () => {
|
||||
const store = new Qdrant({
|
||||
url: "not-a-valid-url",
|
||||
collectionName: "test",
|
||||
embeddingModelDims: 768,
|
||||
dimension: 768,
|
||||
});
|
||||
}).not.toThrow();
|
||||
await expect(store.initialize()).resolves.not.toThrow();
|
||||
|
||||
expect(capturedParams).toBeDefined();
|
||||
expect(capturedParams!.url).toBe("not-a-valid-url");
|
||||
expect(capturedParams!.port).toBe(6333);
|
||||
});
|
||||
|
||||
it("does not pass port when using pre-configured client", () => {
|
||||
it("does not pass port when using pre-configured client", async () => {
|
||||
const mockClient: any = {
|
||||
createCollection: jest.fn().mockResolvedValue(undefined),
|
||||
getCollection: jest.fn().mockResolvedValue({
|
||||
@@ -141,12 +145,13 @@ describe("Qdrant URL port extraction (qdrant/qdrant-js#59 workaround)", () => {
|
||||
deleteCollection: jest.fn().mockResolvedValue(undefined),
|
||||
};
|
||||
|
||||
new Qdrant({
|
||||
const store = new Qdrant({
|
||||
client: mockClient,
|
||||
collectionName: "test",
|
||||
embeddingModelDims: 768,
|
||||
dimension: 768,
|
||||
});
|
||||
await store.initialize();
|
||||
|
||||
// QdrantClient constructor should NOT have been called
|
||||
expect(
|
||||
@@ -154,14 +159,15 @@ describe("Qdrant URL port extraction (qdrant/qdrant-js#59 workaround)", () => {
|
||||
).not.toHaveBeenCalled();
|
||||
});
|
||||
|
||||
it("defaults to 6333 when HTTPS URL uses default port 443", () => {
|
||||
new Qdrant({
|
||||
it("defaults to 6333 when HTTPS URL uses default port 443", async () => {
|
||||
const store = new Qdrant({
|
||||
url: "https://my-cluster.cloud.qdrant.io:443",
|
||||
apiKey: "test-key",
|
||||
collectionName: "test",
|
||||
embeddingModelDims: 768,
|
||||
dimension: 768,
|
||||
});
|
||||
await store.initialize();
|
||||
|
||||
// 443 is default for HTTPS, so URL.port returns empty string — we default to 6333
|
||||
expect(capturedParams).toBeDefined();
|
||||
|
||||
Reference in New Issue
Block a user