Compare commits
28 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 42cf18c4e6 | |||
| d89793b666 | |||
| 17836748d7 | |||
| c9af55986e | |||
| f69f8dcc7b | |||
| 28e4d819f8 | |||
| 1d383c4ee2 | |||
| 8488abe603 | |||
| 49863e9a7a | |||
| 770ce97bd9 | |||
| 44bcfbe1f3 | |||
| 33a0ed7559 | |||
| ba9054e0c8 | |||
| df9d5cc4b1 | |||
| 3b2357bfe0 | |||
| 573b20cec8 | |||
| a781800d3f | |||
| 4470803fe5 | |||
| 6a801bfe2f | |||
| 2a4aa232b2 | |||
| 99206f0c64 | |||
| 5dbf071356 | |||
| b26469e006 | |||
| e72ae96ad4 | |||
| fbdbab805d | |||
| 22f70d50e1 | |||
| f89edb45dc | |||
| 846f25bd39 |
@@ -46,12 +46,12 @@
|
||||
|
||||
| Benchmark | Old | New | Tokens | Latency p50 |
|
||||
| --- | --- | --- | --- | --- |
|
||||
| **LoCoMo** | 71.4 | **91.6** | 7.0K | 0.88s |
|
||||
| **LongMemEval** | 67.8 | **94.8** | 6.8K | 1.09s |
|
||||
| **LoCoMo** | 71.4 | **92.5** | 7.0K | 0.88s |
|
||||
| **LongMemEval** | 67.8 | **94.4** | 6.8K | 1.09s |
|
||||
| **BEAM (1M)** | — | **64.1** | 6.7K | 1.00s |
|
||||
| **BEAM (10M)** | — | **48.6** | 6.9K | 1.05s |
|
||||
|
||||
All benchmarks run on the same production-representative model stack. Single-pass retrieval (one call, no agentic loops).
|
||||
All benchmarks run on the same production-representative model stack. Single-pass retrieval (one call, no agentic loops) at a top_200 retrieval budget. Scores reflect Mem0's managed platform, which includes proprietary optimizations not available in the open-source SDK; open-source users should expect directionally similar gains but not identical numbers.
|
||||
|
||||
**What changed:**
|
||||
- **Single-pass ADD-only extraction** -- one LLM call, no UPDATE/DELETE. Memories accumulate; nothing is overwritten.
|
||||
@@ -63,8 +63,8 @@ All benchmarks run on the same production-representative model stack. Single-pas
|
||||
See the [migration guide](https://docs.mem0.ai/migration/oss-v2-to-v3) for upgrade instructions. The [evaluation framework](https://github.com/mem0ai/memory-benchmarks) is open-sourced so anyone can reproduce the numbers.
|
||||
|
||||
## Research Highlights
|
||||
- **91.6 on LoCoMo** -- +20 points over the previous algorithm
|
||||
- **94.8 on LongMemEval** -- +27 points, with +53.6 on assistant memory recall
|
||||
- **92.5 on LoCoMo** -- +21 points over the previous algorithm
|
||||
- **94.4 on LongMemEval** -- +27 points, with 98.2 on assistant memory recall
|
||||
- **64.1 on BEAM (1M)** -- production-scale memory evaluation at 1M tokens
|
||||
- [Read the full paper](https://mem0.ai/research)
|
||||
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
{
|
||||
"name": "@mem0/cli",
|
||||
"version": "0.2.10",
|
||||
"version": "0.2.11",
|
||||
"description": "The official CLI for mem0 — the memory layer for AI agents",
|
||||
"type": "module",
|
||||
"bin": {
|
||||
|
||||
@@ -4,7 +4,7 @@ build-backend = "hatchling.build"
|
||||
|
||||
[project]
|
||||
name = "mem0-cli"
|
||||
version = "0.2.9"
|
||||
version = "0.2.10"
|
||||
description = "The official CLI for mem0 — the memory layer for AI agents"
|
||||
readme = "README.md"
|
||||
license = "Apache-2.0"
|
||||
|
||||
@@ -1,3 +1,3 @@
|
||||
"""mem0 CLI — the command-line interface for the mem0 memory layer."""
|
||||
|
||||
__version__ = "0.2.9"
|
||||
__version__ = "0.2.10"
|
||||
|
||||
@@ -4,6 +4,22 @@ description: "Major product launches, headline features, and milestones for Mem0
|
||||
mode: "wide"
|
||||
---
|
||||
|
||||
<Update label="2026-07-13" description="TypeScript provider expansion">
|
||||
|
||||
**TypeScript OSS SDK: 26 New Providers, Reranking, and Zero-Dependency Imports**
|
||||
|
||||
TypeScript SDK v3.1.0 is the largest provider release for the OSS SDK so far, closing most of the remaining gap with the Python SDK. Python SDK v2.0.12 ships alongside it with fixes and security patches.
|
||||
|
||||
- **17 new vector stores:** Pinecone, Weaviate, Milvus, Chroma, MongoDB, Elasticsearch, OpenSearch, Databricks, AWS Neptune Analytics, S3 Vectors, Azure MySQL, Google Vertex AI Vector Search, Turbopuffer, Upstash Vector, Valkey, Cassandra, and Baidu Mochow.
|
||||
- **5 new LLM providers:** AWS Bedrock, xAI Grok, Together, vLLM, and Sarvam.
|
||||
- **4 new embedding providers:** Vertex AI, HuggingFace, FastEmbed, and Together.
|
||||
- **Reranking in TypeScript:** Four rerankers (Cohere, ZeroEntropy, cross-encoder, and LLM-based) with per-search rerank via a `rerank` option on `search()`.
|
||||
- **Install only what you use:** Importing `mem0ai/oss` no longer pulls in any provider SDK. Provider packages are resolved lazily on first use, so an app that configures only OpenAI and Qdrant does not need the other provider SDKs installed.
|
||||
|
||||
See [SDK & Tools](/changelog/sdk) for version details and PR links.
|
||||
|
||||
</Update>
|
||||
|
||||
<Update label="2026-06-27" description="SDK memory expiration">
|
||||
|
||||
**SDK Memory Expiration: Expiring Memories Across Python and TypeScript**
|
||||
|
||||
@@ -7,6 +7,34 @@ mode: "wide"
|
||||
<Tabs>
|
||||
<Tab title="Python">
|
||||
|
||||
<Update label="2026-07-13" description="v2.0.12">
|
||||
|
||||
**New Features:**
|
||||
- **Memory (OSS):** Accept `text` in `Memory.update()` and `AsyncMemory.update()`. `data` still works but is now deprecated, so prefer `text` in new code ([#6044](https://github.com/mem0ai/mem0/pull/6044))
|
||||
|
||||
**Bug Fixes:**
|
||||
- **Core:** Coerce non-string entity IDs (`user_id`, `agent_id`, `run_id`) instead of crashing on `.strip()`, so passing an integer ID no longer raises `AttributeError` ([#6206](https://github.com/mem0ai/mem0/pull/6206))
|
||||
- **Core:** Stop requiring `langchain-core` for the default async procedural memory path. The optional dependency is now only imported when you pass a custom LangChain LLM, matching the sync behavior ([#6209](https://github.com/mem0ai/mem0/pull/6209))
|
||||
- **Client:** Encode dynamic URL path segments so IDs containing special characters no longer produce malformed requests ([#5963](https://github.com/mem0ai/mem0/pull/5963))
|
||||
- **LLMs:** Skip `temperature` and `top_p` for newer Anthropic models that reject sampling parameters. Detection is automatic per model family and version, and the new `enable_sampling_parameters` config flag overrides it ([#6211](https://github.com/mem0ai/mem0/pull/6211))
|
||||
- **Vector Stores:** Stop writing internal `OutputData` model fields as properties on Weaviate `update()` ([#6149](https://github.com/mem0ai/mem0/pull/6149))
|
||||
- **Vector Stores:** Improve wildcard search handling in Milvus ([#6187](https://github.com/mem0ai/mem0/pull/6187))
|
||||
- **Vector Stores:** Keep env-resolved Upstash Vector credentials after config validation. An env-var-only config previously passed validation and then failed to build ([#5811](https://github.com/mem0ai/mem0/pull/5811))
|
||||
- **Vector Stores:** Restore the previous payload when a Neptune Analytics vector upsert fails inside `update()`, so a partial write can no longer leave the payload and embedding out of sync ([#5824](https://github.com/mem0ai/mem0/pull/5824))
|
||||
|
||||
**Changes:**
|
||||
- **LLMs:** The Together default model is now `MiniMaxAI/MiniMax-M3` (was `mistralai/Mixtral-8x7B-Instruct-v0.1`) ([#6049](https://github.com/mem0ai/mem0/pull/6049))
|
||||
- **LLMs:** The xAI default model is now `grok-4.3` (was `grok-2-latest`) ([#6115](https://github.com/mem0ai/mem0/pull/6115))
|
||||
- **Embeddings:** The Together default embedding model is now `intfloat/multilingual-e5-large-instruct` at 1024 dimensions (was `togethercomputer/m2-bert-80M-8k-retrieval` at 768). If you use the Together embedder without pinning `model`, existing vectors were written at the old dimension: either re-embed them, or pin `model` and `embedding_dims` to the old values ([#5989](https://github.com/mem0ai/mem0/pull/5989))
|
||||
- **Rerankers:** The Cohere default rerank model is now `rerank-v3.5` (was `rerank-english-v3.0`) ([#6055](https://github.com/mem0ai/mem0/pull/6055))
|
||||
|
||||
**Security:**
|
||||
- **Vector Stores:** Fix SQL and Cypher injection vulnerabilities in the PGVector, Azure MySQL, and Neptune providers ([#4878](https://github.com/mem0ai/mem0/pull/4878))
|
||||
- **Vector Stores:** Validate Elasticsearch filter keys and values to prevent term query injection ([#5980](https://github.com/mem0ai/mem0/pull/5980))
|
||||
- **Dependencies:** Require `transformers>=5.3.0` to remediate GHSA-29pf-2h5f-8g72 (CVE-2026-4372) ([#6110](https://github.com/mem0ai/mem0/pull/6110))
|
||||
|
||||
</Update>
|
||||
|
||||
<Update label="2026-07-01" description="v2.0.11">
|
||||
|
||||
**Bug Fixes:**
|
||||
@@ -1100,6 +1128,32 @@ See the [OSS v2 to v3 migration guide](https://docs.mem0.ai/migration/oss-v2-to-
|
||||
|
||||
<Tab title="TypeScript">
|
||||
|
||||
<Update label="2026-07-13" description="v3.1.0">
|
||||
|
||||
The largest provider release for the TypeScript OSS SDK so far: 17 new vector stores, 5 new LLM providers, 4 new embedders, and reranking support. Importing `mem0ai/oss` no longer pulls in any provider SDK, so you only install what you actually configure.
|
||||
|
||||
**New Features:**
|
||||
- **Rerankers:** Add reranking to the OSS SDK with four providers (Cohere, ZeroEntropy, cross-encoder, and LLM-based), plus per-search rerank via a `rerank` option on `search()` ([#6055](https://github.com/mem0ai/mem0/pull/6055))
|
||||
- **Memory (OSS):** Accept `text` in `Memory.update()`. `data` still works but is now deprecated, so prefer `text` in new code ([#6044](https://github.com/mem0ai/mem0/pull/6044))
|
||||
- **Vector Stores:** Add Pinecone ([#5802](https://github.com/mem0ai/mem0/pull/5802)), Weaviate ([#5800](https://github.com/mem0ai/mem0/pull/5800)), Milvus ([#5889](https://github.com/mem0ai/mem0/pull/5889)), Chroma ([#6145](https://github.com/mem0ai/mem0/pull/6145)), MongoDB ([#5793](https://github.com/mem0ai/mem0/pull/5793)), Elasticsearch ([#5866](https://github.com/mem0ai/mem0/pull/5866)), and OpenSearch ([#5810](https://github.com/mem0ai/mem0/pull/5810))
|
||||
- **Vector Stores:** Add Databricks ([#5824](https://github.com/mem0ai/mem0/pull/5824)), AWS Neptune Analytics ([#5797](https://github.com/mem0ai/mem0/pull/5797)), S3 Vectors ([#5822](https://github.com/mem0ai/mem0/pull/5822)), Azure MySQL ([#5827](https://github.com/mem0ai/mem0/pull/5827)), and Google Vertex AI Vector Search ([#5791](https://github.com/mem0ai/mem0/pull/5791))
|
||||
- **Vector Stores:** Add Turbopuffer ([#5801](https://github.com/mem0ai/mem0/pull/5801)), Upstash Vector ([#5811](https://github.com/mem0ai/mem0/pull/5811)), Valkey ([#5826](https://github.com/mem0ai/mem0/pull/5826)), Cassandra ([#5823](https://github.com/mem0ai/mem0/pull/5823)), and Baidu Mochow ([#5790](https://github.com/mem0ai/mem0/pull/5790))
|
||||
- **LLMs:** Add AWS Bedrock ([#5890](https://github.com/mem0ai/mem0/pull/5890)), xAI Grok ([#6115](https://github.com/mem0ai/mem0/pull/6115)), Together ([#6049](https://github.com/mem0ai/mem0/pull/6049)), vLLM ([#5805](https://github.com/mem0ai/mem0/pull/5805)), and Sarvam ([#6130](https://github.com/mem0ai/mem0/pull/6130))
|
||||
- **Embeddings:** Add Vertex AI ([#5882](https://github.com/mem0ai/mem0/pull/5882)), HuggingFace ([#6027](https://github.com/mem0ai/mem0/pull/6027)), FastEmbed ([#5862](https://github.com/mem0ai/mem0/pull/5862)), and Together ([#5989](https://github.com/mem0ai/mem0/pull/5989))
|
||||
|
||||
**Improvements:**
|
||||
- **Packaging:** Lazy-load optional provider SDKs so importing `mem0ai/oss` never requires them. Provider packages are now resolved on first use, so an app that only configures OpenAI and Qdrant does not need the other provider SDKs installed ([#6280](https://github.com/mem0ai/mem0/pull/6280))
|
||||
|
||||
**Bug Fixes:**
|
||||
- **Memory (OSS):** Re-raise LLM extraction transport failures instead of returning `[]`, so a network error during extraction surfaces as an error rather than a silently empty result ([#6102](https://github.com/mem0ai/mem0/pull/6102))
|
||||
- **Vector Stores:** Prevent an unhandled promise rejection in the Supabase and Redis constructors ([#6111](https://github.com/mem0ai/mem0/pull/6111))
|
||||
- **Client:** Encode dynamic URL path segments so IDs containing special characters no longer produce malformed requests ([#5963](https://github.com/mem0ai/mem0/pull/5963))
|
||||
|
||||
**Security:**
|
||||
- **Dependencies:** Patch the `fast-xml-parser` and `tar` transitive CVEs ([#6160](https://github.com/mem0ai/mem0/pull/6160))
|
||||
|
||||
</Update>
|
||||
|
||||
<Update label="2026-07-01" description="v3.0.13">
|
||||
|
||||
**Bug Fixes:**
|
||||
@@ -1606,6 +1660,13 @@ See the [TypeScript SDK migration guide](https://docs.mem0.ai/migration/ts-v2-to
|
||||
|
||||
<Tab title="CLI">
|
||||
|
||||
<Update label="2026-07-13" description="Python v0.2.10 / Node v0.2.11">
|
||||
|
||||
**Bug Fixes:**
|
||||
- **Platform backend:** Encode dynamic URL path segments so memory and entity IDs containing special characters no longer produce malformed requests (Python and Node [#5963](https://github.com/mem0ai/mem0/pull/5963))
|
||||
|
||||
</Update>
|
||||
|
||||
<Update label="2026-07-01" description="Python v0.2.9 / Node v0.2.10">
|
||||
|
||||
**Bug Fixes:**
|
||||
|
||||
@@ -5,6 +5,10 @@ description: "Configure Hugging Face as an embedding provider in Mem0 for local
|
||||
|
||||
You can use embedding models from Huggingface to run Mem0 locally.
|
||||
|
||||
<Note>
|
||||
The TypeScript SDK supports Hugging Face only through a hosted [Text Embeddings Inference (TEI)](#using-text-embeddings-inference-tei) endpoint, or any OpenAI-compatible Hugging Face endpoint. The local `sentence-transformers` mode shown first is Python-only.
|
||||
</Note>
|
||||
|
||||
### Usage
|
||||
|
||||
```python
|
||||
@@ -34,9 +38,10 @@ m.add(messages, user_id="john")
|
||||
|
||||
### Using Text Embeddings Inference (TEI)
|
||||
|
||||
You can also use Hugging Face's Text Embeddings Inference service for faster and more efficient embeddings:
|
||||
You can also use Hugging Face's Text Embeddings Inference service for faster and more efficient embeddings. This is the mode the TypeScript SDK uses.
|
||||
|
||||
```python
|
||||
<CodeGroup>
|
||||
```python Python
|
||||
import os
|
||||
from mem0 import Memory
|
||||
|
||||
@@ -56,6 +61,24 @@ m = Memory.from_config(config)
|
||||
m.add("This text will be embedded using the TEI service.", user_id="john")
|
||||
```
|
||||
|
||||
```typescript TypeScript
|
||||
import { Memory } from 'mem0ai/oss';
|
||||
|
||||
// Point at a running TEI server, or any OpenAI-compatible HF endpoint
|
||||
const config = {
|
||||
embedder: {
|
||||
provider: 'huggingface',
|
||||
config: {
|
||||
huggingfaceBaseUrl: 'http://localhost:3000/v1',
|
||||
},
|
||||
},
|
||||
};
|
||||
|
||||
const memory = new Memory(config);
|
||||
await memory.add("This text will be embedded using the TEI service.", { userId: "john" });
|
||||
```
|
||||
</CodeGroup>
|
||||
|
||||
To run the TEI service, you can use Docker:
|
||||
|
||||
```bash
|
||||
@@ -66,11 +89,22 @@ docker run -d -p 3000:80 -v huggingfacetei:/data --platform linux/amd64 \
|
||||
|
||||
### Config
|
||||
|
||||
Here are the parameters available for configuring Huggingface embedder:
|
||||
Here are the parameters available for configuring the Hugging Face embedder:
|
||||
|
||||
<Tabs>
|
||||
<Tab title="Python">
|
||||
| Parameter | Description | Default Value |
|
||||
| --- | --- | --- |
|
||||
| `model` | The name of the model to use | `multi-qa-MiniLM-L6-cos-v1` |
|
||||
| `embedding_dims` | Dimensions of the embedding model | `selected_model_dimensions` |
|
||||
| `model_kwargs` | Additional arguments for the model | `None` |
|
||||
| `huggingface_base_url` | URL to connect to Text Embeddings Inference (TEI) API | `None` |
|
||||
| `huggingface_base_url` | URL to connect to Text Embeddings Inference (TEI) API | `None` |
|
||||
</Tab>
|
||||
<Tab title="TypeScript">
|
||||
| Parameter | Description | Default Value |
|
||||
| --- | --- | --- |
|
||||
| `huggingfaceBaseUrl` | TEI or OpenAI-compatible endpoint URL. Required; falls back to `baseURL`, `url`, then the `HUGGINGFACE_BASE_URL` env var | `None` |
|
||||
| `model` | Model name sent to the endpoint (TEI ignores it) | `tei` |
|
||||
| `apiKey` | API key for the endpoint; falls back to the `HUGGINGFACE_API_KEY` env var | `"hf"` |
|
||||
</Tab>
|
||||
</Tabs>
|
||||
@@ -4,11 +4,36 @@ description: "Configure Google Cloud Vertex AI as an embedding provider in Mem0
|
||||
---
|
||||
### Vertex AI
|
||||
|
||||
To use Google Cloud's Vertex AI for text embedding models, set the `GOOGLE_APPLICATION_CREDENTIALS` environment variable to point to the path of your service account's credentials JSON file. These credentials can be created in the [Google Cloud Console](https://console.cloud.google.com/).
|
||||
Google Cloud's Vertex AI serves text embedding models such as `gemini-embedding-001`. Mem0 uses them through the provider's own SDK, which you install alongside Mem0.
|
||||
|
||||
### Installation
|
||||
|
||||
The Vertex AI client is an optional dependency, so install it yourself.
|
||||
|
||||
<CodeGroup>
|
||||
```bash Python
|
||||
pip install vertexai
|
||||
```
|
||||
|
||||
```bash TypeScript
|
||||
npm install @google-cloud/aiplatform
|
||||
```
|
||||
</CodeGroup>
|
||||
|
||||
### Authentication
|
||||
|
||||
Both SDKs authenticate with [Application Default Credentials](https://cloud.google.com/docs/authentication/application-default-credentials). Pick whichever fits your environment:
|
||||
|
||||
- **Local development:** run `gcloud auth application-default login`.
|
||||
- **Service account:** create a key in the [Google Cloud Console](https://console.cloud.google.com/) and point `GOOGLE_APPLICATION_CREDENTIALS` at the JSON file, or pass its path through the embedder config.
|
||||
- **Google Cloud runtimes** (Cloud Run, GKE, Compute Engine): the attached service account is picked up automatically.
|
||||
|
||||
The TypeScript SDK reads the project ID from `googleProjectId`, then the `GCP_PROJECT_ID`, `GOOGLE_CLOUD_PROJECT`, and `GCLOUD_PROJECT` environment variables, and finally from your credentials. Set it explicitly when your credentials cover more than one project.
|
||||
|
||||
### Usage
|
||||
|
||||
```python
|
||||
<CodeGroup>
|
||||
```python Python
|
||||
import os
|
||||
from mem0 import Memory
|
||||
|
||||
@@ -32,28 +57,87 @@ m = Memory.from_config(config)
|
||||
messages = [
|
||||
{"role": "user", "content": "I'm planning to watch a movie tonight. Any recommendations?"},
|
||||
{"role": "assistant", "content": "How about thriller movies? They can be quite engaging."},
|
||||
{"role": "user", "content": "I’m not a big fan of thriller movies but I love sci-fi movies."},
|
||||
{"role": "user", "content": "I'm not a big fan of thriller movies but I love sci-fi movies."},
|
||||
{"role": "assistant", "content": "Got it! I'll avoid thriller recommendations and suggest sci-fi movies in the future."}
|
||||
]
|
||||
m.add(messages, user_id="john")
|
||||
```
|
||||
The embedding types can be one of the following:
|
||||
|
||||
```typescript TypeScript
|
||||
import { Memory } from "mem0ai/oss";
|
||||
|
||||
const config = {
|
||||
embedder: {
|
||||
provider: "vertexai",
|
||||
config: {
|
||||
model: "gemini-embedding-001",
|
||||
// Optional. Falls back to GCP_PROJECT_ID / GOOGLE_CLOUD_PROJECT /
|
||||
// GCLOUD_PROJECT, then to the project on your credentials.
|
||||
googleProjectId: process.env.GCP_PROJECT_ID,
|
||||
location: "us-central1",
|
||||
// Optional. Path to a service account key file, or pass the JSON inline
|
||||
// via googleServiceAccountJson.
|
||||
vertexCredentialsJson: "/path/to/your/credentials.json",
|
||||
embeddingDims: 256,
|
||||
memoryAddEmbeddingType: "RETRIEVAL_DOCUMENT",
|
||||
memoryUpdateEmbeddingType: "RETRIEVAL_DOCUMENT",
|
||||
memorySearchEmbeddingType: "RETRIEVAL_QUERY",
|
||||
},
|
||||
},
|
||||
};
|
||||
|
||||
const memory = new Memory(config);
|
||||
await memory.add("I love sci-fi movies but not thrillers", { userId: "john" });
|
||||
```
|
||||
</CodeGroup>
|
||||
|
||||
### Embedding types
|
||||
|
||||
Vertex AI embeds the same text differently depending on the task you declare. The embedding types can be one of the following:
|
||||
- SEMANTIC_SIMILARITY
|
||||
- CLASSIFICATION
|
||||
- CLUSTERING
|
||||
- RETRIEVAL_DOCUMENT, RETRIEVAL_QUERY, QUESTION_ANSWERING, FACT_VERIFICATION
|
||||
- CODE_RETRIEVAL_QUERY
|
||||
Check out the [Vertex AI documentation](https://cloud.google.com/vertex-ai/generative-ai/docs/embeddings/task-types#supported_task_types) for more information.
|
||||
|
||||
- CODE_RETRIEVAL_QUERY
|
||||
|
||||
Check out the [Vertex AI documentation](https://cloud.google.com/vertex-ai/generative-ai/docs/embeddings/task-types#supported_task_types) for more information.
|
||||
|
||||
<Note>
|
||||
These embedding types map to the add, update, and search memory actions in both the Python and TypeScript SDKs. Stored memories use the add or update type, and searches use the search type.
|
||||
</Note>
|
||||
|
||||
### Choosing a model
|
||||
|
||||
<Warning>
|
||||
`gemini-embedding-001` accepts **one input text per request**. When Mem0 embeds several texts at once, such as the memories extracted from a single conversation turn, it issues one request per text. The older `text-embedding-005` and `text-multilingual-embedding-002` models accept up to 250 texts per request, so they are faster and cheaper for large batches. See [Get text embeddings](https://cloud.google.com/vertex-ai/generative-ai/docs/embeddings/get-text-embeddings).
|
||||
</Warning>
|
||||
|
||||
### Config
|
||||
|
||||
Here are the parameters available for configuring the Vertex AI embedder:
|
||||
|
||||
| Parameter | Description | Default Value |
|
||||
| ------------------------- | ------------------------------------------------ | -------------------- |
|
||||
| `model` | The name of the Vertex AI embedding model to use | `gemini-embedding-001` |
|
||||
| `vertex_credentials_json` | Path to the Google Cloud credentials JSON file | `None` |
|
||||
| `embedding_dims` | Dimensions of the embedding model | `256` |
|
||||
| `memory_add_embedding_type` | The type of embedding to use for the add memory action | `RETRIEVAL_DOCUMENT` |
|
||||
| `memory_update_embedding_type` | The type of embedding to use for the update memory action | `RETRIEVAL_DOCUMENT` |
|
||||
| `memory_search_embedding_type` | The type of embedding to use for the search memory action | `RETRIEVAL_QUERY` |
|
||||
<Tabs>
|
||||
<Tab title="Python">
|
||||
| Parameter | Description | Default Value |
|
||||
| -------------------------------- | ---------------------------------------------------------- | ---------------------- |
|
||||
| `model` | The name of the Vertex AI embedding model to use | `gemini-embedding-001` |
|
||||
| `vertex_credentials_json` | Path to the Google Cloud credentials JSON file | `None` |
|
||||
| `embedding_dims` | Dimensions of the embedding model | `256` |
|
||||
| `memory_add_embedding_type` | The embedding type to use for the add memory action | `RETRIEVAL_DOCUMENT` |
|
||||
| `memory_update_embedding_type` | The embedding type to use for the update memory action | `RETRIEVAL_DOCUMENT` |
|
||||
| `memory_search_embedding_type` | The embedding type to use for the search memory action | `RETRIEVAL_QUERY` |
|
||||
</Tab>
|
||||
<Tab title="TypeScript">
|
||||
| Parameter | Description | Default Value |
|
||||
| ----------------------------- | -------------------------------------------------------------------------- | ---------------------- |
|
||||
| `model` | The name of the Vertex AI embedding model to use | `gemini-embedding-001` |
|
||||
| `googleProjectId` | Google Cloud project ID (falls back to `GCP_PROJECT_ID` env var, then to your credentials) | Resolved from credentials |
|
||||
| `location` | Google Cloud region (falls back to `GCP_LOCATION` env var) | `us-central1` |
|
||||
| `vertexCredentialsJson` | Path to the Google Cloud credentials JSON file | `None` |
|
||||
| `googleServiceAccountJson` | Service account credentials as a JSON string or object | `None` |
|
||||
| `embeddingDims` | Dimensions of the embedding model | `256` |
|
||||
| `memoryAddEmbeddingType` | The embedding type to use for the add memory action | `RETRIEVAL_DOCUMENT` |
|
||||
| `memoryUpdateEmbeddingType` | The embedding type to use for the update memory action | `RETRIEVAL_DOCUMENT` |
|
||||
| `memorySearchEmbeddingType` | The embedding type to use for the search memory action | `RETRIEVAL_QUERY` |
|
||||
</Tab>
|
||||
</Tabs>
|
||||
|
||||
@@ -5,16 +5,18 @@ description: "Configure AWS Bedrock as an LLM provider in Mem0 with IAM authenti
|
||||
|
||||
### Setup
|
||||
- Before using the AWS Bedrock LLM, make sure you have the appropriate model access from [Bedrock Console](https://us-east-1.console.aws.amazon.com/bedrock/home?region=us-east-1#/modelaccess).
|
||||
- You will also need to authenticate the `boto3` client by using a method in the [AWS documentation](https://boto3.amazonaws.com/v1/documentation/api/latest/guide/credentials.html#configuring-credentials)
|
||||
- You will have to export `AWS_REGION`, `AWS_ACCESS_KEY_ID`, and `AWS_SECRET_ACCESS_KEY` to set environment variables.
|
||||
- Model availability is per-region. `anthropic.claude-sonnet-4-20250514-v1:0` supports on-demand inference in `us-east-1` and `ap-southeast-4`; from any other region, use the cross-region inference profile ID `us.anthropic.claude-sonnet-4-20250514-v1:0` instead.
|
||||
- Install the AWS SDK for your language: `pip install boto3` (Python) or `npm install @aws-sdk/client-bedrock-runtime` (TypeScript).
|
||||
- Both SDKs fall back to the standard AWS credential chain (environment variables, `~/.aws/credentials`, or an attached IAM role), so exporting `AWS_REGION`, `AWS_ACCESS_KEY_ID`, and `AWS_SECRET_ACCESS_KEY` is the quickest way to get started. In TypeScript you can also pass credentials inline with `awsRegion`, `awsAccessKeyId`, `awsSecretAccessKey`, and `awsSessionToken`, as shown below.
|
||||
|
||||
### Usage
|
||||
|
||||
```python
|
||||
<CodeGroup>
|
||||
```python Python
|
||||
import os
|
||||
from mem0 import Memory
|
||||
|
||||
os.environ['AWS_REGION'] = 'us-west-2'
|
||||
os.environ['AWS_REGION'] = 'us-east-1'
|
||||
os.environ["AWS_ACCESS_KEY_ID"] = "xx"
|
||||
os.environ["AWS_SECRET_ACCESS_KEY"] = "xx"
|
||||
|
||||
@@ -22,7 +24,7 @@ config = {
|
||||
"llm": {
|
||||
"provider": "aws_bedrock",
|
||||
"config": {
|
||||
"model": "anthropic.claude-3-5-haiku-20241022-v1:0",
|
||||
"model": "anthropic.claude-sonnet-4-20250514-v1:0",
|
||||
"temperature": 0.2,
|
||||
"max_tokens": 2000,
|
||||
}
|
||||
@@ -39,6 +41,43 @@ messages = [
|
||||
m.add(messages, user_id="alice", metadata={"category": "movies"})
|
||||
```
|
||||
|
||||
```typescript TypeScript
|
||||
import { Memory } from 'mem0ai/oss';
|
||||
|
||||
const config = {
|
||||
llm: {
|
||||
provider: 'aws_bedrock',
|
||||
config: {
|
||||
model: 'anthropic.claude-sonnet-4-20250514-v1:0',
|
||||
temperature: 0.2,
|
||||
maxTokens: 2000,
|
||||
// Optional. Omit these to use the default AWS credential chain.
|
||||
awsRegion: process.env.AWS_REGION,
|
||||
awsAccessKeyId: process.env.AWS_ACCESS_KEY_ID,
|
||||
awsSecretAccessKey: process.env.AWS_SECRET_ACCESS_KEY,
|
||||
},
|
||||
},
|
||||
};
|
||||
|
||||
const memory = new Memory(config);
|
||||
const messages = [
|
||||
{"role": "user", "content": "I'm planning to watch a movie tonight. Any recommendations?"},
|
||||
{"role": "assistant", "content": "How about thriller movies? They can be quite engaging."},
|
||||
{"role": "user", "content": "I’m not a big fan of thriller movies but I love sci-fi movies."},
|
||||
{"role": "assistant", "content": "Got it! I'll avoid thriller recommendations and suggest sci-fi movies in the future."}
|
||||
];
|
||||
await memory.add(messages, { userId: 'alice', metadata: { category: 'movies' } });
|
||||
```
|
||||
</CodeGroup>
|
||||
|
||||
<Note>
|
||||
`@aws-sdk/client-bedrock-runtime` is an optional peer dependency of `mem0ai`, so npm will not install it for you. The TypeScript provider loads it lazily and throws a clear error on the first request if the package is missing.
|
||||
</Note>
|
||||
|
||||
<Note>
|
||||
The TypeScript provider calls the Bedrock [Converse API](https://docs.aws.amazon.com/bedrock/latest/userguide/conversation-inference.html), a single uniform interface across the current Bedrock model families. Streaming and `InvokeModel`-only models are not supported yet.
|
||||
</Note>
|
||||
|
||||
### Config
|
||||
|
||||
All available parameters for the `aws_bedrock` config are present in [Master List of All Params in Config](../config).
|
||||
All available parameters for the `aws_bedrock` config are present in [Master List of All Params in Config](../config).
|
||||
|
||||
@@ -9,7 +9,8 @@ To use Sarvam AI's models, please set the `SARVAM_API_KEY` which you can get fro
|
||||
|
||||
## Usage
|
||||
|
||||
```python
|
||||
<CodeGroup>
|
||||
```python Python
|
||||
import os
|
||||
from mem0 import Memory
|
||||
|
||||
@@ -34,8 +35,35 @@ messages = [
|
||||
{"role": "assistant", "content": "Got it! I'll avoid thriller recommendations and suggest sci-fi movies in the future."}
|
||||
]
|
||||
m.add(messages, user_id="alex")
|
||||
|
||||
```
|
||||
|
||||
```typescript TypeScript
|
||||
import { Memory } from 'mem0ai/oss';
|
||||
|
||||
const config = {
|
||||
llm: {
|
||||
provider: 'sarvam',
|
||||
config: {
|
||||
apiKey: process.env.SARVAM_API_KEY || '',
|
||||
model: 'sarvam-m',
|
||||
temperature: 0.7,
|
||||
},
|
||||
},
|
||||
};
|
||||
|
||||
const memory = new Memory(config);
|
||||
const messages = [
|
||||
{"role": "user", "content": "I'm planning to watch a movie tonight. Any recommendations?"},
|
||||
{"role": "assistant", "content": "How about thriller movies? They can be quite engaging."},
|
||||
{"role": "user", "content": "I'm not a big fan of thriller movies but I love sci-fi movies."},
|
||||
{"role": "assistant", "content": "Got it! I'll avoid thriller recommendations and suggest sci-fi movies in the future."}
|
||||
];
|
||||
await memory.add(messages, { userId: 'alex' });
|
||||
```
|
||||
|
||||
</CodeGroup>
|
||||
|
||||
## Advanced Usage with Sarvam-Specific Features
|
||||
|
||||
```python
|
||||
|
||||
@@ -16,7 +16,7 @@ For a comprehensive list of available parameters for llm configuration, please r
|
||||
See the list of supported LLMs below.
|
||||
|
||||
<Note>
|
||||
All LLMs are supported in Python. The following LLMs are also supported in TypeScript: **OpenAI**, **Anthropic**, **Groq**, **Azure OpenAI**, **DeepSeek**, **Google AI**, **Langchain**, **LM Studio**, **Mistral AI**, and **Ollama**.
|
||||
All LLMs are supported in Python. The following LLMs are also supported in TypeScript: **OpenAI**, **Anthropic**, **AWS Bedrock**, **Groq**, **Azure OpenAI**, **DeepSeek**, **Google AI**, **Langchain**, **LM Studio**, **Mistral AI**, and **Ollama**.
|
||||
</Note>
|
||||
|
||||
<CardGroup cols={4}>
|
||||
|
||||
@@ -26,7 +26,7 @@ All rerankers share these common configuration parameters:
|
||||
|
||||
| Parameter | Description | Type | Default |
|
||||
| -------------------- | -------------------------------------------- | ------ | ----------------------- |
|
||||
| `model` | Cohere rerank model | `str` | `"rerank-english-v3.0"` |
|
||||
| `model` | Cohere rerank model | `str` | `"rerank-v3.5"` |
|
||||
| `api_key` | Cohere API key | `str` | `None` |
|
||||
| `return_documents` | Whether to return document texts in response | `bool` | `False` |
|
||||
| `max_chunks_per_doc` | Maximum chunks per document | `int` | `None` |
|
||||
@@ -103,3 +103,30 @@ config = {
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
## TypeScript SDK
|
||||
|
||||
The self-hosted [TypeScript SDK](/open-source/features/reranker-search#typescript-sdk) (`mem0ai/oss`) supports the same five providers. Config keys are camelCase (`apiKey`, `topK`, `maxLength`) and each provider's SDK is a peer dependency you install per reranker.
|
||||
|
||||
| Provider | Install | Default model | Key config fields |
|
||||
| --- | --- | --- | --- |
|
||||
| `cohere` | `pnpm add cohere-ai` | `rerank-v3.5` | `apiKey`, `model`, `topK` |
|
||||
| `zero_entropy` | `pnpm add zeroentropy` | `zerank-1` | `apiKey`, `model`, `topK` |
|
||||
| `sentence_transformer` | `pnpm add @huggingface/transformers` | `Xenova/ms-marco-MiniLM-L-6-v2` | `model`, `device`, `maxLength`, `normalize`, `topK` |
|
||||
| `huggingface` | `pnpm add @huggingface/transformers` | `Xenova/bge-reranker-base` | `model`, `device`, `maxLength`, `normalize`, `topK` |
|
||||
| `llm_reranker` | None (uses your LLM provider's own SDK) | `openai` / `gpt-4o-mini` | `provider`, `model`, `apiKey`, `llm` (nested override), `topK` |
|
||||
|
||||
```typescript
|
||||
import { Memory } from "mem0ai/oss";
|
||||
|
||||
const memory = new Memory({
|
||||
reranker: {
|
||||
provider: "zero_entropy",
|
||||
config: { apiKey: process.env.ZERO_ENTROPY_API_KEY, topK: 5 },
|
||||
},
|
||||
});
|
||||
```
|
||||
|
||||
<Note>
|
||||
The local cross-encoder providers (`sentence_transformer`, `huggingface`) run on [Transformers.js](https://huggingface.co/docs/transformers.js) and default to ONNX (`Xenova/*`) model mirrors, so Python default model strings must be swapped for their ONNX equivalents. `batchSize` and `showProgressBar` are accepted for parity with Python but are no-ops in the TypeScript runtime. See the [reranker feature guide](/open-source/features/reranker-search#typescript-sdk) for full examples.
|
||||
</Note>
|
||||
|
||||
@@ -9,9 +9,9 @@ Cohere provides enterprise-grade reranking models with excellent multilingual su
|
||||
|
||||
Cohere offers several reranking models:
|
||||
|
||||
- **`rerank-english-v3.0`**: Latest English reranker with best performance
|
||||
- **`rerank-multilingual-v3.0`**: Multilingual support for global applications
|
||||
- **`rerank-english-v2.0`**: Previous generation English reranker
|
||||
- **`rerank-v3.5`** (default): Latest reranker, multilingual, best performance
|
||||
- **`rerank-english-v3.0`**: Previous generation, English only
|
||||
- **`rerank-multilingual-v3.0`**: Previous generation, multilingual
|
||||
|
||||
## Installation
|
||||
|
||||
@@ -41,7 +41,7 @@ config = {
|
||||
"reranker": {
|
||||
"provider": "cohere",
|
||||
"config": {
|
||||
"model": "rerank-english-v3.0",
|
||||
"model": "rerank-v3.5",
|
||||
"api_key": "your-cohere-api-key", # or set COHERE_API_KEY
|
||||
"top_k": 5,
|
||||
"return_documents": False,
|
||||
@@ -53,6 +53,34 @@ config = {
|
||||
memory = Memory.from_config(config)
|
||||
```
|
||||
|
||||
## TypeScript (self-hosted)
|
||||
|
||||
The [TypeScript OSS SDK](/open-source/features/reranker-search#typescript-sdk) (`mem0ai/oss`) ships the Cohere reranker. Config keys are camelCase, it defaults to the `rerank-v3.5` model, and you opt in per search with `rerank: true`.
|
||||
|
||||
```bash
|
||||
pnpm add cohere-ai
|
||||
```
|
||||
|
||||
```typescript
|
||||
import { Memory } from "mem0ai/oss";
|
||||
|
||||
const memory = new Memory({
|
||||
reranker: {
|
||||
provider: "cohere",
|
||||
config: {
|
||||
apiKey: process.env.COHERE_API_KEY, // or set COHERE_API_KEY
|
||||
// model: "rerank-v3.5", // default
|
||||
topK: 5,
|
||||
},
|
||||
},
|
||||
});
|
||||
|
||||
const results = await memory.search("What is the user's profession?", {
|
||||
filters: { userId: "bob" },
|
||||
rerank: true,
|
||||
});
|
||||
```
|
||||
|
||||
## Environment Variables
|
||||
|
||||
Set your API key as an environment variable:
|
||||
@@ -77,7 +105,7 @@ config = {
|
||||
"rerank": {
|
||||
"provider": "cohere",
|
||||
"config": {
|
||||
"model": "rerank-english-v3.0",
|
||||
"model": "rerank-v3.5",
|
||||
"top_k": 3
|
||||
}
|
||||
}
|
||||
@@ -124,7 +152,7 @@ config = {
|
||||
|
||||
| Parameter | Description | Type | Default |
|
||||
| -------------------- | -------------------------------- | ------ | ----------------------- |
|
||||
| `model` | Cohere rerank model to use | `str` | `"rerank-english-v3.0"` |
|
||||
| `model` | Cohere rerank model to use | `str` | `"rerank-v3.5"` |
|
||||
| `api_key` | Cohere API key | `str` | `None` |
|
||||
| `top_k` | Maximum documents to return | `int` | `None` |
|
||||
| `return_documents` | Whether to return document texts | `bool` | `False` |
|
||||
@@ -139,7 +167,7 @@ config = {
|
||||
|
||||
## Best Practices
|
||||
|
||||
1. **Model Selection**: Use `rerank-english-v3.0` for English, `rerank-multilingual-v3.0` for other languages
|
||||
1. **Model Selection**: `rerank-v3.5` handles English and multilingual workloads; pin an older `v3.0` model only if you need to reproduce prior results
|
||||
2. **Batch Processing**: Process multiple queries efficiently
|
||||
3. **Error Handling**: Implement retry logic for production systems
|
||||
4. **Monitoring**: Track reranking performance and costs
|
||||
|
||||
@@ -57,6 +57,40 @@ config = {
|
||||
}
|
||||
```
|
||||
|
||||
## TypeScript (self-hosted)
|
||||
|
||||
The [TypeScript OSS SDK](/open-source/features/reranker-search#typescript-sdk) (`mem0ai/oss`) runs this reranker locally with [Transformers.js](https://huggingface.co/docs/transformers.js), the same cross-encoder path as `sentence_transformer`, just a different default model. It executes ONNX weights, so the default is the ONNX mirror `Xenova/bge-reranker-base`. Point `model` at any ONNX-exported reranker on the Hub (a raw `BAAI/bge-reranker-*` PyTorch checkpoint will not load in this runtime).
|
||||
|
||||
```bash
|
||||
pnpm add @huggingface/transformers
|
||||
```
|
||||
|
||||
```typescript
|
||||
import { Memory } from "mem0ai/oss";
|
||||
|
||||
const memory = new Memory({
|
||||
reranker: {
|
||||
provider: "huggingface",
|
||||
config: {
|
||||
// model: "Xenova/bge-reranker-base", // default (ONNX)
|
||||
device: "cpu", // "cpu" | "wasm" | "webgpu"
|
||||
maxLength: 512, // max tokens per query-document pair
|
||||
normalize: true, // sigmoid-normalize logits to [0, 1] (default)
|
||||
topK: 5,
|
||||
},
|
||||
},
|
||||
});
|
||||
|
||||
const results = await memory.search("What are the user's interests?", {
|
||||
filters: { userId: "alice" },
|
||||
rerank: true,
|
||||
});
|
||||
```
|
||||
|
||||
<Note>
|
||||
`batchSize` and `showProgressBar` are accepted for parity with the Python SDK but are no-ops in the TypeScript runtime. `trust_remote_code` and `model_kwargs` are Python-only.
|
||||
</Note>
|
||||
|
||||
## Popular Models
|
||||
|
||||
### BGE Rerankers (Recommended)
|
||||
|
||||
@@ -67,6 +67,43 @@ config = {
|
||||
}
|
||||
```
|
||||
|
||||
## TypeScript (self-hosted)
|
||||
|
||||
The [TypeScript OSS SDK](/open-source/features/reranker-search#typescript-sdk) (`mem0ai/oss`) ships the LLM reranker under the provider name `llm_reranker`. It does **not** reuse the Memory's main `llm` instance; it builds its own LLM from the reranker's own config, defaulting to `openai` / `gpt-4o-mini`. Set `provider`/`model`/`apiKey` directly on `config`, or nest a fully separate `config.llm: { provider, config }` (its `provider`/`config` take priority over the top-level fields, which only backfill values missing from the nested config).
|
||||
|
||||
```typescript
|
||||
import { Memory } from "mem0ai/oss";
|
||||
|
||||
const memory = new Memory({
|
||||
reranker: {
|
||||
provider: "llm_reranker",
|
||||
config: { apiKey: process.env.OPENAI_API_KEY },
|
||||
},
|
||||
});
|
||||
|
||||
const results = await memory.search("What movies do I like?", {
|
||||
filters: { userId: "alice" },
|
||||
rerank: true,
|
||||
});
|
||||
```
|
||||
|
||||
To rerank with a different LLM provider than the Memory's main `llm`, nest it under `config.llm`:
|
||||
|
||||
```typescript
|
||||
const memory = new Memory({
|
||||
llm: { provider: "openai", config: { apiKey: process.env.OPENAI_API_KEY } },
|
||||
reranker: {
|
||||
provider: "llm_reranker",
|
||||
config: {
|
||||
llm: {
|
||||
provider: "anthropic",
|
||||
config: { apiKey: process.env.ANTHROPIC_API_KEY },
|
||||
},
|
||||
},
|
||||
},
|
||||
});
|
||||
```
|
||||
|
||||
## Supported LLM Providers
|
||||
|
||||
### OpenAI
|
||||
|
||||
@@ -54,6 +54,40 @@ config = {
|
||||
memory = Memory.from_config(config)
|
||||
```
|
||||
|
||||
## TypeScript (self-hosted)
|
||||
|
||||
The [TypeScript OSS SDK](/open-source/features/reranker-search#typescript-sdk) (`mem0ai/oss`) runs this reranker locally with [Transformers.js](https://huggingface.co/docs/transformers.js). Because it executes ONNX weights, the default model is the ONNX mirror of the Python default: `Xenova/ms-marco-MiniLM-L-6-v2`. Point `model` at any ONNX-exported cross-encoder on the Hub (a raw `cross-encoder/...` PyTorch checkpoint will not load in this runtime).
|
||||
|
||||
```bash
|
||||
pnpm add @huggingface/transformers
|
||||
```
|
||||
|
||||
```typescript
|
||||
import { Memory } from "mem0ai/oss";
|
||||
|
||||
const memory = new Memory({
|
||||
reranker: {
|
||||
provider: "sentence_transformer",
|
||||
config: {
|
||||
// model: "Xenova/ms-marco-MiniLM-L-6-v2", // default (ONNX)
|
||||
device: "cpu", // "cpu" | "wasm" | "webgpu"
|
||||
maxLength: 512, // max tokens per query-document pair
|
||||
normalize: true, // sigmoid-normalize logits to [0, 1] (default)
|
||||
topK: 5,
|
||||
},
|
||||
},
|
||||
});
|
||||
|
||||
const results = await memory.search("What books does the user like?", {
|
||||
filters: { userId: "charlie" },
|
||||
rerank: true,
|
||||
});
|
||||
```
|
||||
|
||||
<Note>
|
||||
`batchSize` and `showProgressBar` are accepted for parity with the Python SDK but are no-ops in the TypeScript runtime, because a search reranks a small candidate set in a single in-process forward pass. The model downloads once and is cached in-process.
|
||||
</Note>
|
||||
|
||||
## GPU Acceleration
|
||||
|
||||
For better performance, use GPU acceleration:
|
||||
|
||||
@@ -50,6 +50,34 @@ config = {
|
||||
memory = Memory.from_config(config)
|
||||
```
|
||||
|
||||
## TypeScript (self-hosted)
|
||||
|
||||
The [TypeScript OSS SDK](/open-source/features/reranker-search#typescript-sdk) (`mem0ai/oss`) ships the Zero Entropy reranker under the same provider name as Python, `zero_entropy`. It reads the key from config or `ZERO_ENTROPY_API_KEY` and defaults to the `zerank-1` model.
|
||||
|
||||
```bash
|
||||
pnpm add zeroentropy
|
||||
```
|
||||
|
||||
```typescript
|
||||
import { Memory } from "mem0ai/oss";
|
||||
|
||||
const memory = new Memory({
|
||||
reranker: {
|
||||
provider: "zero_entropy",
|
||||
config: {
|
||||
apiKey: process.env.ZERO_ENTROPY_API_KEY,
|
||||
// model: "zerank-1", // default (or "zerank-1-small")
|
||||
topK: 5,
|
||||
},
|
||||
},
|
||||
});
|
||||
|
||||
const results = await memory.search("What Italian food does the user like?", {
|
||||
filters: { userId: "alice" },
|
||||
rerank: true,
|
||||
});
|
||||
```
|
||||
|
||||
## Environment Variables
|
||||
|
||||
Set your API key as an environment variable:
|
||||
|
||||
@@ -47,7 +47,7 @@ config = {
|
||||
"reranker": {
|
||||
"provider": "cohere",
|
||||
"config": {
|
||||
"model": "rerank-english-v3.0",
|
||||
"model": "rerank-v3.5",
|
||||
"top_n": 10,
|
||||
"max_chunks_per_doc": 10, # Limit chunk processing
|
||||
"return_documents": False # Reduce response size
|
||||
@@ -280,7 +280,7 @@ config = {
|
||||
```python
|
||||
def benchmark_rerankers():
|
||||
configs = [
|
||||
{"provider": "cohere", "model": "rerank-english-v3.0"},
|
||||
{"provider": "cohere", "model": "rerank-v3.5"},
|
||||
{"provider": "sentence_transformer", "model": "cross-encoder/ms-marco-MiniLM-L-6-v2"},
|
||||
{"provider": "huggingface", "model": "BAAI/bge-reranker-base"}
|
||||
]
|
||||
|
||||
@@ -19,6 +19,10 @@ Reranking trades extra latency for better precision. Start once you have baselin
|
||||
<Card title="Zero Entropy" icon="/images/provider-icons/zeroentropy.svg" href="/components/rerankers/models/zero_entropy" />
|
||||
</CardGroup>
|
||||
|
||||
<Note>
|
||||
All five rerankers are available in both the Python and the [TypeScript](/open-source/features/reranker-search#typescript-sdk) self-hosted SDKs. Each provider page has a **TypeScript (self-hosted)** section with the camelCase config.
|
||||
</Note>
|
||||
|
||||
## Reranking Workflow
|
||||
|
||||
<CardGroup cols={3}>
|
||||
|
||||
@@ -5,10 +5,22 @@ description: "Use Baidu Mochow as an enterprise vector database in Mem0 for high
|
||||
|
||||
[Baidu VectorDB](https://cloud.baidu.com/doc/VDB/index.html) is an enterprise-level distributed vector database service developed by Baidu Intelligent Cloud. It is powered by Baidu's proprietary "Mochow" vector database kernel, providing high performance, availability, and security for vector search.
|
||||
|
||||
### Installation
|
||||
|
||||
<CodeGroup>
|
||||
```bash Python
|
||||
pip install pymochow
|
||||
```
|
||||
|
||||
```bash TypeScript
|
||||
npm install @mochow/mochow-sdk-node
|
||||
```
|
||||
|
||||
</CodeGroup>
|
||||
|
||||
### Usage
|
||||
|
||||
```python
|
||||
import os
|
||||
from mem0 import Memory
|
||||
|
||||
config = {
|
||||
@@ -36,19 +48,63 @@ messages = [
|
||||
m.add(messages, user_id="alice", metadata={"category": "movies"})
|
||||
```
|
||||
|
||||
```typescript
|
||||
import { Memory } from "mem0ai/oss";
|
||||
|
||||
const memory = new Memory({
|
||||
embedder: {
|
||||
provider: "openai",
|
||||
config: {
|
||||
apiKey: process.env.OPENAI_API_KEY || "",
|
||||
model: "text-embedding-3-small",
|
||||
embeddingDims: 1536,
|
||||
},
|
||||
},
|
||||
vectorStore: {
|
||||
provider: "baidu",
|
||||
config: {
|
||||
endpoint: process.env.BAIDU_ENDPOINT || "",
|
||||
account: process.env.BAIDU_ACCOUNT || "root",
|
||||
apiKey: process.env.BAIDU_API_KEY || "",
|
||||
databaseName: "mem0",
|
||||
tableName: "mem0_table",
|
||||
embeddingModelDims: 1536,
|
||||
metricType: "COSINE",
|
||||
},
|
||||
},
|
||||
llm: {
|
||||
provider: "openai",
|
||||
config: {
|
||||
apiKey: process.env.OPENAI_API_KEY || "",
|
||||
model: "gpt-5-mini",
|
||||
},
|
||||
},
|
||||
});
|
||||
```
|
||||
|
||||
### Config
|
||||
|
||||
Here are the parameters available for configuring Baidu VectorDB:
|
||||
|
||||
| Parameter | Description | Default Value |
|
||||
| --- | --- | --- |
|
||||
| `endpoint` | Endpoint URL for your Baidu VectorDB instance | Required |
|
||||
| `account` | Baidu VectorDB account name | `root` |
|
||||
| `api_key` | API key for accessing Baidu VectorDB | Required |
|
||||
| `database_name` | Name of the database | `mem0` |
|
||||
| `table_name` | Name of the table | `mem0` |
|
||||
| `embedding_model_dims` | Dimensions of the embedding model | `1536` |
|
||||
| `metric_type` | Distance metric for similarity search | `L2` |
|
||||
| Parameter | Description | Default Value |
|
||||
| ---------------------- | --------------------------------------------- | ------------- |
|
||||
| `endpoint` | Endpoint URL for your Baidu VectorDB instance | Required |
|
||||
| `account` | Baidu VectorDB account name | `root` |
|
||||
| `api_key` | API key for accessing Baidu VectorDB | Required |
|
||||
| `database_name` | Name of the database | `mem0` |
|
||||
| `table_name` | Name of the table | `mem0` |
|
||||
| `embedding_model_dims` | Dimensions of the embedding model | `1536` |
|
||||
| `metric_type` | Distance metric for similarity search | `L2` |
|
||||
| `client` | Prebuilt Mochow client (TypeScript SDK only) | `None` |
|
||||
|
||||
For the TypeScript OSS SDK, use the camelCase equivalents:
|
||||
|
||||
- `databaseName`
|
||||
- `tableName`
|
||||
- `embeddingModelDims`
|
||||
- `metricType`
|
||||
|
||||
For OSS TS usage, `endpoint`, `account`, `apiKey`, `databaseName`, `tableName`, and `embeddingModelDims` are required unless you inject a prebuilt client. `metricType` defaults to `L2`, matching the Python SDK.
|
||||
|
||||
### Distance Metrics
|
||||
|
||||
@@ -66,3 +122,5 @@ The vector index is automatically configured with the following HNSW parameters:
|
||||
- `efconstruction`: 200 (size of the dynamic candidate list)
|
||||
- `auto_build`: true (automatically build index)
|
||||
- `auto_build_index_policy`: Incremental build with 10000 rows increment
|
||||
|
||||
The TypeScript provider also creates a BM25 inverted index over a `textLemmatized` column so `keywordSearch()` runs against a real full-text index. Mem0 lemmatizes the query before it reaches the vector store, so only the lemmatized form of each memory is indexed. If you point `tableName` at a table created before this index existed, `keywordSearch()` returns `null` and search falls back to vector similarity alone; recreate the table to enable it.
|
||||
|
||||
@@ -6,9 +6,8 @@ description: "Use Chroma as a vector database in Mem0 for local or cloud-hosted
|
||||
|
||||
### Usage
|
||||
|
||||
#### Local Installation
|
||||
|
||||
```python
|
||||
<CodeGroup>
|
||||
```python Python
|
||||
import os
|
||||
from mem0 import Memory
|
||||
|
||||
@@ -37,10 +36,46 @@ messages = [
|
||||
m.add(messages, user_id="alice", metadata={"category": "movies"})
|
||||
```
|
||||
|
||||
```typescript TypeScript
|
||||
import { Memory } from 'mem0ai/oss';
|
||||
|
||||
// The Node.js client connects to a running Chroma server.
|
||||
// Start one locally with: chroma run --host localhost --port 8000
|
||||
const config = {
|
||||
vectorStore: {
|
||||
provider: 'chroma',
|
||||
config: {
|
||||
collectionName: 'memories',
|
||||
host: 'localhost',
|
||||
port: 8000,
|
||||
// Optional: ChromaDB Cloud configuration
|
||||
// apiKey: 'your-chroma-cloud-api-key',
|
||||
// tenant: 'your-chroma-cloud-tenant-id',
|
||||
},
|
||||
},
|
||||
};
|
||||
|
||||
const memory = new Memory(config);
|
||||
const messages = [
|
||||
{"role": "user", "content": "I'm planning to watch a movie tonight. Any recommendations?"},
|
||||
{"role": "assistant", "content": "How about thriller movies? They can be quite engaging."},
|
||||
{"role": "user", "content": "I’m not a big fan of thriller movies but I love sci-fi movies."},
|
||||
{"role": "assistant", "content": "Got it! I'll avoid thriller recommendations and suggest sci-fi movies in the future."}
|
||||
]
|
||||
await memory.add(messages, { userId: "alice", metadata: { category: "movies" } });
|
||||
```
|
||||
</CodeGroup>
|
||||
|
||||
<Note>
|
||||
The Node.js SDK uses the `chromadb` v3 client, which talks to a Chroma server over HTTP (local server or ChromaDB Cloud). Install it with `npm install chromadb`. Mem0 supplies the embeddings, so the collection is created without an embedding function.
|
||||
</Note>
|
||||
|
||||
### Config
|
||||
|
||||
Here are the parameters available for configuring Chroma:
|
||||
|
||||
<Tabs>
|
||||
<Tab title="Python">
|
||||
| Parameter | Description | Default Value |
|
||||
| --- | --- | --- |
|
||||
| `collection_name` | The name of the collection | `mem0` |
|
||||
@@ -49,4 +84,19 @@ Here are the parameters available for configuring Chroma:
|
||||
| `host` | The host where the Chroma server is running | `None` |
|
||||
| `port` | The port where the Chroma server is running | `None` |
|
||||
| `api_key` | ChromaDB Cloud API key (for cloud usage) | `None` |
|
||||
| `tenant` | ChromaDB Cloud tenant ID (for cloud usage) | `None` |
|
||||
| `tenant` | ChromaDB Cloud tenant ID (for cloud usage) | `None` |
|
||||
</Tab>
|
||||
<Tab title="TypeScript">
|
||||
| Parameter | Description | Default Value |
|
||||
| --- | --- | --- |
|
||||
| `collectionName` | The name of the collection | `mem0` |
|
||||
| `client` | Pre-configured `ChromaClient` or `CloudClient` instance | `None` |
|
||||
| `host` | The host where the Chroma server is running | `None` |
|
||||
| `port` | The port where the Chroma server is running | `None` |
|
||||
| `ssl` | Whether to use SSL when connecting to the Chroma server | `false` |
|
||||
| `path` | Full URL of a Chroma server, e.g. `http://localhost:8000` (alternative to `host` and `port`) | `None` |
|
||||
| `apiKey` | ChromaDB Cloud API key (for cloud usage) | `None` |
|
||||
| `tenant` | ChromaDB Cloud tenant ID (for cloud usage) | `None` |
|
||||
| `database` | ChromaDB Cloud database name (for cloud usage) | `mem0` |
|
||||
</Tab>
|
||||
</Tabs>
|
||||
|
||||
@@ -6,7 +6,8 @@ description: "Use Databricks Vector Search as a serverless vector store in Mem0
|
||||
|
||||
### Usage
|
||||
|
||||
```python
|
||||
<CodeGroup>
|
||||
```python Python
|
||||
import os
|
||||
from mem0 import Memory
|
||||
|
||||
@@ -36,10 +37,44 @@ messages = [
|
||||
m.add(messages, user_id="alice", metadata={"category": "movies"})
|
||||
```
|
||||
|
||||
```typescript TypeScript
|
||||
// Requires the Databricks SQL driver (peer dependency): pnpm add @databricks/sql
|
||||
import { Memory } from 'mem0ai/oss';
|
||||
|
||||
const config = {
|
||||
vectorStore: {
|
||||
provider: 'databricks',
|
||||
config: {
|
||||
workspaceUrl: 'https://your-workspace.databricks.com',
|
||||
// SQL warehouse HTTP path, used for index writes (required)
|
||||
httpPath: '/sql/1.0/warehouses/your-warehouse-id',
|
||||
accessToken: 'your-access-token',
|
||||
catalog: 'your_catalog',
|
||||
schema: 'your_schema',
|
||||
tableName: 'your_table',
|
||||
collectionName: 'your_index_name',
|
||||
embeddingModelDims: 1536,
|
||||
},
|
||||
},
|
||||
};
|
||||
|
||||
const memory = new Memory(config);
|
||||
const messages = [
|
||||
{"role": "user", "content": "I'm planning to watch a movie tonight. Any recommendations?"},
|
||||
{"role": "assistant", "content": "How about thriller movies? They can be quite engaging."},
|
||||
{"role": "user", "content": "I'm not a big fan of thriller movies but I love sci-fi movies."},
|
||||
{"role": "assistant", "content": "Got it! I'll avoid thriller recommendations and suggest sci-fi movies in the future."}
|
||||
]
|
||||
await memory.add(messages, { userId: "alice", metadata: { category: "movies" } });
|
||||
```
|
||||
</CodeGroup>
|
||||
|
||||
### Config
|
||||
|
||||
Here are the parameters available for configuring Databricks Vector Search:
|
||||
|
||||
<Tabs>
|
||||
<Tab title="Python">
|
||||
| Parameter | Description | Default Value |
|
||||
| --- | --- | --- |
|
||||
| `workspace_url` | The URL of your Databricks workspace | **Required** |
|
||||
@@ -60,6 +95,32 @@ Here are the parameters available for configuring Databricks Vector Search:
|
||||
| `pipeline_type` | Sync pipeline type: `TRIGGERED` or `CONTINUOUS` | `TRIGGERED` |
|
||||
| `warehouse_name` | Databricks SQL warehouse name (if using SQL warehouse) | `None` |
|
||||
| `query_type` | Query type: `ANN` or `HYBRID` | `ANN` |
|
||||
</Tab>
|
||||
<Tab title="TypeScript">
|
||||
| Parameter | Description | Default Value |
|
||||
| --- | --- | --- |
|
||||
| `workspaceUrl` | The URL of your Databricks workspace (or pass `host`) | **Required** |
|
||||
| `httpPath` | SQL warehouse HTTP path, used for index writes | **Required** |
|
||||
| `accessToken` | Personal Access Token for authentication | `None` |
|
||||
| `clientId` | Service principal client ID (alternative to `accessToken`) | `None` |
|
||||
| `clientSecret` | Service principal client secret (required with `clientId`) | `None` |
|
||||
| `endpointName` | Name of the Vector Search endpoint | `mem0_vector_search` |
|
||||
| `endpointType` | Type of endpoint (`STANDARD` or `STORAGE_OPTIMIZED`) | `STANDARD` |
|
||||
| `pipelineType` | Delta Sync pipeline type: `TRIGGERED` or `CONTINUOUS` | `TRIGGERED` |
|
||||
| `queryType` | Query type: `ANN` or `HYBRID` | `ANN` |
|
||||
| `catalog` | Unity Catalog catalog name | `main` |
|
||||
| `schema` | Unity Catalog schema name | `default` |
|
||||
| `collectionName` | Vector Search index name | `mem0` |
|
||||
| `tableName` | Source Delta table name | falls back to `collectionName` |
|
||||
| `embeddingModelDims` | Dimension of self-managed embeddings | `1536` |
|
||||
| `syncPollIntervalMs` | Poll interval while waiting for a `TRIGGERED` sync | `1000` |
|
||||
| `syncTimeoutMs` | Timeout while waiting for an index sync | `300000` |
|
||||
|
||||
<Note>
|
||||
The TypeScript provider uses `DELTA_SYNC` indexes with self-managed embeddings: pass vectors directly. `DIRECT_ACCESS` indexes, Databricks-computed embeddings (`embedding_model_endpoint_name`), and Azure AD auth are Python-only today. It writes to the index through a SQL warehouse, so `httpPath` is required, and `@databricks/sql` must be installed as a peer dependency.
|
||||
</Note>
|
||||
</Tab>
|
||||
</Tabs>
|
||||
|
||||
### Authentication
|
||||
|
||||
|
||||
@@ -6,7 +6,14 @@ description: "Use Milvus as an open-source vector database in Mem0, scalable fro
|
||||
|
||||
### Usage
|
||||
|
||||
```python
|
||||
The TypeScript SDK loads the Milvus client lazily. Install it alongside `mem0ai` when you use this provider:
|
||||
|
||||
```bash
|
||||
npm install @zilliz/milvus2-sdk-node
|
||||
```
|
||||
|
||||
<CodeGroup>
|
||||
```python Python
|
||||
import os
|
||||
from mem0 import Memory
|
||||
|
||||
@@ -33,10 +40,39 @@ messages = [
|
||||
m.add(messages, user_id="alice", metadata={"category": "movies"})
|
||||
```
|
||||
|
||||
```typescript TypeScript
|
||||
import { Memory } from 'mem0ai/oss';
|
||||
|
||||
const config = {
|
||||
vectorStore: {
|
||||
provider: 'milvus',
|
||||
config: {
|
||||
collectionName: 'test',
|
||||
embeddingModelDims: 1536,
|
||||
url: 'http://localhost:19530',
|
||||
token: '8e4b8ca8cf2c67',
|
||||
dbName: 'my_database',
|
||||
},
|
||||
},
|
||||
};
|
||||
|
||||
const memory = new Memory(config);
|
||||
const messages = [
|
||||
{ role: "user", content: "I'm planning to watch a movie tonight. Any recommendations?" },
|
||||
{ role: "assistant", content: "How about thriller movies? They can be quite engaging." },
|
||||
{ role: "user", content: "I'm not a big fan of thriller movies but I love sci-fi movies." },
|
||||
{ role: "assistant", content: "Got it! I'll avoid thriller recommendations and suggest sci-fi movies in the future." },
|
||||
];
|
||||
await memory.add(messages, { userId: "alice", metadata: { category: "movies" } });
|
||||
```
|
||||
</CodeGroup>
|
||||
|
||||
### Config
|
||||
|
||||
Here are the parameters available for configuring Milvus:
|
||||
|
||||
<Tabs>
|
||||
<Tab title="Python">
|
||||
| Parameter | Description | Default Value |
|
||||
| --- | --- | --- |
|
||||
| `url` | Full URL/Uri for Milvus/Zilliz server | `http://localhost:19530` |
|
||||
@@ -45,3 +81,15 @@ Here are the parameters available for configuring Milvus:
|
||||
| `embedding_model_dims` | Dimensions of the embedding model | `1536` |
|
||||
| `metric_type` | Metric type for similarity search | `L2` |
|
||||
| `db_name` | Name of the database | `""` |
|
||||
</Tab>
|
||||
<Tab title="TypeScript">
|
||||
| Parameter | Description | Default Value |
|
||||
| --- | --- | --- |
|
||||
| `url` | Full URL/Uri for Milvus/Zilliz server | `http://localhost:19530` |
|
||||
| `token` | Token for Zilliz Cloud (optional for a local setup) | `undefined` |
|
||||
| `collectionName` | The name of the collection | `mem0` |
|
||||
| `embeddingModelDims` | Dimensions of the embedding model | `1536` |
|
||||
| `metricType` | Metric type for similarity search (`L2`, `IP`, `COSINE`, `HAMMING`, `JACCARD`) | `L2` |
|
||||
| `dbName` | Name of the database | `undefined` |
|
||||
</Tab>
|
||||
</Tabs>
|
||||
|
||||
@@ -2,26 +2,37 @@
|
||||
title: "Neptune Analytics"
|
||||
description: "Use AWS Neptune Analytics as a vector store in Mem0, combining graph analytics with vector search capabilities."
|
||||
---
|
||||
# Neptune Analytics Vector Store
|
||||
|
||||
[Neptune Analytics](https://docs.aws.amazon.com/neptune-analytics/latest/userguide/what-is-neptune-analytics.html/) is a memory-optimized graph database engine for analytics. With Neptune Analytics, you can get insights and find trends by processing large amounts of graph data in seconds, including vector search.
|
||||
[Neptune Analytics](https://docs.aws.amazon.com/neptune-analytics/latest/userguide/what-is-neptune-analytics.html) is a memory-optimized graph database engine for analytics. With Neptune Analytics, you can get insights and find trends by processing large amounts of graph data in seconds, including vector search.
|
||||
|
||||
### Installation
|
||||
|
||||
## Installation
|
||||
The Neptune Analytics provider needs the AWS Neptune Graph client. Install it alongside `mem0ai`:
|
||||
|
||||
```bash
|
||||
<CodeGroup>
|
||||
```bash Python
|
||||
pip install mem0ai[vector-stores]
|
||||
```
|
||||
|
||||
## Usage
|
||||
```bash TypeScript
|
||||
npm install @aws-sdk/client-neptune-graph
|
||||
```
|
||||
</CodeGroup>
|
||||
|
||||
### Usage
|
||||
|
||||
Configure AWS credentials in your environment (environment variables, shared config file, an IAM role, or an instance profile). Both SDKs pick them up automatically through the standard AWS credential chain.
|
||||
|
||||
<CodeGroup>
|
||||
```python Python
|
||||
from mem0 import Memory
|
||||
|
||||
```python
|
||||
config = {
|
||||
"vector_store": {
|
||||
"provider": "neptune",
|
||||
"config": {
|
||||
"collection_name": "mem0",
|
||||
"endpoint": f"neptune-graph://my-graph-identifier",
|
||||
"endpoint": "neptune-graph://g-abc123xyz0",
|
||||
},
|
||||
},
|
||||
}
|
||||
@@ -29,18 +40,90 @@ config = {
|
||||
m = Memory.from_config(config)
|
||||
messages = [
|
||||
{"role": "user", "content": "I'm planning to watch a movie tonight. Any recommendations?"},
|
||||
{"role": "assistant", "content": "How about a thriller movies? They can be quite engaging."},
|
||||
{"role": "assistant", "content": "How about a thriller movie? They can be quite engaging."},
|
||||
{"role": "user", "content": "I'm not a big fan of thriller movies but I love sci-fi movies."},
|
||||
{"role": "assistant", "content": "Got it! I'll avoid thriller recommendations and suggest sci-fi movies in the future."}
|
||||
]
|
||||
m.add(messages, user_id="alice", metadata={"category": "movies"})
|
||||
```
|
||||
|
||||
## Parameters
|
||||
```typescript TypeScript
|
||||
import { Memory } from 'mem0ai/oss';
|
||||
|
||||
Let's see the available parameters for the `neptune` config:
|
||||
const config = {
|
||||
vectorStore: {
|
||||
provider: 'neptune',
|
||||
config: {
|
||||
collectionName: 'mem0',
|
||||
graphIdentifier: 'g-abc123xyz0',
|
||||
// Any other key here (region, credentials, maxAttempts, ...) is
|
||||
// forwarded to the underlying NeptuneGraphClient constructor.
|
||||
region: 'us-east-1',
|
||||
},
|
||||
},
|
||||
};
|
||||
|
||||
const memory = new Memory(config);
|
||||
const messages = [
|
||||
{ role: "user", content: "I'm planning to watch a movie tonight. Any recommendations?" },
|
||||
{ role: "assistant", content: "How about a thriller movie? They can be quite engaging." },
|
||||
{ role: "user", content: "I'm not a big fan of thriller movies but I love sci-fi movies." },
|
||||
{ role: "assistant", content: "Got it! I'll avoid thriller recommendations and suggest sci-fi movies in the future." },
|
||||
];
|
||||
await memory.add(messages, { userId: "alice", metadata: { category: "movies" } });
|
||||
```
|
||||
</CodeGroup>
|
||||
|
||||
### Config
|
||||
|
||||
<Tabs>
|
||||
<Tab title="Python">
|
||||
| Parameter | Description | Default Value |
|
||||
| --- | --- | --- |
|
||||
| `collection_name` | The name of the collection to store the vectors | `mem0` |
|
||||
| `endpoint` | Connection URL for the Neptune Analytics service | `neptune-graph://my-graph-identifier` |
|
||||
| `endpoint` | Connection URL for the Neptune Analytics service, must be `neptune-graph://<graph-id>` | Required |
|
||||
</Tab>
|
||||
<Tab title="TypeScript">
|
||||
| Parameter | Description | Default Value |
|
||||
| --- | --- | --- |
|
||||
| `collectionName` | The name of the collection to store the vectors | `memories` |
|
||||
| `graphIdentifier` | Graph ID, e.g. `g-abc123xyz0`. Takes priority over `endpoint`. | Required, unless `endpoint` supplies it |
|
||||
| `endpoint` | Either `neptune-graph://<graph-id>` (or a bare graph ID) to supply the graph ID, or an `https://` service endpoint to override the AWS endpoint. An `https://` value must be paired with `graphIdentifier`. | `undefined` |
|
||||
| `dimension` | Embedding vector dimension | Auto-detected from the embedder when omitted |
|
||||
| `client` | A pre-built `NeptuneGraphClient` to use instead of constructing one | `undefined` |
|
||||
| any other key | Forwarded as-is to the [`NeptuneGraphClient`](https://www.npmjs.com/package/@aws-sdk/client-neptune-graph) constructor, e.g. `region`, `credentials`, `maxAttempts` | N/A |
|
||||
</Tab>
|
||||
</Tabs>
|
||||
|
||||
Both SDKs store vectors on graph nodes labeled `MEM0_VECTOR_<collection_name>`. Point them at the same
|
||||
graph with the same `collection_name` — the defaults differ, `mem0` in Python and `memories` in
|
||||
TypeScript — and `get()`, `list()`, and `delete()` interoperate across SDKs.
|
||||
|
||||
<Note>
|
||||
`search()` is not currently cross-SDK compatible. The TypeScript provider filters on Neptune's reserved
|
||||
`~label` metafield, while the Python provider filters on a synthetic `label` property that only Python's
|
||||
own `insert()` writes. Python's `search()` therefore cannot see nodes written by the TypeScript provider.
|
||||
</Note>
|
||||
|
||||
### IAM Permissions
|
||||
|
||||
Your AWS identity (user or role) needs a policy that allows the [`ExecuteQuery`](https://docs.aws.amazon.com/neptune-analytics/latest/apiref/API_ExecuteQuery.html) actions used for reads, writes, and deletes:
|
||||
|
||||
```json
|
||||
{
|
||||
"Version": "2012-10-17",
|
||||
"Statement": [
|
||||
{
|
||||
"Effect": "Allow",
|
||||
"Action": [
|
||||
"neptune-graph:ReadDataViaQuery",
|
||||
"neptune-graph:WriteDataViaQuery",
|
||||
"neptune-graph:DeleteDataViaQuery"
|
||||
],
|
||||
"Resource": "*"
|
||||
}
|
||||
]
|
||||
}
|
||||
```
|
||||
|
||||
For production, scope the resource ARN down to your specific graph.
|
||||
|
||||
@@ -42,7 +42,7 @@ const config = {
|
||||
provider: 'qdrant',
|
||||
config: {
|
||||
collectionName: 'memories',
|
||||
embeddingModelDims: 1536,
|
||||
dimension: 1536,
|
||||
host: 'localhost',
|
||||
port: 6333,
|
||||
},
|
||||
@@ -83,7 +83,7 @@ Let's see the available parameters for the `qdrant` config:
|
||||
| Parameter | Description | Default Value |
|
||||
| --- | --- | --- |
|
||||
| `collectionName` | The name of the collection to store the vectors | `mem0` |
|
||||
| `embeddingModelDims` | Dimensions of the embedding model | `1536` |
|
||||
| `dimension` | Dimensions of the embedding model | `1536` |
|
||||
| `host` | The host where the Qdrant server is running | `None` |
|
||||
| `port` | The port where the Qdrant server is running | `None` |
|
||||
| `path` | Path for the Qdrant database | `/tmp/qdrant` |
|
||||
|
||||
@@ -4,14 +4,21 @@ description: "Use Weaviate as an open-source vector search engine in Mem0 for st
|
||||
---
|
||||
[Weaviate](https://weaviate.io/) is an open-source vector search engine. It allows efficient storage and retrieval of high-dimensional vector embeddings, enabling powerful search and retrieval capabilities.
|
||||
|
||||
|
||||
### Installation
|
||||
```bash
|
||||
|
||||
<CodeGroup>
|
||||
```bash Python
|
||||
pip install weaviate-client
|
||||
```
|
||||
|
||||
```bash TypeScript
|
||||
npm install weaviate-client
|
||||
```
|
||||
</CodeGroup>
|
||||
|
||||
### Usage
|
||||
|
||||
<CodeGroup>
|
||||
```python Python
|
||||
import os
|
||||
from mem0 import Memory
|
||||
@@ -33,20 +40,73 @@ m = Memory.from_config(config)
|
||||
messages = [
|
||||
{"role": "user", "content": "I'm planning to watch a movie tonight. Any recommendations?"},
|
||||
{"role": "assistant", "content": "How about a thriller movie? They can be quite engaging."},
|
||||
{"role": "user", "content": "I’m not a big fan of thriller movies but I love sci-fi movies."},
|
||||
{"role": "user", "content": "I'm not a big fan of thriller movies but I love sci-fi movies."},
|
||||
{"role": "assistant", "content": "Got it! I'll avoid thriller recommendations and suggest sci-fi movies in the future."}
|
||||
]
|
||||
m.add(messages, user_id="alice", metadata={"category": "movies"})
|
||||
```
|
||||
|
||||
```typescript TypeScript
|
||||
import { Memory } from "mem0ai/oss";
|
||||
|
||||
const config = {
|
||||
vectorStore: {
|
||||
provider: "weaviate",
|
||||
config: {
|
||||
collectionName: "test",
|
||||
embeddingModelDims: 1536,
|
||||
clusterUrl: "http://localhost:8080",
|
||||
},
|
||||
},
|
||||
};
|
||||
|
||||
const memory = new Memory(config);
|
||||
|
||||
const messages = [
|
||||
{
|
||||
role: "user",
|
||||
content: "I'm planning to watch a movie tonight. Any recommendations?",
|
||||
},
|
||||
{
|
||||
role: "assistant",
|
||||
content: "How about a thriller movie? They can be quite engaging.",
|
||||
},
|
||||
{
|
||||
role: "user",
|
||||
content: "I'm not a big fan of thriller movies but I love sci-fi movies.",
|
||||
},
|
||||
{
|
||||
role: "assistant",
|
||||
content:
|
||||
"Got it! I'll avoid thriller recommendations and suggest sci-fi movies in the future.",
|
||||
},
|
||||
];
|
||||
|
||||
await memory.add(messages, {
|
||||
userId: "alice",
|
||||
metadata: {
|
||||
category: "movies",
|
||||
},
|
||||
});
|
||||
```
|
||||
</CodeGroup>
|
||||
|
||||
The TypeScript SDK picks the connection mode from the config you pass:
|
||||
|
||||
- `clusterUrl` pointing at `localhost` connects to a local instance.
|
||||
- `clusterUrl` plus `apiKey` connects to a Weaviate Cloud cluster (for example `https://my-cluster.weaviate.cloud`).
|
||||
- Any other `clusterUrl` without an `apiKey` connects to a custom deployment, using the host and port from the URL.
|
||||
|
||||
You can also pass a pre-configured `client` (a `WeaviateClient` instance) to reuse an existing connection.
|
||||
|
||||
### Config
|
||||
|
||||
Here are the parameters available for configuring Weaviate:
|
||||
|
||||
| Parameter | Description | Default Value |
|
||||
| --- | --- | --- |
|
||||
| `collection_name` | The name of the collection to store the vectors | `mem0` |
|
||||
| `embedding_model_dims` | Dimensions of the embedding model | `1536` |
|
||||
| `cluster_url` | URL for the Weaviate server | `None` |
|
||||
| `auth_client_secret` | API key for Weaviate authentication | `None` |
|
||||
| `additional_headers` | Additional headers to include in requests (`Dict[str, str]`) | `None` |
|
||||
| Python | TypeScript | Description | Default Value |
|
||||
| --- | --- | --- | --- |
|
||||
| `collection_name` | `collectionName` | The name of the collection to store the vectors | `mem0` |
|
||||
| `embedding_model_dims` | `embeddingModelDims` | Dimensions of the embedding model | `1536` |
|
||||
| `cluster_url` | `clusterUrl` | URL for the Weaviate server | `None` |
|
||||
| `auth_client_secret` | `apiKey` | API key for Weaviate authentication | `None` |
|
||||
| `additional_headers` | `additionalHeaders` | Additional headers to include in requests | `None` |
|
||||
|
||||
@@ -10,7 +10,7 @@ Mem0 includes built-in support for various popular databases. Memory can utilize
|
||||
See the list of supported vector databases below.
|
||||
|
||||
<Note>
|
||||
The following vector databases are supported in the Python implementation. The TypeScript implementation currently supports Qdrant, Redis, PGVector, Supabase, LangChain, Azure AI Search, Vectorize, Amazon S3 Vectors, and an in-memory store.
|
||||
The following vector databases are supported in the Python implementation. The TypeScript implementation currently supports Qdrant, Redis, PGVector, Supabase, LangChain, Azure AI Search, Vectorize, Amazon S3 Vectors, Milvus, Neptune Analytics, and an in-memory store.
|
||||
</Note>
|
||||
|
||||
<CardGroup cols={3}>
|
||||
@@ -32,6 +32,7 @@ See the list of supported vector databases below.
|
||||
<Card title="FAISS" icon="layer-group" href="/components/vectordbs/dbs/faiss"></Card>
|
||||
<Card title="LangChain" icon="/images/provider-icons/langchain-color.svg" href="/components/vectordbs/dbs/langchain"></Card>
|
||||
<Card title="Amazon S3 Vectors" icon="/images/provider-icons/aws-color.svg" href="/components/vectordbs/dbs/s3_vectors"></Card>
|
||||
<Card title="Neptune Analytics" icon="/images/provider-icons/aws-color.svg" href="/components/vectordbs/dbs/neptune_analytics"></Card>
|
||||
<Card title="Databricks" icon="/images/provider-icons/databricks.svg" href="/components/vectordbs/dbs/databricks"></Card>
|
||||
<Card title="Turbopuffer" icon="/images/provider-icons/turbopuffer.svg" href="/components/vectordbs/dbs/turbopuffer"></Card>
|
||||
</CardGroup>
|
||||
|
||||
@@ -588,9 +588,9 @@ mem0_client.project.update(
|
||||
Exclude: greetings, filler, casual chat
|
||||
""",
|
||||
custom_categories=[
|
||||
{"name": "goals", "description": "Training targets"},
|
||||
{"name": "constraints", "description": "Injuries and limitations"},
|
||||
{"name": "preferences", "description": "Training style"}
|
||||
{"goals": "Training targets"},
|
||||
{"constraints": "Injuries and limitations"},
|
||||
{"preferences": "Training style"}
|
||||
]
|
||||
)
|
||||
```
|
||||
|
||||
@@ -19,7 +19,7 @@ client = MemoryClient(api_key="your-api-key")
|
||||
```
|
||||
|
||||
<Note>
|
||||
Define custom categories at the **project level** with `client.project.update()` before adding memories. Categories apply to all future memories: Mem0 auto-assigns them based on content semantics.
|
||||
Define custom categories at the **project level** with `client.project.update()` before adding memories. Categories apply to all future memories: Mem0 auto-assigns them based on content semantics. You can also pass `custom_categories` on a single `client.add()` call to override the project list for just those memories. See [Custom Categories](/platform/features/custom-categories).
|
||||
</Note>
|
||||
|
||||
---
|
||||
@@ -96,6 +96,10 @@ Start with 3-5 clear categories that match how your team thinks. Too many catego
|
||||
|
||||
These categories are now available project-wide. Every memory can be tagged with one or more categories.
|
||||
|
||||
<Tip>
|
||||
Need a different vocabulary for one tenant or one kind of conversation? Pass `custom_categories=[...]` to `client.add()`. That list replaces the project list for the memories created by that call, and it does not change the project configuration.
|
||||
</Tip>
|
||||
|
||||
---
|
||||
|
||||
## Tagging Memories
|
||||
|
||||
@@ -15,6 +15,7 @@ Adding memory is how Mem0 captures useful details from a conversation so your ag
|
||||
- **Infer**: Controls whether Mem0 extracts structured memories (`infer=True`, default) or stores raw messages.
|
||||
- **Metadata**: Optional filters (e.g., `{"category": "movie_recommendations"}`) that improve retrieval later.
|
||||
- **User / Session identifiers**: `user_id`, `agent_id`, `app_id`, or `run_id` that scope the memory for future searches.
|
||||
- **expiration_date**: Optional `YYYY-MM-DD` date after which the memory is treated as expired. Use `expirationDate` in the JavaScript SDKs. Expired memories are hidden from `search` and `get_all` unless you pass `show_expired` (`showExpired` in JavaScript); fetching by ID still returns them.
|
||||
|
||||
## How does it work?
|
||||
|
||||
@@ -105,6 +106,9 @@ result = m.add(messages, user_id="alice", metadata={"category": "movie_recommend
|
||||
|
||||
# Optionally store raw messages without inference
|
||||
result = m.add(messages, user_id="alice", metadata={"category": "movie_recommendations"}, infer=False)
|
||||
|
||||
# Optionally set an expiration date (YYYY-MM-DD)
|
||||
result = m.add(messages, user_id="alice", expiration_date="2030-01-31")
|
||||
```
|
||||
|
||||
```javascript JavaScript
|
||||
@@ -123,6 +127,12 @@ const result = memory.add(messages, {
|
||||
userId: "alice",
|
||||
metadata: { category: "preferences" }
|
||||
});
|
||||
|
||||
// Optionally set an expiration date (YYYY-MM-DD)
|
||||
const expiring = memory.add(messages, {
|
||||
userId: "alice",
|
||||
expirationDate: "2030-01-31",
|
||||
});
|
||||
```
|
||||
</CodeGroup>
|
||||
|
||||
|
||||
@@ -188,11 +188,16 @@ memory = Memory()
|
||||
memory.delete(memory_id="mem_123")
|
||||
memory.delete_all(user_id="alice")
|
||||
```
|
||||
</CodeGroup>
|
||||
|
||||
<Note>
|
||||
The OSS JavaScript SDK does not yet expose deletion helpers: use the REST API or Python SDK when self-hosting.
|
||||
</Note>
|
||||
```typescript TypeScript
|
||||
import { Memory } from "mem0ai/oss";
|
||||
|
||||
const memory = new Memory();
|
||||
|
||||
await memory.delete("mem_123");
|
||||
await memory.deleteAll({ userId: "alice" });
|
||||
```
|
||||
</CodeGroup>
|
||||
|
||||
## Use cases recap
|
||||
|
||||
|
||||
@@ -12,7 +12,7 @@ Mem0’s update operation lets you fix or enrich an existing memory without dele
|
||||
## Key terms
|
||||
|
||||
- **memory_id**: Unique identifier returned by `add` or `search` results.
|
||||
- **text** / **data**: New content that replaces the stored memory value.
|
||||
- **text**: New content that replaces the stored memory value. In the Python OSS SDK, `data` is a deprecated alias for `text`.
|
||||
- **metadata**: Optional key-value pairs you update alongside the text.
|
||||
- **timestamp**: Unix epoch (int/float) or ISO 8601 string to override the memory's timestamp.
|
||||
- **batch_update**: Platform API that edits multiple memories in a single request.
|
||||
@@ -110,17 +110,47 @@ from mem0 import Memory
|
||||
|
||||
memory = Memory()
|
||||
|
||||
# Replace the content
|
||||
memory.update(
|
||||
memory_id="mem_123",
|
||||
data="Alex now prefers decaf coffee",
|
||||
text="Alex now prefers decaf coffee",
|
||||
)
|
||||
|
||||
# Update content plus metadata and an expiration date (None clears it)
|
||||
memory.update(
|
||||
memory_id="mem_123",
|
||||
text="Alex now prefers decaf coffee",
|
||||
metadata={"category": "preferences"},
|
||||
expiration_date="2030-01-31",
|
||||
)
|
||||
```
|
||||
```
|
||||
|
||||
```javascript JavaScript
|
||||
import { Memory } from "mem0ai/oss";
|
||||
|
||||
const memory = new Memory();
|
||||
|
||||
// Replace the content
|
||||
await memory.update("mem_123", { text: "Alex now prefers decaf coffee" });
|
||||
|
||||
// Update content plus metadata and an expiration date (null clears it)
|
||||
await memory.update("mem_123", {
|
||||
text: "Alex now prefers decaf coffee",
|
||||
metadata: { category: "preferences" },
|
||||
expirationDate: "2030-01-31",
|
||||
});
|
||||
|
||||
// Update metadata only, leaving the stored text untouched
|
||||
await memory.update("mem_123", { metadata: { category: "preferences" } });
|
||||
```
|
||||
</CodeGroup>
|
||||
|
||||
<Note>
|
||||
OSS JavaScript SDK does not expose `update` yet: use the REST API or Python SDK when self-hosting.
|
||||
In both OSS SDKs the content is optional: pass only `metadata` and/or an expiration date to update those while keeping the existing content. At least one of the three must be provided, otherwise the call raises.
|
||||
</Note>
|
||||
|
||||
<Note>
|
||||
`data` is a deprecated alias for `text` in both OSS SDKs (`data=` in Python, `{ data: ... }` in JavaScript). It still works but logs a warning; prefer `text`. In JavaScript, passing a bare string is shorthand for `{ text }`, so `update(memoryId, "new text")` also still works.
|
||||
</Note>
|
||||
|
||||
## Tips
|
||||
@@ -138,7 +168,7 @@ memory.update(
|
||||
|
||||
| Capability | Mem0 Platform | Mem0 OSS |
|
||||
| --- | --- | --- |
|
||||
| Update call | `client.update(memory_id, {...})` | `memory.update(memory_id, data=...)` |
|
||||
| Update call | `client.update(memory_id, {...})` | `memory.update(memory_id, text=...)` |
|
||||
| Batch updates | `client.batch_update` (up to 1000 memories) | Script your own loop or bulk job |
|
||||
| Dashboard visibility | Inspect updates in the UI | Inspect via logs or custom tooling |
|
||||
| Immutable handling | Returns descriptive error | Raises exception: delete and re-add |
|
||||
|
||||
+4
-2
@@ -96,7 +96,8 @@
|
||||
"pages": [
|
||||
"platform/features/direct-import",
|
||||
"platform/features/memory-export",
|
||||
"platform/features/timestamp"
|
||||
"platform/features/timestamp",
|
||||
"platform/features/memory-expiration"
|
||||
]
|
||||
},
|
||||
{
|
||||
@@ -152,7 +153,8 @@
|
||||
"open-source/features/multimodal-support",
|
||||
"open-source/features/custom-instructions",
|
||||
"open-source/features/rest-api",
|
||||
"open-source/features/openai_compatibility"
|
||||
"open-source/features/openai_compatibility",
|
||||
"platform/features/memory-expiration"
|
||||
]
|
||||
},
|
||||
{
|
||||
|
||||
+5
-3
@@ -61,7 +61,7 @@ client.get_all(user_id="alice")
|
||||
client.get(memory_id="<id>")
|
||||
|
||||
# Update
|
||||
client.update(memory_id="<id>", data="Alice loves mountain hiking")
|
||||
client.update(memory_id="<id>", text="Alice loves mountain hiking")
|
||||
|
||||
# Delete
|
||||
client.delete(memory_id="<id>")
|
||||
@@ -118,7 +118,7 @@ m.get_all(user_id="alice")
|
||||
m.get(memory_id="<id>")
|
||||
|
||||
# Update
|
||||
m.update(memory_id="<id>", data="Alice loves mountain hiking")
|
||||
m.update(memory_id="<id>", text="Alice loves mountain hiking")
|
||||
|
||||
# Delete
|
||||
m.delete(memory_id="<id>")
|
||||
@@ -215,6 +215,7 @@ If the user is on a pre-current major (Python < 2, TS < 3, or Platform `output_f
|
||||
- [Direct Import](https://docs.mem0.ai/platform/features/direct-import) [Platform]: Use when seeding a Mem0 project from existing data.
|
||||
- [Memory Export](https://docs.mem0.ai/platform/features/memory-export) [Platform]: Use when exporting memories via a Pydantic schema.
|
||||
- [Timestamp Support](https://docs.mem0.ai/platform/features/timestamp) [Platform]: Use when temporal queries or time-based filtering matter.
|
||||
- [Memory Expiration](https://docs.mem0.ai/platform/features/memory-expiration) [Both]: Use when a memory should stop surfacing after a known date without being deleted, e.g. trial facts, seasonal preferences, or retention windows.
|
||||
|
||||
### Features - Integration & Ops
|
||||
- [Webhooks](https://docs.mem0.ai/platform/features/webhooks) [Platform]: Use when another system needs to react to memory changes in real time.
|
||||
@@ -501,5 +502,6 @@ Everything below is OSS-only provider configuration. Skip this entire section wh
|
||||
- [Custom Reranker Prompts](https://docs.mem0.ai/components/rerankers/custom-prompts) [OSS]: Use when rewriting reranker prompts.
|
||||
- [Cohere Reranker](https://docs.mem0.ai/components/rerankers/models/cohere) [OSS]: Use for Cohere Rerank.
|
||||
- [Sentence Transformer Reranker](https://docs.mem0.ai/components/rerankers/models/sentence_transformer) [OSS]: Use for local cross-encoder rerankers.
|
||||
- [Hugging Face Reranker](https://docs.mem0.ai/components/rerankers/models/huggingface) [OSS]: Use for HF-hosted reranker models.- [LLM Reranker](https://docs.mem0.ai/components/rerankers/models/llm_reranker) [OSS]: Use when the reranker is a prompted LLM (implementation reference).
|
||||
- [Hugging Face Reranker](https://docs.mem0.ai/components/rerankers/models/huggingface) [OSS]: Use for HF-hosted reranker models.
|
||||
- [LLM Reranker](https://docs.mem0.ai/components/rerankers/models/llm_reranker) [OSS]: Use when the reranker is a prompted LLM (implementation reference).
|
||||
- [Zero Entropy Reranker](https://docs.mem0.ai/components/rerankers/models/zero_entropy) [OSS]: Use for the Zero Entropy reranker.
|
||||
|
||||
@@ -339,21 +339,22 @@ The Platform introduces powerful capabilities not available in OSS:
|
||||
</Info>
|
||||
```python
|
||||
# Set custom categories for your project
|
||||
client.projects.update_categories(
|
||||
project_id="proj_123",
|
||||
categories=[
|
||||
"Customer Preferences",
|
||||
"Product Feedback",
|
||||
"Support Issues",
|
||||
"Feature Requests"
|
||||
client.project.update(
|
||||
custom_categories=[
|
||||
{"customer_preferences": "Likes, dislikes, and product preferences"},
|
||||
{"product_feedback": "Feature requests and complaints about the product"},
|
||||
{"support_issues": "Problems reported and how they were resolved"}
|
||||
]
|
||||
)
|
||||
|
||||
# Memories will use these categories
|
||||
# Mem0 assigns these categories automatically as memories come in
|
||||
client.add("User wants dark mode in dashboard", user_id="alex")
|
||||
|
||||
# Or pass a different catalog for a single call
|
||||
client.add(
|
||||
"User wants dark mode in dashboard",
|
||||
user_id="alex",
|
||||
categories=["Customer Preferences"]
|
||||
custom_categories=[{"ui_requests": "Requests about interface and appearance"}]
|
||||
)
|
||||
```
|
||||
</Accordion>
|
||||
|
||||
@@ -59,7 +59,7 @@ config = {
|
||||
},
|
||||
"reranker": {
|
||||
"provider": "cohere",
|
||||
"config": {"model": "rerank-english-v3.0"},
|
||||
"config": {"model": "rerank-v3.5"},
|
||||
},
|
||||
}
|
||||
|
||||
@@ -119,9 +119,9 @@ Change the `provider` string to switch backends. The most common options:
|
||||
|
||||
| Component | Python | TypeScript |
|
||||
| --- | --- | --- |
|
||||
| LLM | `openai`, `anthropic`, `gemini`, `groq`, `ollama`, `aws_bedrock`, `azure_openai`, `litellm` | `openai`, `anthropic`, `gemini`, `groq`, `ollama`, `azure_openai`, `mistral`, `deepseek` |
|
||||
| LLM | `openai`, `anthropic`, `gemini`, `groq`, `ollama`, `aws_bedrock`, `azure_openai`, `litellm` | `openai`, `anthropic`, `gemini`, `groq`, `ollama`, `aws_bedrock`, `azure_openai`, `mistral`, `deepseek` |
|
||||
| Embedder | `openai`, `gemini`, `azure_openai`, `ollama`, `huggingface`, `vertexai`, `aws_bedrock` | `openai`, `gemini`, `azure_openai`, `ollama` |
|
||||
| Vector store | `qdrant`, `pgvector`, `chroma`, `pinecone`, `redis`, `weaviate`, `milvus`, `elasticsearch` | `memory`, `qdrant`, `pgvector`, `redis`, `supabase`, `azure-ai-search`, `vectorize` |
|
||||
| Vector store | `qdrant`, `pgvector`, `chroma`, `pinecone`, `redis`, `weaviate`, `milvus`, `elasticsearch` | `memory`, `qdrant`, `pgvector`, `redis`, `supabase`, `azure-ai-search`, `vectorize`, `milvus` |
|
||||
|
||||
See the full catalog in <Link href="/components/llms/overview">Components</Link>.
|
||||
|
||||
|
||||
@@ -36,7 +36,7 @@ icon: "bolt"
|
||||
| Search memories | `await memory.search(...)` | Returns dict with `results`, identical shape. |
|
||||
| List memories | `await memory.get_all(...)` | Filter by `user_id`, `agent_id`, `run_id`. |
|
||||
| Retrieve memory | `await memory.get(memory_id=...)` | Raises `ValueError` if ID is invalid. |
|
||||
| Update memory | `await memory.update(memory_id=..., data=...)` | Accepts partial updates. |
|
||||
| Update memory | `await memory.update(memory_id=..., text=...)` | Accepts partial updates. |
|
||||
| Delete memory | `await memory.delete(memory_id=...)` | Returns confirmation payload. |
|
||||
| Delete in bulk | `await memory.delete_all(...)` | Requires at least one scope filter. |
|
||||
| History | `await memory.history(memory_id=...)` | Fetches change log for auditing. |
|
||||
@@ -185,7 +185,7 @@ specific_memory = await memory.get(memory_id="memory-id-here")
|
||||
# Update a memory
|
||||
updated_memory = await memory.update(
|
||||
memory_id="memory-id-here",
|
||||
data="I'm travelling to Seattle"
|
||||
text="I'm travelling to Seattle"
|
||||
)
|
||||
|
||||
# Delete a memory
|
||||
|
||||
@@ -18,7 +18,129 @@ Reranker-enhanced search adds a second scoring pass after vector retrieval so Me
|
||||
</Warning>
|
||||
|
||||
<Note>
|
||||
All configuration snippets translate directly to the TypeScript SDK: swap dictionaries for objects while keeping the same keys (`provider`, `config`, `rerank` flags).
|
||||
The `Configure it` and `See it in action` snippets below use the Python SDK. The self-hosted **TypeScript SDK** supports the Cohere, Zero Entropy, Sentence Transformer, Hugging Face, and LLM rerankers; see [TypeScript SDK](#typescript-sdk).
|
||||
</Note>
|
||||
|
||||
---
|
||||
|
||||
## TypeScript SDK
|
||||
|
||||
The self-hosted TypeScript SDK (`mem0ai/oss`) ships five rerankers: **Cohere**, **Zero Entropy**, **Sentence Transformer**, **Hugging Face**, and the **LLM reranker**. Configure one under `reranker`, then opt in per search with `rerank: true`. Keys are camelCase (`apiKey`, not `api_key`).
|
||||
|
||||
Provider SDKs are peer dependencies. Install the one your reranker needs:
|
||||
|
||||
```bash
|
||||
pnpm add cohere-ai # cohere
|
||||
pnpm add zeroentropy # zero_entropy
|
||||
pnpm add @huggingface/transformers # sentence_transformer, huggingface
|
||||
# llm_reranker defaults to openai (already a core dependency); install another
|
||||
# provider's SDK only if you nest a different one under config.llm
|
||||
```
|
||||
|
||||
### Hosted rerankers (Cohere, Zero Entropy)
|
||||
|
||||
Both call a hosted API and read their key from config or the provider's environment variable (`COHERE_API_KEY`, `ZERO_ENTROPY_API_KEY`).
|
||||
|
||||
```typescript
|
||||
import { Memory } from "mem0ai/oss";
|
||||
|
||||
// Cohere reranker (defaults to the rerank-v3.5 model)
|
||||
const memory = new Memory({
|
||||
reranker: {
|
||||
provider: "cohere",
|
||||
config: { apiKey: process.env.COHERE_API_KEY },
|
||||
},
|
||||
});
|
||||
|
||||
const results = await memory.search("What are my food preferences?", {
|
||||
filters: { userId: "alice" },
|
||||
rerank: true,
|
||||
});
|
||||
```
|
||||
|
||||
```typescript
|
||||
// Zero Entropy reranker (defaults to the zerank-1 model)
|
||||
const memory = new Memory({
|
||||
reranker: {
|
||||
provider: "zero_entropy",
|
||||
config: { apiKey: process.env.ZERO_ENTROPY_API_KEY },
|
||||
},
|
||||
});
|
||||
```
|
||||
|
||||
### Local cross-encoders (Sentence Transformer, Hugging Face)
|
||||
|
||||
Both run a cross-encoder locally with [Transformers.js](https://huggingface.co/docs/transformers.js): no API key, no network at inference time. Because Transformers.js runs ONNX weights, the default models are the ONNX mirrors of the Python SDK's defaults (`sentence_transformer` → `Xenova/ms-marco-MiniLM-L-6-v2`, `huggingface` → `Xenova/bge-reranker-base`). Point `model` at any ONNX-exported cross-encoder on the Hub to override.
|
||||
|
||||
```typescript
|
||||
const memory = new Memory({
|
||||
reranker: {
|
||||
provider: "sentence_transformer", // or "huggingface"
|
||||
config: {
|
||||
// model: "Xenova/bge-reranker-base", // override the default
|
||||
device: "cpu", // Transformers.js device: "cpu" | "wasm" | "webgpu"
|
||||
maxLength: 512, // max tokens per query-document pair
|
||||
normalize: true, // sigmoid-normalize logits to [0, 1] (default)
|
||||
},
|
||||
},
|
||||
});
|
||||
|
||||
const results = await memory.search("What movies do I like?", {
|
||||
filters: { userId: "alice" },
|
||||
rerank: true,
|
||||
});
|
||||
```
|
||||
|
||||
<Note>
|
||||
`batchSize` and `showProgressBar` are accepted for config parity with the Python SDK but are no-ops in this runtime, because a memory search reranks a small candidate set in a single in-process forward pass. The model is downloaded once and cached in-process on first use.
|
||||
</Note>
|
||||
|
||||
### LLM reranker
|
||||
|
||||
To score with an LLM instead of a dedicated reranker, use the `llm_reranker` provider. It builds its own LLM from the reranker's config (defaulting to `openai` / `gpt-4o-mini`) rather than reusing the Memory's main `llm`:
|
||||
|
||||
```typescript
|
||||
const memory = new Memory({
|
||||
reranker: {
|
||||
provider: "llm_reranker",
|
||||
config: { apiKey: process.env.OPENAI_API_KEY },
|
||||
},
|
||||
});
|
||||
|
||||
const results = await memory.search("What movies do I like?", {
|
||||
filters: { userId: "alice" },
|
||||
rerank: true,
|
||||
});
|
||||
```
|
||||
|
||||
Nest a different provider under `config.llm` to override the default:
|
||||
|
||||
```typescript
|
||||
const memory = new Memory({
|
||||
reranker: {
|
||||
provider: "llm_reranker",
|
||||
config: {
|
||||
llm: {
|
||||
provider: "anthropic",
|
||||
config: { apiKey: process.env.ANTHROPIC_API_KEY },
|
||||
},
|
||||
},
|
||||
},
|
||||
});
|
||||
```
|
||||
|
||||
### Config reference
|
||||
|
||||
| Provider | Default model | Key config fields |
|
||||
| --- | --- | --- |
|
||||
| `cohere` | `rerank-v3.5` | `apiKey`, `model`, `topK` |
|
||||
| `zero_entropy` | `zerank-1` | `apiKey`, `model`, `topK` |
|
||||
| `sentence_transformer` | `Xenova/ms-marco-MiniLM-L-6-v2` | `model`, `device`, `maxLength`, `normalize`, `topK` |
|
||||
| `huggingface` | `Xenova/bge-reranker-base` | `model`, `device`, `maxLength`, `normalize`, `topK` |
|
||||
| `llm_reranker` | `openai` / `gpt-4o-mini` | `provider`, `model`, `apiKey`, `llm` (nested override), `topK` |
|
||||
|
||||
<Note>
|
||||
`rerank` is opt-in per search and a no-op when no `reranker` is configured. If the reranker call fails, Mem0 logs a warning and returns the original vector-ranked results.
|
||||
</Note>
|
||||
|
||||
---
|
||||
@@ -61,7 +183,7 @@ config = {
|
||||
"reranker": {
|
||||
"provider": "cohere",
|
||||
"config": {
|
||||
"model": "rerank-english-v3.0",
|
||||
"model": "rerank-v3.5",
|
||||
"api_key": "your-cohere-api-key"
|
||||
}
|
||||
}
|
||||
@@ -86,7 +208,7 @@ config = {
|
||||
"reranker": {
|
||||
"provider": "cohere",
|
||||
"config": {
|
||||
"model": "rerank-english-v3.0",
|
||||
"model": "rerank-v3.5",
|
||||
"api_key": "your-cohere-api-key",
|
||||
"top_k": 10,
|
||||
"return_documents": True
|
||||
@@ -164,7 +286,7 @@ config = {
|
||||
"reranker": {
|
||||
"provider": "cohere",
|
||||
"config": {
|
||||
"model": "rerank-english-v3.0",
|
||||
"model": "rerank-v3.5",
|
||||
"api_key": "your-cohere-api-key",
|
||||
"top_k": 15,
|
||||
"return_documents": True
|
||||
@@ -338,7 +460,7 @@ config = {
|
||||
"reranker": {
|
||||
"provider": "cohere",
|
||||
"config": {
|
||||
"model": "rerank-english-v3.0",
|
||||
"model": "rerank-v3.5",
|
||||
"api_key": "your-cohere-api-key"
|
||||
}
|
||||
}
|
||||
|
||||
+27
-16
@@ -1232,7 +1232,7 @@
|
||||
"tags": [
|
||||
"memories"
|
||||
],
|
||||
"description": "Delete memories by filter. At least one filter is required — previously omitting all filters silently deleted everything; now it returns a validation error.",
|
||||
"description": "Delete memories by filter. At least one filter is required. Previously, omitting all filters silently deleted everything; now it returns a validation error.",
|
||||
"operationId": "memories_delete_all",
|
||||
"parameters": [
|
||||
{
|
||||
@@ -1315,15 +1315,15 @@
|
||||
"x-code-samples": [
|
||||
{
|
||||
"lang": "Python",
|
||||
"source": "# To use the Python SDK, install the package:\n# pip install mem0ai\n\nfrom mem0 import MemoryClient\nclient = MemoryClient(api_key=\"your_api_key\")\n\n# Delete all memories for a specific user\nclient.delete_all(user_id=\"<user_id>\")\n\n# Delete all memories for every user in the project (wildcard)\nclient.delete_all(user_id=\"*\")\n\n# Full project wipe — all four filters must be explicitly set to \"*\"\nclient.delete_all(user_id=\"*\", agent_id=\"*\", app_id=\"*\", run_id=\"*\")\n\n# NOTE: Calling delete_all() with no filters raises a validation error.\n# At least one filter is required to prevent accidental data loss."
|
||||
"source": "# To use the Python SDK, install the package:\n# pip install mem0ai\n\nfrom mem0 import MemoryClient\nclient = MemoryClient(api_key=\"your_api_key\")\n\n# Delete all memories for a specific user\nclient.delete_all(user_id=\"<user_id>\")\n\n# Delete all memories for every user in the project (wildcard)\nclient.delete_all(user_id=\"*\")\n\n# Full project wipe: all four filters must be explicitly set to \"*\"\nclient.delete_all(user_id=\"*\", agent_id=\"*\", app_id=\"*\", run_id=\"*\")\n\n# NOTE: Calling delete_all() with no filters raises a validation error.\n# At least one filter is required to prevent accidental data loss."
|
||||
},
|
||||
{
|
||||
"lang": "JavaScript",
|
||||
"source": "// To use the JavaScript SDK, install the package:\n// npm i mem0ai\n\nimport MemoryClient from 'mem0ai';\nconst client = new MemoryClient({ apiKey: \"your-api-key\" });\n\n// Delete all memories for a specific user\nclient.deleteAll({ user_id: \"<user_id>\" })\n .then(result => console.log(result))\n .catch(error => console.error(error));\n\n// Delete all memories for every user in the project (wildcard)\nclient.deleteAll({ user_id: \"*\" })\n .then(result => console.log(result))\n .catch(error => console.error(error));\n\n// Full project wipe — all four filters must be explicitly set to \"*\"\nclient.deleteAll({ user_id: \"*\", agent_id: \"*\", app_id: \"*\", run_id: \"*\" })\n .then(result => console.log(result))\n .catch(error => console.error(error));"
|
||||
"source": "// To use the JavaScript SDK, install the package:\n// npm i mem0ai\n\nimport MemoryClient from 'mem0ai';\nconst client = new MemoryClient({ apiKey: \"your-api-key\" });\n\n// Delete all memories for a specific user\nclient.deleteAll({ user_id: \"<user_id>\" })\n .then(result => console.log(result))\n .catch(error => console.error(error));\n\n// Delete all memories for every user in the project (wildcard)\nclient.deleteAll({ user_id: \"*\" })\n .then(result => console.log(result))\n .catch(error => console.error(error));\n\n// Full project wipe: all four filters must be explicitly set to \"*\"\nclient.deleteAll({ user_id: \"*\", agent_id: \"*\", app_id: \"*\", run_id: \"*\" })\n .then(result => console.log(result))\n .catch(error => console.error(error));"
|
||||
},
|
||||
{
|
||||
"lang": "cURL",
|
||||
"source": "# Delete memories for a specific user\ncurl --request DELETE \\\n --url 'https://api.mem0.ai/v1/memories/?user_id=<user_id>' \\\n --header 'Authorization: Token <api-key>'\n\n# Delete memories for all users (wildcard)\ncurl --request DELETE \\\n --url 'https://api.mem0.ai/v1/memories/?user_id=*' \\\n --header 'Authorization: Token <api-key>'\n\n# Full project wipe — all four filters must be set to *\ncurl --request DELETE \\\n --url 'https://api.mem0.ai/v1/memories/?user_id=*&agent_id=*&app_id=*&run_id=*' \\\n --header 'Authorization: Token <api-key>'"
|
||||
"source": "# Delete memories for a specific user\ncurl --request DELETE \\\n --url 'https://api.mem0.ai/v1/memories/?user_id=<user_id>' \\\n --header 'Authorization: Token <api-key>'\n\n# Delete memories for all users (wildcard)\ncurl --request DELETE \\\n --url 'https://api.mem0.ai/v1/memories/?user_id=*' \\\n --header 'Authorization: Token <api-key>'\n\n# Full project wipe: all four filters must be set to *\ncurl --request DELETE \\\n --url 'https://api.mem0.ai/v1/memories/?user_id=*&agent_id=*&app_id=*&run_id=*' \\\n --header 'Authorization: Token <api-key>'"
|
||||
},
|
||||
{
|
||||
"lang": "Go",
|
||||
@@ -1740,7 +1740,7 @@
|
||||
"memories"
|
||||
],
|
||||
"summary": "Get all memories (V3, paginated)",
|
||||
"description": "List memories scoped by filters, paginated. Entity IDs **must** be passed inside the `filters` object — top-level `user_id` / `agent_id` / `run_id` are rejected with 400. `filters` supports the same operator set as V2 search (`AND`, `OR`, `NOT`, `in`, `gte`, `lte`, etc.). Response is a paginated envelope; pass `page` and `page_size` as query parameters to step through results.",
|
||||
"description": "List memories scoped by filters, paginated. Entity IDs **must** be passed inside the `filters` object. Top-level `user_id` / `agent_id` / `run_id` are rejected with 400. `filters` supports the same operator set as V2 search (`AND`, `OR`, `NOT`, `in`, `gte`, `lte`, etc.). Response is a paginated envelope; pass `page` and `page_size` as query parameters to step through results.",
|
||||
"operationId": "memories_list_v3",
|
||||
"parameters": [
|
||||
{
|
||||
@@ -1896,10 +1896,10 @@
|
||||
}
|
||||
},
|
||||
"400": {
|
||||
"description": "Validation error — e.g. empty `filters` or no positively-scoped entity ID."
|
||||
"description": "Validation error, e.g. empty `filters` or no positively-scoped entity ID."
|
||||
},
|
||||
"401": {
|
||||
"description": "Unauthorized — missing or invalid API key."
|
||||
"description": "Unauthorized: missing or invalid API key."
|
||||
}
|
||||
},
|
||||
"security": [
|
||||
@@ -1992,6 +1992,17 @@
|
||||
"type": "string",
|
||||
"description": "Project-level instructions that guide extraction for this call."
|
||||
},
|
||||
"custom_categories": {
|
||||
"type": "array",
|
||||
"description": "Category catalog for this call. Replaces the project-level list rather than merging with it. Omit to fall back to the project list, then the default catalog.",
|
||||
"items": {
|
||||
"type": "object",
|
||||
"additionalProperties": {
|
||||
"type": "string"
|
||||
},
|
||||
"description": "Maps a category name to the description the classifier matches against."
|
||||
}
|
||||
},
|
||||
"infer": {
|
||||
"type": "boolean",
|
||||
"default": true,
|
||||
@@ -2007,7 +2018,7 @@
|
||||
},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": "Got it — I'll update your location."
|
||||
"content": "Got it, I'll update your location."
|
||||
}
|
||||
],
|
||||
"user_id": "alice"
|
||||
@@ -2049,10 +2060,10 @@
|
||||
}
|
||||
},
|
||||
"400": {
|
||||
"description": "Validation error — e.g. missing `messages` or no entity ID supplied."
|
||||
"description": "Validation error, e.g. missing `messages` or no entity ID supplied."
|
||||
},
|
||||
"401": {
|
||||
"description": "Unauthorized — missing or invalid API key."
|
||||
"description": "Unauthorized: missing or invalid API key."
|
||||
}
|
||||
},
|
||||
"security": [
|
||||
@@ -2063,15 +2074,15 @@
|
||||
"x-codeSamples": [
|
||||
{
|
||||
"lang": "cURL",
|
||||
"source": "curl -X POST https://api.mem0.ai/v3/memories/add/ \\\n -H \"Authorization: Token <api-key>\" \\\n -H \"Content-Type: application/json\" \\\n -d '{\n \"messages\": [\n {\"role\": \"user\", \"content\": \"I just moved to San Francisco from New York.\"},\n {\"role\": \"assistant\", \"content\": \"Got it — I\\u0027ll update your location.\"}\n ],\n \"user_id\": \"alice\"\n }'"
|
||||
"source": "curl -X POST https://api.mem0.ai/v3/memories/add/ \\\n -H \"Authorization: Token <api-key>\" \\\n -H \"Content-Type: application/json\" \\\n -d '{\n \"messages\": [\n {\"role\": \"user\", \"content\": \"I just moved to San Francisco from New York.\"},\n {\"role\": \"assistant\", \"content\": \"Got it, I\\u0027ll update your location.\"}\n ],\n \"user_id\": \"alice\"\n }'"
|
||||
},
|
||||
{
|
||||
"lang": "Python",
|
||||
"source": "from mem0 import MemoryClient\n\nclient = MemoryClient(api_key=\"your-api-key\")\n\nresult = client.add(\n messages=[\n {\"role\": \"user\", \"content\": \"I just moved to San Francisco from New York.\"},\n {\"role\": \"assistant\", \"content\": \"Got it — I'll update your location.\"}\n ],\n user_id=\"alice\",\n)\nprint(result)"
|
||||
"source": "from mem0 import MemoryClient\n\nclient = MemoryClient(api_key=\"your-api-key\")\n\nresult = client.add(\n messages=[\n {\"role\": \"user\", \"content\": \"I just moved to San Francisco from New York.\"},\n {\"role\": \"assistant\", \"content\": \"Got it, I'll update your location.\"}\n ],\n user_id=\"alice\",\n)\nprint(result)"
|
||||
},
|
||||
{
|
||||
"lang": "JavaScript",
|
||||
"source": "import MemoryClient from \"mem0ai\";\n\nconst client = new MemoryClient({ apiKey: \"your-api-key\" });\n\nconst result = await client.add(\n [\n { role: \"user\", content: \"I just moved to San Francisco from New York.\" },\n { role: \"assistant\", content: \"Got it — I'll update your location.\" },\n ],\n { userId: \"alice\" }\n);\nconsole.log(result);"
|
||||
"source": "import MemoryClient from \"mem0ai\";\n\nconst client = new MemoryClient({ apiKey: \"your-api-key\" });\n\nconst result = await client.add(\n [\n { role: \"user\", content: \"I just moved to San Francisco from New York.\" },\n { role: \"assistant\", content: \"Got it, I'll update your location.\" },\n ],\n { userId: \"alice\" }\n);\nconsole.log(result);"
|
||||
}
|
||||
]
|
||||
}
|
||||
@@ -2082,7 +2093,7 @@
|
||||
"memories"
|
||||
],
|
||||
"summary": "Search memories (V3)",
|
||||
"description": "Relevance-ranked search across stored memories. V3 uses hybrid retrieval and can also apply temporal reasoning for time-aware queries. Entity IDs **must** be passed inside the `filters` object — top-level `user_id` / `agent_id` / `run_id` are rejected with 400. At least one entity ID is required.",
|
||||
"description": "Relevance-ranked search across stored memories. V3 uses hybrid retrieval and can also apply temporal reasoning for time-aware queries. Entity IDs **must** be passed inside the `filters` object. Top-level `user_id` / `agent_id` / `run_id` are rejected with 400. At least one entity ID is required.",
|
||||
"operationId": "memories_search_v3",
|
||||
"requestBody": {
|
||||
"required": true,
|
||||
@@ -2236,10 +2247,10 @@
|
||||
}
|
||||
},
|
||||
"400": {
|
||||
"description": "Validation error — e.g. empty `query`, missing `filters`, or no positively-scoped entity ID."
|
||||
"description": "Validation error, e.g. empty `query`, missing `filters`, or no positively-scoped entity ID."
|
||||
},
|
||||
"401": {
|
||||
"description": "Unauthorized — missing or invalid API key."
|
||||
"description": "Unauthorized: missing or invalid API key."
|
||||
}
|
||||
},
|
||||
"security": [
|
||||
|
||||
@@ -129,7 +129,7 @@ matches = await memory.search(
|
||||
```python
|
||||
await memory.update(
|
||||
memory_id=matches["results"][0]["id"],
|
||||
data="Morgan avoids shellfish and prefers boutique hotels in central Tokyo.",
|
||||
text="Morgan avoids shellfish and prefers boutique hotels in central Tokyo.",
|
||||
)
|
||||
```
|
||||
</Step>
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
---
|
||||
title: Custom Categories
|
||||
description: "Replace default memory tags with custom category labels that match your product terminology at the project level."
|
||||
description: "Replace default memory tags with custom category labels that match your product terminology, set once per project or per individual add call."
|
||||
---
|
||||
|
||||
# Custom Categories
|
||||
@@ -14,9 +14,7 @@ Mem0 automatically tags every memory, but the default labels (travel, sports, mu
|
||||
- You’re moving from the open-source version and want the same labels here.
|
||||
</Info>
|
||||
|
||||
<Warning>
|
||||
Per-request overrides (`custom_categories=...` on `client.add`) are not supported on the managed API yet. Set categories at the project level, then ingest memories as usual.
|
||||
</Warning>
|
||||
You can set the list once for the whole project, or pass a different list on an individual `add` call.
|
||||
|
||||
## Configure access
|
||||
|
||||
@@ -27,7 +25,20 @@ Mem0 automatically tags every memory, but the default labels (travel, sports, mu
|
||||
|
||||
- **Default list**: Each project starts with 15 broad categories like `travel`, `sports`, and `music`.
|
||||
- **Project override**: When you call `project.update(custom_categories=[...])`, that list replaces the defaults for future memories.
|
||||
- **Automatic tags**: As new memories come in, Mem0 picks the closest matches from your list and saves them in the `categories` field.
|
||||
- **Per-call override**: When you pass `custom_categories=[...]` to `client.add(...)`, that list is used for the memories extracted from that call.
|
||||
- **Automatic tags**: As new memories come in, Mem0 picks the closest matches from the active list and saves them in the `categories` field.
|
||||
|
||||
### Which list wins
|
||||
|
||||
Mem0 resolves the category catalog for each `add` call in this order, and stops at the first one it finds:
|
||||
|
||||
1. `custom_categories` passed on the `add` call
|
||||
2. `custom_categories` set on the project
|
||||
3. The built-in default catalog
|
||||
|
||||
A per-call list **fully replaces** the project list for that call. The two are not merged, so a memory added with a per-call list can only be tagged with categories from that list.
|
||||
|
||||
Categories are applied at ingestion time. Changing the project list, or passing a new per-call list, does not re-tag memories that already exist.
|
||||
|
||||
<Note>
|
||||
Default catalog: `personal_details`, `family`, `professional_details`, `sports`, `travel`, `food`, `music`, `health`, `technology`, `hobbies`, `fashion`, `entertainment`, `milestones`, `user_preferences`, `misc`.
|
||||
@@ -84,6 +95,82 @@ print(categories)
|
||||
```
|
||||
</CodeGroup>
|
||||
|
||||
`get` echoes back the shape you set. `update` also accepts a plain list of names, such as `["billing", "support"]`, in which case `get` returns that same list of names. Descriptions are optional here, and the classifier uses them to disambiguate when it has them.
|
||||
|
||||
<Warning>
|
||||
`add` is stricter than `update`. Every entry in a per-call `custom_categories` list must be an object mapping a name to a description. Passing bare names to `add` fails with `400 Expected a dictionary of items but got type "str"`.
|
||||
</Warning>
|
||||
|
||||
### 3. Override categories on a single add call
|
||||
|
||||
Pass `custom_categories` directly to `add` when one call needs a different catalog than the project default. The memories created by that call are tagged from the list you pass, and the project list is left untouched.
|
||||
|
||||
<CodeGroup>
|
||||
```python Python
|
||||
health_messages = [
|
||||
{"role": "user", "content": "My doctor bumped my metformin to 1000mg and I see her again on the 14th."},
|
||||
{"role": "assistant", "content": "Noted the new dosage and the follow-up appointment."},
|
||||
]
|
||||
|
||||
health_categories = [
|
||||
{"symptoms": "Reported physical or mental symptoms"},
|
||||
{"medications": "Prescriptions, dosages, and adherence"},
|
||||
{"appointments": "Scheduled visits and follow-ups"},
|
||||
]
|
||||
|
||||
client.add(
|
||||
health_messages,
|
||||
user_id="alice",
|
||||
custom_categories=health_categories,
|
||||
)
|
||||
```
|
||||
|
||||
```javascript JavaScript
|
||||
const healthMessages = [
|
||||
{ role: "user", content: "My doctor bumped my metformin to 1000mg and I see her again on the 14th." },
|
||||
{ role: "assistant", content: "Noted the new dosage and the follow-up appointment." },
|
||||
];
|
||||
|
||||
const healthCategories = [
|
||||
{ symptoms: "Reported physical or mental symptoms" },
|
||||
{ medications: "Prescriptions, dosages, and adherence" },
|
||||
{ appointments: "Scheduled visits and follow-ups" },
|
||||
];
|
||||
|
||||
await client.add(healthMessages, {
|
||||
userId: "alice",
|
||||
customCategories: healthCategories,
|
||||
});
|
||||
```
|
||||
|
||||
```text Resulting categories
|
||||
["medications", "appointments"]
|
||||
```
|
||||
</CodeGroup>
|
||||
|
||||
The memory is tagged from `health_categories` alone. The project catalog is not consulted for this call, and it is not modified.
|
||||
|
||||
#### Per-user categories inside one project
|
||||
|
||||
The main reason to reach for a per-call list is to give different users, tenants, or entities their own vocabulary without splitting them across projects. Keep one project, and pass the list that fits the entity you are writing for.
|
||||
|
||||
<CodeGroup>
|
||||
```python Python
|
||||
patient_categories = [
|
||||
{"symptoms": "Reported physical or mental symptoms"},
|
||||
{"medications": "Prescriptions, dosages, and adherence"},
|
||||
]
|
||||
|
||||
clinician_categories = [
|
||||
{"caseload": "Patients under this clinician's care"},
|
||||
{"availability": "Shift patterns and on-call windows"},
|
||||
]
|
||||
|
||||
client.add("My metformin is now 1000mg.", user_id="alice", custom_categories=patient_categories)
|
||||
client.add("I'm on call Tuesdays and Thursdays.", user_id="dr-reyes", custom_categories=clinician_categories)
|
||||
```
|
||||
</CodeGroup>
|
||||
|
||||
## See it in action
|
||||
|
||||
### Add a memory (uses the project catalog automatically)
|
||||
@@ -104,46 +191,57 @@ client.add(messages, user_id="alice")
|
||||
|
||||
### Retrieve memories and inspect categories
|
||||
|
||||
`get_all` returns a paginated object. The memories are under `results`, and each one carries its own `categories` list.
|
||||
|
||||
<CodeGroup>
|
||||
```python Code
|
||||
memories = client.get_all(filters={"user_id": "alice"})
|
||||
response = client.get_all(filters={"user_id": "alice"})
|
||||
|
||||
for memory in response["results"]:
|
||||
print(memory["memory"], memory["categories"])
|
||||
```
|
||||
|
||||
```json Output
|
||||
["lifestyle_management_concerns", "seeking_structure"]
|
||||
```text Output
|
||||
User introduced herself as Alice and expressed a desire for help organizing her daily schedule. ['lifestyle_management_concerns', 'seeking_structure', 'personal_information']
|
||||
User feels overwhelmed trying to balance work responsibilities, regular exercise, and a social life, indicating difficulty managing time across these areas. ['lifestyle_management_concerns']
|
||||
User's goals include becoming more productive at work, maintaining a consistent workout routine, and preserving enough energy for friends and hobbies. ['lifestyle_management_concerns', 'seeking_structure']
|
||||
```
|
||||
</CodeGroup>
|
||||
|
||||
Extraction is model driven, so the exact wording and the number of memories vary between runs. The categories are drawn from the active list.
|
||||
|
||||
<Info>
|
||||
**Sample memory payload**
|
||||
```json
|
||||
{
|
||||
"id": "33d2***",
|
||||
"memory": "Trying to balance work and workouts",
|
||||
"id": "638008c4-***",
|
||||
"memory": "User is seeking to balance work responsibilities with regular workout sessions and requests a personalized schedule to manage both.",
|
||||
"user_id": "alice",
|
||||
"metadata": null,
|
||||
"categories": ["wellness"], // ← matches the custom category we set
|
||||
"created_at": "2025-11-01T02:13:32.828364-07:00",
|
||||
"updated_at": "2025-11-01T02:13:32.830896-07:00",
|
||||
"categories": ["lifestyle_management_concerns", "seeking_structure"],
|
||||
"created_at": "2026-07-10T06:13:12-07:00",
|
||||
"updated_at": "2026-07-10T06:13:20-07:00",
|
||||
"expiration_date": null,
|
||||
"structured_attributes": {
|
||||
"day": 1,
|
||||
"hour": 9,
|
||||
"year": 2025,
|
||||
"month": 11,
|
||||
"year": 2026,
|
||||
"month": 7,
|
||||
"day": 10,
|
||||
"hour": 13,
|
||||
"minute": 13,
|
||||
"quarter": 4,
|
||||
"is_weekend": true,
|
||||
"day_of_week": "saturday",
|
||||
"day_of_year": 305,
|
||||
"week_of_year": 44
|
||||
"day_of_week": "friday",
|
||||
"week_of_year": 28,
|
||||
"day_of_year": 191,
|
||||
"quarter": 3,
|
||||
"is_weekend": false
|
||||
}
|
||||
}
|
||||
```
|
||||
</Info>
|
||||
|
||||
Categorization runs asynchronously, a moment after the memory itself is written. A memory fetched immediately after `add` may not show up in `get_all` yet, or can come back with `categories: null` and pick up its tags a moment later. Poll until `categories` is populated rather than reading once.
|
||||
|
||||
<Note>
|
||||
Need ad-hoc labels for a single call? Store them in `metadata` until per-request overrides become available.
|
||||
Need ad-hoc labels for a single call? Pass `custom_categories` on that `add` call. Use `metadata` instead when the label is a fixed value you already know, rather than something the classifier should infer.
|
||||
</Note>
|
||||
|
||||
## Default categories (fallback)
|
||||
@@ -208,11 +306,13 @@ client.project.get(["custom_categories"])
|
||||
|
||||
```json Output
|
||||
{
|
||||
"custom_categories": None
|
||||
"custom_categories": null
|
||||
}
|
||||
```
|
||||
</CodeGroup>
|
||||
|
||||
A project that has never set a list returns `null`. One you have reset with `project.update(custom_categories=[])` returns `[]`. Both mean the default catalog is active.
|
||||
|
||||
## Verify the feature is working
|
||||
|
||||
- `client.project.get(["custom_categories"])` returns the category list you set.
|
||||
@@ -223,7 +323,8 @@ client.project.get(["custom_categories"])
|
||||
|
||||
- Keep category descriptions concise but specific; the classifier uses them to disambiguate.
|
||||
- Review memories with empty `categories` to see where you might extend or rename your list.
|
||||
- Stick with project-level overrides until per-request support is released; mixing approaches causes confusion.
|
||||
- Set the catalog your app uses most often at the project level, and reserve per-call lists for the calls that genuinely need a different vocabulary.
|
||||
- If a per-call list should also keep the project categories, include them in the list you pass. Passing a list replaces, it does not extend.
|
||||
|
||||
<CardGroup cols={2}>
|
||||
<Card title="Advanced Memory Operations" icon="wand-magic-sparkles" href="/platform/advanced-memory-operations">
|
||||
|
||||
@@ -0,0 +1,175 @@
|
||||
---
|
||||
title: Memory Expiration
|
||||
description: "Give a memory a shelf life: set an expiration date and it stops surfacing in search once that date passes, without deleting the record."
|
||||
---
|
||||
|
||||
# Memory Expiration
|
||||
|
||||
Some facts are only true for a while. A trial plan ends, a seasonal preference goes stale, a support ticket ages past its retention window. Set an `expiration_date` on a memory and Mem0 stops surfacing it once that date passes, so you don't need a cleanup job hunting for rows to delete.
|
||||
|
||||
**Expiration hides a memory, it does not delete it.** The record stays in storage untouched. `search()` and `get_all()` skip it, fetching it by ID still returns it, and clearing the date brings it straight back.
|
||||
|
||||
## How it works
|
||||
|
||||
- **Format**: a plain `YYYY-MM-DD` date. No time component, no timezone offset.
|
||||
- **Evaluated in UTC**, never against the caller's local timezone.
|
||||
- **Inclusive of the date itself**: a memory set to expire on `2030-01-31` stays visible all through `2030-01-31` UTC and disappears on `2030-02-01`.
|
||||
- **Only list-shaped reads filter**: `search()` and `get_all()` (`getAll()` in TypeScript) hide expired memories. <Link href="/api-reference/memory/get-memory">`get(memory_id)`</Link> always returns the memory, so there is no `show_expired` parameter on that path.
|
||||
- **No expiration date means never expires.** That is the default for every memory.
|
||||
- **Malformed dates fail open**: a stored value Mem0 can't parse is treated as *not* expired. A bad date never makes a memory silently vanish.
|
||||
|
||||
## Set an expiration date
|
||||
|
||||
Set it when you add the memory:
|
||||
|
||||
<CodeGroup>
|
||||
```python Python
|
||||
# Mem0 Platform
|
||||
from mem0 import MemoryClient
|
||||
|
||||
client = MemoryClient(api_key="your-api-key")
|
||||
messages = [{"role": "user", "content": "My Pro trial ends soon."}]
|
||||
|
||||
client.add(messages, user_id="alice", expiration_date="2030-01-31")
|
||||
|
||||
# Mem0 OSS
|
||||
from mem0 import Memory
|
||||
|
||||
memory = Memory()
|
||||
memory.add(messages, user_id="alice", expiration_date="2030-01-31")
|
||||
```
|
||||
|
||||
```javascript JavaScript
|
||||
// Mem0 Platform
|
||||
import { MemoryClient } from "mem0ai";
|
||||
|
||||
const client = new MemoryClient({ apiKey: "your-api-key" });
|
||||
const messages = [{ role: "user", content: "My Pro trial ends soon." }];
|
||||
|
||||
await client.add(messages, { userId: "alice", expirationDate: "2030-01-31" });
|
||||
|
||||
// Mem0 OSS
|
||||
import { Memory } from "mem0ai/oss";
|
||||
|
||||
const memory = new Memory();
|
||||
await memory.add(messages, { userId: "alice", expirationDate: "2030-01-31" });
|
||||
```
|
||||
</CodeGroup>
|
||||
|
||||
Or attach it to a memory that already exists, using `update()`:
|
||||
|
||||
<CodeGroup>
|
||||
```python Python
|
||||
client.update("mem_123", expiration_date="2030-01-31") # Platform
|
||||
memory.update("mem_123", expiration_date="2030-01-31") # OSS
|
||||
```
|
||||
|
||||
```javascript JavaScript
|
||||
await client.update("mem_123", { expirationDate: "2030-01-31" }); // Platform
|
||||
await memory.update("mem_123", { expirationDate: "2030-01-31" }); // OSS
|
||||
```
|
||||
</CodeGroup>
|
||||
|
||||
Full field lists live in the <Link href="/api-reference/memory/add-memories">Add Memories</Link> and <Link href="/api-reference/memory/update-memory">Update Memory</Link> references.
|
||||
|
||||
## Read expired memories back
|
||||
|
||||
Pass `show_expired` (`showExpired` in TypeScript) to include them. It defaults to `false` on every client.
|
||||
|
||||
<CodeGroup>
|
||||
```python Python
|
||||
client.get_all(filters={"user_id": "alice"}, show_expired=True)
|
||||
client.search("What plan is Alice on?", filters={"user_id": "alice"}, show_expired=True)
|
||||
```
|
||||
|
||||
```javascript JavaScript
|
||||
await client.getAll({ filters: { user_id: "alice" }, showExpired: true });
|
||||
await client.search("What plan is Alice on?", {
|
||||
filters: { user_id: "alice" },
|
||||
showExpired: true,
|
||||
});
|
||||
```
|
||||
</CodeGroup>
|
||||
|
||||
The same parameter and spelling work on the OSS `Memory` class. See <Link href="/api-reference/memory/search-memories">Search Memories</Link> and <Link href="/api-reference/memory/get-memories">Get Memories</Link>.
|
||||
|
||||
<Note>
|
||||
Expired memories are dropped *before* your `top_k` is applied, so Mem0 widens the internal candidate pool first and short result sets are rare. They are not impossible: if nearly every memory in a scope has expired, a call can still return fewer than `top_k` results. Pass `show_expired: true` to get the full set back.
|
||||
</Note>
|
||||
|
||||
## Clear an expiration date
|
||||
|
||||
Pass an explicit `None` (Python) or `null` (TypeScript) to make the memory permanent again. The SDKs deliberately preserve that null instead of treating it as "argument not supplied".
|
||||
|
||||
<CodeGroup>
|
||||
```python Python
|
||||
client.update("mem_123", expiration_date=None) # Platform
|
||||
memory.update("mem_123", expiration_date=None) # OSS
|
||||
```
|
||||
|
||||
```javascript JavaScript
|
||||
await client.update("mem_123", { expirationDate: null }); // Platform
|
||||
await memory.update("mem_123", { expirationDate: null }); // OSS
|
||||
```
|
||||
</CodeGroup>
|
||||
|
||||
<Note>
|
||||
`update()` needs at least one of `text`, `metadata`, or `expiration_date`, and raises if you pass none of them. Clearing the date satisfies that on its own: the memory's content and metadata are left alone.
|
||||
</Note>
|
||||
|
||||
## What each client accepts
|
||||
|
||||
| Client | Accepted input | Notes |
|
||||
| --- | --- | --- |
|
||||
| Python (Platform and OSS) | `str` in `YYYY-MM-DD` form, or a `date` / `datetime` object | Normalized to `YYYY-MM-DD` before storage. |
|
||||
| TypeScript (Platform and OSS) | `string` in `YYYY-MM-DD` form only | Stricter than `new Date()`: rejects `12/31/2099`, `2099-12-31T23:00:00`, and non-days like `2099-02-30` or `2100-02-29`. |
|
||||
| Self-hosted REST server | `string` in `YYYY-MM-DD` form | Same normalization as OSS Python underneath. |
|
||||
| CLI (`mem0 add --expires`) | `string` in `YYYY-MM-DD` form | Must be strictly in the future, checked against the local system date. The SDKs have no such restriction. Platform only: the CLI has no OSS backend. |
|
||||
|
||||
## Reading the field back
|
||||
|
||||
Most clients return expiration as a top-level field on the memory: `expiration_date` in Python (Platform and OSS) and in the REST API, `expirationDate` in the Platform TypeScript SDK.
|
||||
|
||||
<Info>
|
||||
The OSS TypeScript SDK is the one exception. There, expiration round-trips under **`result.metadata.expiration_date`**, not `result.expirationDate`, on both `get()` and `getAll()`.
|
||||
</Info>
|
||||
|
||||
## Expiration, decay, and delete
|
||||
|
||||
These three get conflated. They solve different problems:
|
||||
|
||||
| | Memory Expiration | Memory Decay | Delete |
|
||||
| --- | --- | --- | --- |
|
||||
| What it does | Hides a memory once a date you set passes | Re-ranks results by how recently a memory was used | Removes a memory permanently |
|
||||
| Data still stored? | Yes | Yes | No |
|
||||
| Filters results? | Yes, after the date | Never, it only reorders scores | Yes, permanently |
|
||||
| Reversible? | Yes, clear or push back the date | Yes, toggle `decay` off | No |
|
||||
| Set where | Per memory, by you | Per project, opt-in | Per call |
|
||||
| Available in | Platform and OSS | Platform only | Platform and OSS |
|
||||
|
||||
Reach for <Link href="/platform/features/memory-decay">Memory Decay</Link> when old memories should rank lower but stay searchable, expiration when a memory should stop appearing after a specific known date, and <Link href="/core-concepts/memory-operations/delete">Delete</Link> when it should be gone for good.
|
||||
|
||||
## Common patterns
|
||||
|
||||
**Trial and subscription facts.** "Alice is on the Pro trial" is true until the trial ends. Set `expiration_date` to that end date when you write the fact. If she upgrades, clear the date and the memory becomes permanent. If she doesn't, it stops surfacing the next day on its own.
|
||||
|
||||
**Seasonal preferences.** "Alex wants gift ideas for the holidays" matters in December and is noise in July. A short-lived expiration date keeps it from competing with evergreen preferences in every search.
|
||||
|
||||
**Retention windows.** Data-retention policies usually want a soft window before a hard delete: keep a ticket's memories searchable for 90 days, stop surfacing them, purge them later on a schedule. Expiration is the soft step, and a scheduled <Link href="/core-concepts/memory-operations/delete">delete</Link> is the permanent one.
|
||||
|
||||
<CardGroup cols={2}>
|
||||
<Card
|
||||
title="Memory Decay"
|
||||
description="Rank stale memories lower instead of hiding them outright."
|
||||
icon="chart-line"
|
||||
href="/platform/features/memory-decay"
|
||||
/>
|
||||
<Card
|
||||
title="Delete Memories"
|
||||
description="Remove memories permanently instead of hiding them."
|
||||
icon="trash"
|
||||
href="/core-concepts/memory-operations/delete"
|
||||
/>
|
||||
</CardGroup>
|
||||
|
||||
<Snippet file="get-help.mdx" />
|
||||
@@ -102,9 +102,29 @@ await client.updateProject({ customCategories: newCategories });
|
||||
categories = client.project.get(fields=["custom_categories"])
|
||||
```
|
||||
|
||||
### Key Constraint
|
||||
**Override categories for a single add call:**
|
||||
```python
|
||||
client.add(messages, user_id="alice", custom_categories=per_call_categories)
|
||||
```
|
||||
|
||||
Per-request overrides (`custom_categories=...` on `client.add`) are **not supported** on the managed API. Only project-level configuration works. Workaround: store ad-hoc labels in `metadata` field.
|
||||
```javascript
|
||||
await client.add(messages, { userId: "alice", customCategories: perCallCategories });
|
||||
```
|
||||
|
||||
### Resolution Order
|
||||
|
||||
1. `custom_categories` passed on the `add` call
|
||||
2. `custom_categories` set on the project
|
||||
3. Built-in default catalog
|
||||
|
||||
### Key Constraints
|
||||
|
||||
- A per-call list **fully replaces** the project list for that call. The lists are not merged.
|
||||
- Categories are applied at ingestion time. Changing the list later does not re-tag existing memories.
|
||||
|
||||
### Main Use Case
|
||||
|
||||
Per-call lists give different users or entities their own vocabulary inside a single project, without splitting them across projects.
|
||||
|
||||
---
|
||||
|
||||
|
||||
+83
-4
@@ -1,6 +1,6 @@
|
||||
{
|
||||
"name": "mem0ai",
|
||||
"version": "3.0.13",
|
||||
"version": "3.1.0",
|
||||
"description": "The Memory Layer For Your AI Apps",
|
||||
"main": "./dist/index.js",
|
||||
"module": "./dist/index.mjs",
|
||||
@@ -98,7 +98,8 @@
|
||||
"ts-node": "^10.9.2",
|
||||
"tsup": "^8.3.0",
|
||||
"typescript": "5.5.4",
|
||||
"iovalkey": "^0.3.3"
|
||||
"iovalkey": "^0.3.3",
|
||||
"@mochow/mochow-sdk-node": "^2.1.5"
|
||||
},
|
||||
"dependencies": {
|
||||
"axios": "^1.16.0",
|
||||
@@ -108,12 +109,16 @@
|
||||
},
|
||||
"peerDependencies": {
|
||||
"@anthropic-ai/sdk": "^0.40.1",
|
||||
"@aws-sdk/client-neptune-graph": ">=3.0.0 <3.968.0",
|
||||
"@aws-sdk/client-s3vectors": "3.967.0",
|
||||
"@mochow/mochow-sdk-node": "^2.1.5",
|
||||
"@azure/identity": "^4.0.0",
|
||||
"@azure/search-documents": "^12.0.0",
|
||||
"@cloudflare/workers-types": "^4.20250504.0",
|
||||
"@databricks/sql": "^1.16.0",
|
||||
"@google-cloud/aiplatform": "^6.8.0",
|
||||
"@google/genai": "^1.40.0",
|
||||
"@huggingface/transformers": "^3.0.0 || ^4.0.0",
|
||||
"@langchain/core": "^1.1.47",
|
||||
"@mistralai/mistralai": "^1.5.2",
|
||||
"@opensearch-project/opensearch": "^3.5.1",
|
||||
@@ -126,10 +131,13 @@
|
||||
"@upstash/vector": "^1.2.3",
|
||||
"better-sqlite3": "^12.6.2",
|
||||
"cassandra-driver": "4.8.0",
|
||||
"chromadb": "^3.5.0",
|
||||
"cloudflare": "^4.2.0",
|
||||
"cohere-ai": "^7.17.0 || ^8.0.0",
|
||||
"fastembed": "^2.1.0",
|
||||
"groq-sdk": "0.3.0",
|
||||
"mongodb": "^7.0.0",
|
||||
"weaviate-client": "^3.0.0",
|
||||
"ollama": "^0.5.14",
|
||||
"pg": "8.11.3",
|
||||
"redis": "^4.6.13",
|
||||
@@ -137,11 +145,77 @@
|
||||
"iovalkey": "^0.3.3",
|
||||
"compromise": "^14.0.0",
|
||||
"natural": "^8.0.1",
|
||||
"mysql2": "^3.0.0"
|
||||
"zeroentropy": "^0.1.0-alpha.10",
|
||||
"mysql2": "^3.0.0",
|
||||
"@zilliz/milvus2-sdk-node": "^2.4.0 || ^3.0.0",
|
||||
"@aws-sdk/client-bedrock-runtime": ">=3.0.0 <3.968.0"
|
||||
},
|
||||
"peerDependenciesMeta": {
|
||||
"mysql2": {
|
||||
"optional": true
|
||||
},
|
||||
"@zilliz/milvus2-sdk-node": {
|
||||
"optional": true
|
||||
},
|
||||
"@mochow/mochow-sdk-node": {
|
||||
"optional": true
|
||||
},
|
||||
"@aws-sdk/client-bedrock-runtime": {
|
||||
"optional": true
|
||||
},
|
||||
"@databricks/sql": {
|
||||
"optional": true
|
||||
},
|
||||
"@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": {
|
||||
@@ -169,13 +243,18 @@
|
||||
"path-to-regexp@>=8.0.0 <8.4.0": "^8.4.0",
|
||||
"postcss@<8.5.10": ">=8.5.10",
|
||||
"uuid@<11.1.1": ">=11.1.1",
|
||||
"weaviate-client>uuid": "^11.1.1",
|
||||
"ws@>=8.0.0 <8.20.1": ">=8.20.1",
|
||||
"rollup@>=4.0.0 <4.59.0": "^4.59.0",
|
||||
"tar-fs@>=2.0.0 <2.1.4": "^2.1.4",
|
||||
"glob@>=10.2.0 <10.5.0": "^10.5.0",
|
||||
"fast-xml-parser@>=5.0.0 <5.7.0": "^5.9.3",
|
||||
"tar@<=7.5.15": "^7.5.19",
|
||||
"@modelcontextprotocol/sdk": "^1.25.4",
|
||||
"esbuild": ">=0.28.1",
|
||||
"undici@<6.27.0": ">=6.27.0 <8.0.0"
|
||||
"undici@<6.27.0": ">=6.27.0 <8.0.0",
|
||||
"@aws-sdk/client-bedrock-runtime": "3.967.0",
|
||||
"@aws-sdk/client-neptune-graph": "3.966.0"
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
Generated
+2573
-122
File diff suppressed because it is too large
Load Diff
@@ -110,14 +110,42 @@ You only need to provide API keys - all other settings are optional.
|
||||
### Methods
|
||||
|
||||
- `add(messages: string | Message[], userId?: string, ...): Promise<SearchResult>`
|
||||
- Options include `metadata`, `infer`, and `expirationDate` (a `YYYY-MM-DD` date after which the memory is treated as expired).
|
||||
- `search(query: string, userId?: string, ...): Promise<SearchResult>`
|
||||
- Expired memories are omitted unless you pass `showExpired: true`.
|
||||
- `get(memoryId: string): Promise<MemoryItem | null>`
|
||||
- `update(memoryId: string, data: string): Promise<{ message: string }>`
|
||||
- Fetching by ID returns the memory even if it has expired.
|
||||
- `getAll(options): Promise<SearchResult>`
|
||||
- Expired memories are omitted unless you pass `showExpired: true`.
|
||||
- `update(memoryId: string, config: string | UpdateMemoryOptions): Promise<{ message: string }>`
|
||||
- `UpdateMemoryOptions` is `{ text?, data?, metadata?, expirationDate? }`. At least one must be
|
||||
provided; omitted fields are left untouched, and `expirationDate: null` clears an existing
|
||||
expiry. `data` is a deprecated alias for `text`.
|
||||
- Passing a bare string is shorthand for `{ text }`, so `update(memoryId, "new text")` still works.
|
||||
- `delete(memoryId: string): Promise<{ message: string }>`
|
||||
- `deleteAll(userId?: string, ...): Promise<{ message: string }>`
|
||||
- `history(memoryId: string): Promise<any[]>`
|
||||
- `reset(): Promise<void>`
|
||||
|
||||
```typescript
|
||||
// Replace the content
|
||||
await memory.update(memoryId, { text: "Alex now prefers decaf coffee" });
|
||||
|
||||
// Update metadata only, leaving the stored text untouched
|
||||
await memory.update(memoryId, { metadata: { category: "preferences" } });
|
||||
|
||||
// Expire the memory on a given day, or clear an existing expiry
|
||||
await memory.update(memoryId, { expirationDate: "2030-01-31" });
|
||||
await memory.update(memoryId, { expirationDate: null });
|
||||
|
||||
// Include expired memories in reads
|
||||
await memory.getAll({ filters: { user_id: "alice" }, showExpired: true });
|
||||
await memory.search("coffee", {
|
||||
filters: { user_id: "alice" },
|
||||
showExpired: true,
|
||||
});
|
||||
```
|
||||
|
||||
### Try the Example
|
||||
|
||||
We provide a comprehensive example in `examples/basic.ts` that demonstrates all the features including:
|
||||
|
||||
@@ -73,10 +73,9 @@ async function runTests(memory: Memory) {
|
||||
}
|
||||
|
||||
// Updating this memory
|
||||
const result4 = await memory.update(
|
||||
result1.results[0].id,
|
||||
"I love India, it is my favorite country.",
|
||||
);
|
||||
const result4 = await memory.update(result1.results[0].id, {
|
||||
text: "I love India, it is my favorite country.",
|
||||
});
|
||||
console.log("Updated memory:", result4);
|
||||
|
||||
// Get all memories
|
||||
|
||||
@@ -55,10 +55,9 @@ export async function runTests(memory: Memory) {
|
||||
}
|
||||
|
||||
// Updating this memory
|
||||
const result4 = await memory.update(
|
||||
result1.results[0].id,
|
||||
"I love India, it is my favorite country.",
|
||||
);
|
||||
const result4 = await memory.update(result1.results[0].id, {
|
||||
text: "I love India, it is my favorite country.",
|
||||
});
|
||||
console.log("Updated memory:", result4);
|
||||
|
||||
// Get all memories
|
||||
|
||||
@@ -40,6 +40,10 @@ export class ConfigManager {
|
||||
| undefined);
|
||||
|
||||
return {
|
||||
// Spread first so provider-specific keys (e.g. the Vertex AI
|
||||
// project/location/credentials) survive the merge, while the
|
||||
// normalized values below still win.
|
||||
...userConf,
|
||||
apiKey:
|
||||
userConf?.apiKey !== undefined
|
||||
? userConf.apiKey
|
||||
@@ -56,9 +60,14 @@ export class ConfigManager {
|
||||
})(),
|
||||
},
|
||||
vectorStore: {
|
||||
provider:
|
||||
// Every factory already matches the provider case-insensitively, so a capitalized
|
||||
// name constructs the right store -- but the `provider === "memory"` comparisons that
|
||||
// pick per-provider entity-store settings do not. Normalize once, here, so those
|
||||
// comparisons cannot silently miss.
|
||||
provider: (
|
||||
userConfig.vectorStore?.provider ||
|
||||
DEFAULT_MEMORY_CONFIG.vectorStore.provider,
|
||||
DEFAULT_MEMORY_CONFIG.vectorStore.provider
|
||||
).toLowerCase(),
|
||||
config: (() => {
|
||||
const defaultConf = DEFAULT_MEMORY_CONFIG.vectorStore.config;
|
||||
const userConf = userConfig.vectorStore?.config;
|
||||
@@ -133,6 +142,11 @@ export class ConfigManager {
|
||||
userConf?.maxTokens ?? (llmRaw?.max_tokens as number | undefined);
|
||||
|
||||
return {
|
||||
// Spread user-provided config first so any additional fields
|
||||
// (e.g. future aws_bedrock options) pass through without a
|
||||
// manager.ts edit, matching the vectorStore.config pattern above
|
||||
// and making the schema's .passthrough() on llm.config meaningful.
|
||||
...userConf,
|
||||
baseURL: llmBaseURL,
|
||||
url: userConf?.url,
|
||||
apiKey:
|
||||
@@ -147,6 +161,20 @@ export class ConfigManager {
|
||||
temperature,
|
||||
topP,
|
||||
maxTokens,
|
||||
// Pass through AWS Bedrock fields so the aws_bedrock provider works
|
||||
// through the standard Memory config path (snake_case tolerated).
|
||||
awsRegion:
|
||||
userConf?.awsRegion ?? (llmRaw?.aws_region as string | undefined),
|
||||
awsAccessKeyId:
|
||||
userConf?.awsAccessKeyId ??
|
||||
(llmRaw?.aws_access_key_id as string | undefined),
|
||||
awsSecretAccessKey:
|
||||
userConf?.awsSecretAccessKey ??
|
||||
(llmRaw?.aws_secret_access_key as string | undefined),
|
||||
awsSessionToken:
|
||||
userConf?.awsSessionToken ??
|
||||
(llmRaw?.aws_session_token as string | undefined),
|
||||
client: userConf?.client,
|
||||
};
|
||||
})(),
|
||||
},
|
||||
@@ -177,6 +205,7 @@ export class ConfigManager {
|
||||
})(),
|
||||
disableHistory:
|
||||
userConfig.disableHistory || DEFAULT_MEMORY_CONFIG.disableHistory,
|
||||
reranker: userConfig.reranker,
|
||||
};
|
||||
|
||||
// Validate the merged config
|
||||
|
||||
@@ -1,4 +1,10 @@
|
||||
export interface Embedder {
|
||||
embed(text: string): Promise<number[]>;
|
||||
embedBatch(texts: string[]): Promise<number[][]>;
|
||||
embed(
|
||||
text: string,
|
||||
memoryAction?: "add" | "update" | "search",
|
||||
): Promise<number[]>;
|
||||
embedBatch(
|
||||
texts: string[],
|
||||
memoryAction?: "add" | "update" | "search",
|
||||
): Promise<number[][]>;
|
||||
}
|
||||
|
||||
@@ -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, " ");
|
||||
}
|
||||
|
||||
@@ -0,0 +1,78 @@
|
||||
import OpenAI from "openai";
|
||||
import { Embedder } from "./base";
|
||||
import { EmbeddingConfig } from "../types";
|
||||
|
||||
/**
|
||||
* HuggingFace embedding provider (hosted inference mode).
|
||||
*
|
||||
* Mirrors the `huggingface_base_url` branch of the Python provider
|
||||
* (`mem0/embeddings/huggingface.py`): a HuggingFace Text Embeddings Inference
|
||||
* (TEI) server, or any HuggingFace OpenAI-compatible inference endpoint,
|
||||
* exposes a `/v1/embeddings` route, so this embedder reuses the existing
|
||||
* `openai` client pointed at that base URL. No new dependency is required.
|
||||
*
|
||||
* A base URL is required. The Python provider's alternative local
|
||||
* `sentence-transformers` path has no lightweight TypeScript equivalent, so
|
||||
* hosted inference is the supported TS mode.
|
||||
*/
|
||||
export class HuggingFaceEmbedder implements Embedder {
|
||||
private openai: OpenAI;
|
||||
private model: string;
|
||||
|
||||
constructor(config: EmbeddingConfig) {
|
||||
const baseURL =
|
||||
config.huggingfaceBaseUrl ||
|
||||
config.baseURL ||
|
||||
config.url ||
|
||||
process.env.HUGGINGFACE_BASE_URL;
|
||||
|
||||
if (!baseURL) {
|
||||
throw new Error(
|
||||
"HuggingFace embedder requires an inference endpoint. Set " +
|
||||
"`huggingfaceBaseUrl` (or `baseURL`) in the embedder config, or the " +
|
||||
"HUGGINGFACE_BASE_URL environment variable (e.g. a TEI server at " +
|
||||
"http://localhost:8080/v1).",
|
||||
);
|
||||
}
|
||||
|
||||
this.openai = new OpenAI({
|
||||
apiKey: config.apiKey || process.env.HUGGINGFACE_API_KEY || "hf",
|
||||
baseURL,
|
||||
});
|
||||
// TEI ignores the model field; default mirrors the Python provider.
|
||||
this.model = config.model || "tei";
|
||||
}
|
||||
|
||||
async embed(text: string): Promise<number[]> {
|
||||
const response = await this.openai.embeddings.create({
|
||||
model: this.model,
|
||||
input: text,
|
||||
});
|
||||
if (!response.data || response.data.length === 0) {
|
||||
throw new Error(
|
||||
`HuggingFace embed() returned no embeddings for model '${this.model}'`,
|
||||
);
|
||||
}
|
||||
return response.data[0].embedding;
|
||||
}
|
||||
|
||||
async embedBatch(texts: string[]): Promise<number[][]> {
|
||||
if (texts.length === 0) {
|
||||
return [];
|
||||
}
|
||||
const response = await this.openai.embeddings.create({
|
||||
model: this.model,
|
||||
input: texts,
|
||||
});
|
||||
const embeddings = response.data
|
||||
.sort((a, b) => a.index - b.index)
|
||||
.map((item) => item.embedding);
|
||||
if (embeddings.length !== texts.length) {
|
||||
throw new Error(
|
||||
`HuggingFace embedBatch() returned ${embeddings.length} embeddings ` +
|
||||
`for ${texts.length} texts using model '${this.model}'`,
|
||||
);
|
||||
}
|
||||
return embeddings;
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,254 @@
|
||||
import type { PredictionServiceClient } from "@google-cloud/aiplatform";
|
||||
import { Embedder } from "./base";
|
||||
import { VertexAIConfig } from "../types";
|
||||
|
||||
type AIPlatform = typeof import("@google-cloud/aiplatform");
|
||||
type ClientOptions = NonNullable<
|
||||
ConstructorParameters<AIPlatform["PredictionServiceClient"]>[0]
|
||||
>;
|
||||
|
||||
interface EmbeddingResponse {
|
||||
embeddings: {
|
||||
values: number[];
|
||||
};
|
||||
}
|
||||
|
||||
/**
|
||||
* Vertex AI caps how many input texts one `predict()` call may carry, and the
|
||||
* cap depends on the model family. `gemini-embedding-*` accepts exactly one
|
||||
* text per request; the older `text-embedding-*` / `text-multilingual-*`
|
||||
* models accept up to 250.
|
||||
* https://cloud.google.com/vertex-ai/generative-ai/docs/embeddings/get-text-embeddings
|
||||
*/
|
||||
function maxInstancesPerRequest(model: string): number {
|
||||
return model.startsWith("gemini-embedding") ? 1 : 250;
|
||||
}
|
||||
|
||||
function isValidEmbedding(value: unknown): value is EmbeddingResponse {
|
||||
if (typeof value !== "object" || value === null) return false;
|
||||
const obj = value as Record<string, unknown>;
|
||||
if (typeof obj.embeddings !== "object" || obj.embeddings === null)
|
||||
return false;
|
||||
const embeddings = obj.embeddings as Record<string, unknown>;
|
||||
const values = embeddings.values;
|
||||
return (
|
||||
Array.isArray(values) &&
|
||||
values.every((v) => typeof v === "number" && Number.isFinite(v))
|
||||
);
|
||||
}
|
||||
|
||||
export class VertexAIEmbedder implements Embedder {
|
||||
private client: PredictionServiceClient | undefined;
|
||||
private helpers: AIPlatform["helpers"] | undefined;
|
||||
private initPromise: Promise<void> | undefined;
|
||||
private clientOptions: ClientOptions;
|
||||
private model: string;
|
||||
private embeddingDims: number;
|
||||
private location: string;
|
||||
private projectId: string;
|
||||
private embeddingTypes: {
|
||||
add: string;
|
||||
update: string;
|
||||
search: string;
|
||||
};
|
||||
|
||||
constructor(config: VertexAIConfig) {
|
||||
this.model = config.model || "gemini-embedding-001";
|
||||
this.embeddingDims = config.embeddingDims || 256;
|
||||
this.location =
|
||||
config.location || process.env.GCP_LOCATION || "us-central1";
|
||||
|
||||
// Left empty when unset: initClient() resolves it from Application Default
|
||||
// Credentials or the service account key file, the way the Python SDK does.
|
||||
this.projectId =
|
||||
config.googleProjectId ||
|
||||
process.env.GCP_PROJECT_ID ||
|
||||
process.env.GOOGLE_CLOUD_PROJECT ||
|
||||
process.env.GCLOUD_PROJECT ||
|
||||
"";
|
||||
|
||||
this.embeddingTypes = {
|
||||
add: config.memoryAddEmbeddingType || "RETRIEVAL_DOCUMENT",
|
||||
update: config.memoryUpdateEmbeddingType || "RETRIEVAL_DOCUMENT",
|
||||
search: config.memorySearchEmbeddingType || "RETRIEVAL_QUERY",
|
||||
};
|
||||
|
||||
const endpoint = `${this.location}-aiplatform.googleapis.com`;
|
||||
this.clientOptions = { apiEndpoint: endpoint };
|
||||
|
||||
if (config.vertexCredentialsJson) {
|
||||
this.clientOptions.keyFilename = config.vertexCredentialsJson;
|
||||
} else if (config.googleServiceAccountJson) {
|
||||
try {
|
||||
this.clientOptions.credentials =
|
||||
typeof config.googleServiceAccountJson === "string"
|
||||
? JSON.parse(config.googleServiceAccountJson)
|
||||
: config.googleServiceAccountJson;
|
||||
} catch (err) {
|
||||
throw new Error(
|
||||
"Failed to parse googleServiceAccountJson: " + (err as Error).message,
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
private async initClient(): Promise<void> {
|
||||
// Memoized so concurrent embed() calls share one client instead of each
|
||||
// racing to build (and leak) their own gRPC channel.
|
||||
if (!this.initPromise) {
|
||||
this.initPromise = this.createClient().catch((err) => {
|
||||
this.initPromise = undefined;
|
||||
throw err;
|
||||
});
|
||||
}
|
||||
await this.initPromise;
|
||||
}
|
||||
|
||||
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 client = new aiplatform.PredictionServiceClient(this.clientOptions);
|
||||
|
||||
if (!this.projectId) {
|
||||
try {
|
||||
this.projectId = await client.getProjectId();
|
||||
} catch (err) {
|
||||
throw new Error(
|
||||
"Vertex AI could not determine a Google Cloud project ID. Set googleProjectId in config, " +
|
||||
"one of the GCP_PROJECT_ID / GOOGLE_CLOUD_PROJECT / GCLOUD_PROJECT env vars, or configure " +
|
||||
"Application Default Credentials: " +
|
||||
(err as Error).message,
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
this.client = client;
|
||||
this.helpers = aiplatform.helpers;
|
||||
}
|
||||
|
||||
private endpoint(): string {
|
||||
return `projects/${this.projectId}/locations/${this.location}/publishers/google/models/${this.model}`;
|
||||
}
|
||||
|
||||
private formatInstance(text: string, taskType: string) {
|
||||
// task_type must live on the instance (snake_case), not in `parameters`.
|
||||
// Vertex silently ignores an unknown `parameters.taskType`, which would
|
||||
// fall back to the model's default task type. This mirrors the Python SDK's
|
||||
// TextEmbeddingInput(text=..., task_type=...).
|
||||
return {
|
||||
content: text,
|
||||
task_type: taskType,
|
||||
};
|
||||
}
|
||||
|
||||
async embed(
|
||||
text: string,
|
||||
memoryAction?: "add" | "update" | "search",
|
||||
): Promise<number[]> {
|
||||
await this.initClient();
|
||||
if (!this.client || !this.helpers) {
|
||||
throw new Error("Client not initialized");
|
||||
}
|
||||
|
||||
let embeddingType = "SEMANTIC_SIMILARITY";
|
||||
if (memoryAction !== undefined) {
|
||||
if (!(memoryAction in this.embeddingTypes)) {
|
||||
throw new Error(`Invalid memory action: ${memoryAction}`);
|
||||
}
|
||||
embeddingType = this.embeddingTypes[memoryAction];
|
||||
}
|
||||
|
||||
const instance = this.formatInstance(text, embeddingType);
|
||||
const parameters = {
|
||||
outputDimensionality: this.embeddingDims,
|
||||
};
|
||||
|
||||
const [response] = await this.client.predict({
|
||||
endpoint: this.endpoint(),
|
||||
instances: [this.helpers.toValue(instance) as any],
|
||||
parameters: this.helpers.toValue(parameters) as any,
|
||||
});
|
||||
|
||||
if (!response.predictions || response.predictions.length === 0) {
|
||||
throw new Error("No predictions returned from Vertex AI");
|
||||
}
|
||||
|
||||
const decoded = this.helpers.fromValue(response.predictions[0] as any);
|
||||
if (!isValidEmbedding(decoded)) {
|
||||
throw new Error("Failed to extract embedding values from response");
|
||||
}
|
||||
|
||||
return decoded.embeddings.values;
|
||||
}
|
||||
|
||||
async embedBatch(
|
||||
texts: string[],
|
||||
memoryAction: "add" | "update" | "search" = "add",
|
||||
): Promise<number[][]> {
|
||||
if (!texts || texts.length === 0) {
|
||||
return [];
|
||||
}
|
||||
|
||||
await this.initClient();
|
||||
if (!this.client || !this.helpers) {
|
||||
throw new Error("Client not initialized");
|
||||
}
|
||||
|
||||
if (!(memoryAction in this.embeddingTypes)) {
|
||||
throw new Error(`Invalid memory action: ${memoryAction}`);
|
||||
}
|
||||
const embeddingType = this.embeddingTypes[memoryAction];
|
||||
|
||||
const allEmbeddings: number[][] = [];
|
||||
const batchSize = maxInstancesPerRequest(this.model);
|
||||
|
||||
for (let i = 0; i < texts.length; i += batchSize) {
|
||||
const chunk = texts.slice(i, i + batchSize);
|
||||
const instances = chunk.map(
|
||||
(text) =>
|
||||
this.helpers!.toValue(
|
||||
this.formatInstance(text, embeddingType),
|
||||
) as any,
|
||||
);
|
||||
const parameters = {
|
||||
outputDimensionality: this.embeddingDims,
|
||||
};
|
||||
|
||||
const [response] = await this.client.predict({
|
||||
endpoint: this.endpoint(),
|
||||
instances,
|
||||
parameters: this.helpers.toValue(parameters) as any,
|
||||
});
|
||||
|
||||
if (!response.predictions || response.predictions.length === 0) {
|
||||
throw new Error("No predictions returned from Vertex AI batch request");
|
||||
}
|
||||
|
||||
for (const prediction of response.predictions) {
|
||||
const decoded = this.helpers.fromValue(prediction as any);
|
||||
if (!isValidEmbedding(decoded)) {
|
||||
throw new Error(
|
||||
"Failed to extract embedding values from batch response",
|
||||
);
|
||||
}
|
||||
allEmbeddings.push(decoded.embeddings.values);
|
||||
}
|
||||
}
|
||||
|
||||
if (allEmbeddings.length !== texts.length) {
|
||||
throw new Error(
|
||||
`Vertex AI embedBatch() returned ${allEmbeddings.length} embeddings for ${texts.length} texts using model '${this.model}'`,
|
||||
);
|
||||
}
|
||||
|
||||
return allEmbeddings;
|
||||
}
|
||||
}
|
||||
@@ -2,6 +2,7 @@ export * from "./memory";
|
||||
export * from "./memory/memory.types";
|
||||
export * from "./types";
|
||||
export * from "./embeddings/base";
|
||||
export * from "./embeddings/huggingface";
|
||||
export * from "./embeddings/openai";
|
||||
export * from "./embeddings/ollama";
|
||||
export * from "./embeddings/lmstudio";
|
||||
@@ -9,6 +10,7 @@ export * from "./embeddings/together";
|
||||
export * from "./embeddings/google";
|
||||
export * from "./embeddings/azure";
|
||||
export * from "./embeddings/langchain";
|
||||
export * from "./embeddings/vertexai";
|
||||
export * from "./embeddings/fastembed";
|
||||
export * from "./llms/base";
|
||||
export * from "./llms/openai";
|
||||
@@ -22,8 +24,10 @@ export * from "./llms/mistral";
|
||||
export * from "./llms/langchain";
|
||||
export * from "./llms/litellm";
|
||||
export * from "./llms/vllm";
|
||||
export * from "./llms/aws_bedrock";
|
||||
export * from "./vector_stores/base";
|
||||
export * from "./vector_stores/memory";
|
||||
export * from "./vector_stores/baidu";
|
||||
export * from "./vector_stores/qdrant";
|
||||
export * from "./vector_stores/redis";
|
||||
export * from "./vector_stores/valkey";
|
||||
@@ -32,6 +36,8 @@ export * from "./vector_stores/langchain";
|
||||
export * from "./vector_stores/vectorize";
|
||||
export * from "./vector_stores/azure_ai_search";
|
||||
export * from "./vector_stores/pgvector";
|
||||
export * from "./vector_stores/databricks";
|
||||
export * from "./vector_stores/neptune_analytics";
|
||||
export * from "./vector_stores/elasticsearch";
|
||||
export * from "./vector_stores/upstash_vector";
|
||||
export * from "./vector_stores/azure_mysql";
|
||||
@@ -40,6 +46,13 @@ export * from "./vector_stores/s3_vectors";
|
||||
export * from "./vector_stores/vertex_ai_vector_search";
|
||||
export * from "./vector_stores/pinecone";
|
||||
export * from "./vector_stores/turbopuffer";
|
||||
export * from "./vector_stores/milvus";
|
||||
export * from "./vector_stores/mongodb";
|
||||
export * from "./vector_stores/opensearch";
|
||||
export * from "./vector_stores/weaviate";
|
||||
export * from "./rerankers/base";
|
||||
export * from "./rerankers/cohere";
|
||||
export * from "./rerankers/llm";
|
||||
export * from "./rerankers/zeroentropy";
|
||||
export * from "./rerankers/cross_encoder";
|
||||
export * from "./utils/factory";
|
||||
|
||||
@@ -0,0 +1,294 @@
|
||||
import { LLM, LLMResponse } from "./base";
|
||||
import { LLMConfig, Message } from "../types";
|
||||
|
||||
/**
|
||||
* Providers recognised in Bedrock model identifiers, mirroring the Python
|
||||
* provider's `PROVIDERS` list (`mem0/llms/aws_bedrock.py`).
|
||||
*/
|
||||
const PROVIDERS = [
|
||||
"ai21",
|
||||
"amazon",
|
||||
"anthropic",
|
||||
"cohere",
|
||||
"meta",
|
||||
"mistral",
|
||||
"stability",
|
||||
"writer",
|
||||
"deepseek",
|
||||
"gpt-oss",
|
||||
"perplexity",
|
||||
"snowflake",
|
||||
"titan",
|
||||
"command",
|
||||
"j2",
|
||||
"llama",
|
||||
"minimax",
|
||||
];
|
||||
|
||||
/**
|
||||
* Extract the model-family provider from a Bedrock model id
|
||||
* (e.g. `anthropic.claude-3-sonnet-...` -> `anthropic`).
|
||||
*/
|
||||
export function extractProvider(model: string): string {
|
||||
for (const provider of PROVIDERS) {
|
||||
const re = new RegExp(
|
||||
`\\b${provider.replace(/[.*+?^${}()|[\]\\]/g, "\\$&")}\\b`,
|
||||
);
|
||||
if (re.test(model)) return provider;
|
||||
}
|
||||
throw new Error(`Unknown provider in model: ${model}`);
|
||||
}
|
||||
|
||||
/**
|
||||
* AWS Bedrock fields (awsRegion / awsAccessKeyId / awsSecretAccessKey /
|
||||
* awsSessionToken / client) now live on the shared `LLMConfig`, so the
|
||||
* provider is configurable through the standard typed `Memory` config path.
|
||||
*/
|
||||
type AWSBedrockConfig = LLMConfig;
|
||||
|
||||
/**
|
||||
* AWS Bedrock LLM provider for the TypeScript OSS SDK.
|
||||
*
|
||||
* Mirrors `mem0/llms/aws_bedrock.py`. Uses the Bedrock **Converse API**
|
||||
* (`ConverseCommand`), which provides a uniform message/tool interface across
|
||||
* the Anthropic / Amazon (Nova) / Meta / Mistral / Cohere model families, so a
|
||||
* single code path serves them all (the Python provider keeps per-family
|
||||
* `invoke_model` branches for legacy reasons; Converse supersedes them).
|
||||
*
|
||||
* The `@aws-sdk/client-bedrock-runtime` dependency is loaded on first use via
|
||||
* dynamic `import()` so the package stays optional. Credentials resolve via the
|
||||
* standard AWS chain unless provided explicitly in config.
|
||||
*/
|
||||
interface BedrockSDK {
|
||||
BedrockRuntimeClient: new (config: Record<string, any>) => any;
|
||||
ConverseCommand: new (input: Record<string, any>) => any;
|
||||
}
|
||||
|
||||
export class AWSBedrockLLM implements LLM {
|
||||
private model: string;
|
||||
private provider: string;
|
||||
private temperature: number;
|
||||
private maxTokens: number;
|
||||
private topP?: number;
|
||||
private clientConfig: Record<string, any>;
|
||||
private clientOverride?: any;
|
||||
private sdkPromise?: Promise<BedrockSDK>;
|
||||
private clientPromise?: Promise<any>;
|
||||
|
||||
constructor(config: AWSBedrockConfig = {}) {
|
||||
this.model =
|
||||
(typeof config.model === "string" && config.model) ||
|
||||
"anthropic.claude-3-5-sonnet-20240620-v1:0";
|
||||
this.provider = extractProvider(this.model);
|
||||
this.temperature = config.temperature ?? 0.1;
|
||||
this.maxTokens = config.maxTokens ?? 2000;
|
||||
this.topP = config.topP;
|
||||
|
||||
const region =
|
||||
config.awsRegion ||
|
||||
process.env.AWS_REGION ||
|
||||
process.env.AWS_DEFAULT_REGION;
|
||||
const clientConfig: Record<string, any> = {};
|
||||
if (region) clientConfig.region = region;
|
||||
if (config.awsAccessKeyId && config.awsSecretAccessKey) {
|
||||
clientConfig.credentials = {
|
||||
accessKeyId: config.awsAccessKeyId,
|
||||
secretAccessKey: config.awsSecretAccessKey,
|
||||
...(config.awsSessionToken && {
|
||||
sessionToken: config.awsSessionToken,
|
||||
}),
|
||||
};
|
||||
}
|
||||
|
||||
this.clientConfig = clientConfig;
|
||||
this.clientOverride = config.client;
|
||||
}
|
||||
|
||||
/**
|
||||
* Load the optional AWS SDK on first use.
|
||||
*
|
||||
* This MUST be a dynamic `import()`, never `require()`: tsup/esbuild rewrite
|
||||
* `require()` in the published ESM bundle (`dist/oss/index.mjs`) into a
|
||||
* `__require` shim that throws `Dynamic require of "..." is not supported`,
|
||||
* so every ESM consumer would hit a dead provider even with the SDK installed.
|
||||
*/
|
||||
private async getSDK(): Promise<BedrockSDK> {
|
||||
if (!this.sdkPromise) {
|
||||
this.sdkPromise = import("@aws-sdk/client-bedrock-runtime").then(
|
||||
(sdk) => sdk as unknown as BedrockSDK,
|
||||
(err) => {
|
||||
// Let a later call retry rather than caching the rejection forever.
|
||||
this.sdkPromise = undefined;
|
||||
const detail = err instanceof Error ? err.message : String(err);
|
||||
throw new Error(
|
||||
"The '@aws-sdk/client-bedrock-runtime' package is required to use the AWS Bedrock LLM provider. " +
|
||||
`Install it with: npm install @aws-sdk/client-bedrock-runtime (original error: ${detail})`,
|
||||
);
|
||||
},
|
||||
);
|
||||
}
|
||||
return this.sdkPromise;
|
||||
}
|
||||
|
||||
/** Memoized Bedrock client; an injected `config.client` short-circuits the SDK. */
|
||||
private async getClient(): Promise<any> {
|
||||
if (this.clientOverride) return this.clientOverride;
|
||||
if (!this.clientPromise) {
|
||||
this.clientPromise = this.getSDK().then(
|
||||
({ BedrockRuntimeClient }) =>
|
||||
new BedrockRuntimeClient(this.clientConfig),
|
||||
);
|
||||
}
|
||||
return this.clientPromise;
|
||||
}
|
||||
|
||||
/**
|
||||
* Split messages into a top-level `system` block (Converse passes system
|
||||
* prompts separately) and role-tagged content blocks for everything else.
|
||||
*/
|
||||
private formatMessages(messages: Message[]): {
|
||||
system?: { text: string }[];
|
||||
converseMessages: { role: string; content: { text: string }[] }[];
|
||||
} {
|
||||
const systemParts: string[] = [];
|
||||
const converseMessages: { role: string; content: { text: string }[] }[] =
|
||||
[];
|
||||
|
||||
for (const msg of messages) {
|
||||
const role = msg.role;
|
||||
const content =
|
||||
typeof msg.content === "string"
|
||||
? msg.content
|
||||
: JSON.stringify(msg.content);
|
||||
if (role === "system") {
|
||||
systemParts.push(content);
|
||||
} else {
|
||||
converseMessages.push({
|
||||
role: role === "assistant" ? "assistant" : "user",
|
||||
content: [{ text: content }],
|
||||
});
|
||||
}
|
||||
}
|
||||
|
||||
if (converseMessages.length === 0) {
|
||||
converseMessages.push({ role: "user", content: [{ text: "" }] });
|
||||
}
|
||||
|
||||
return {
|
||||
system: systemParts.length
|
||||
? [{ text: systemParts.join("\n") }]
|
||||
: undefined,
|
||||
converseMessages,
|
||||
};
|
||||
}
|
||||
|
||||
/**
|
||||
* Build the Converse `inferenceConfig`. Anthropic and MiniMax reasoning
|
||||
* models reject requests carrying both `temperature` and `topP`, so `topP`
|
||||
* is omitted for those families (mirrors the Python `_build_inference_config`).
|
||||
*/
|
||||
private buildInferenceConfig(): Record<string, any> {
|
||||
const inferenceConfig: Record<string, any> = {
|
||||
maxTokens: this.maxTokens,
|
||||
temperature: this.temperature,
|
||||
};
|
||||
if (
|
||||
this.topP != null &&
|
||||
!["anthropic", "minimax"].includes(this.provider)
|
||||
) {
|
||||
inferenceConfig.topP = this.topP;
|
||||
}
|
||||
return inferenceConfig;
|
||||
}
|
||||
|
||||
/** Convert OpenAI-style tools to the Converse `toolConfig` shape. */
|
||||
private convertToolsToConverse(tools: any[]): any | undefined {
|
||||
if (!tools || tools.length === 0) return undefined;
|
||||
const converseTools = tools
|
||||
.filter((t) => t?.type === "function" && t.function)
|
||||
.map((t) => ({
|
||||
toolSpec: {
|
||||
name: t.function.name,
|
||||
description: t.function.description || "",
|
||||
inputSchema: { json: t.function.parameters || {} },
|
||||
},
|
||||
}));
|
||||
return converseTools.length ? { tools: converseTools } : undefined;
|
||||
}
|
||||
|
||||
private async converse(messages: Message[], tools?: any[]): Promise<any> {
|
||||
const { system, converseMessages } = this.formatMessages(messages);
|
||||
const input: Record<string, any> = {
|
||||
modelId: this.model,
|
||||
messages: converseMessages,
|
||||
inferenceConfig: this.buildInferenceConfig(),
|
||||
};
|
||||
if (system) input.system = system;
|
||||
const toolConfig = tools ? this.convertToolsToConverse(tools) : undefined;
|
||||
if (toolConfig) input.toolConfig = toolConfig;
|
||||
|
||||
const [{ ConverseCommand }, client] = await Promise.all([
|
||||
this.getSDK(),
|
||||
this.getClient(),
|
||||
]);
|
||||
return client.send(new ConverseCommand(input));
|
||||
}
|
||||
|
||||
/** Pull the first text block out of a Converse response. */
|
||||
private parseText(response: any): string {
|
||||
const content = response?.output?.message?.content || [];
|
||||
for (const block of content) {
|
||||
if (block && typeof block.text === "string") return block.text;
|
||||
}
|
||||
return "";
|
||||
}
|
||||
|
||||
/** Collect any toolUse blocks out of a Converse response. */
|
||||
private parseToolCalls(response: any): { name: string; arguments: string }[] {
|
||||
const content = response?.output?.message?.content || [];
|
||||
const calls: { name: string; arguments: string }[] = [];
|
||||
for (const block of content) {
|
||||
if (block?.toolUse) {
|
||||
calls.push({
|
||||
name: block.toolUse.name,
|
||||
arguments: JSON.stringify(block.toolUse.input ?? {}),
|
||||
});
|
||||
}
|
||||
}
|
||||
return calls;
|
||||
}
|
||||
|
||||
async generateResponse(
|
||||
messages: Message[],
|
||||
_responseFormat?: { type: string },
|
||||
tools?: any[],
|
||||
): Promise<string | LLMResponse> {
|
||||
try {
|
||||
const response = await this.converse(messages, tools);
|
||||
if (tools && tools.length) {
|
||||
const toolCalls = this.parseToolCalls(response);
|
||||
if (toolCalls.length) {
|
||||
return {
|
||||
content: this.parseText(response),
|
||||
role: "assistant",
|
||||
toolCalls,
|
||||
};
|
||||
}
|
||||
}
|
||||
return this.parseText(response);
|
||||
} catch (err) {
|
||||
const message = err instanceof Error ? err.message : String(err);
|
||||
throw new Error(`AWS Bedrock LLM failed: ${message}`);
|
||||
}
|
||||
}
|
||||
|
||||
async generateChat(messages: Message[]): Promise<LLMResponse> {
|
||||
try {
|
||||
const response = await this.converse(messages);
|
||||
return { content: this.parseText(response), role: "assistant" };
|
||||
} catch (err) {
|
||||
const message = err instanceof Error ? err.message : String(err);
|
||||
throw new Error(`AWS Bedrock LLM failed: ${message}`);
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,52 @@
|
||||
import { OpenAILLM } from "./openai";
|
||||
import { LLMConfig, Message } from "../types";
|
||||
import { LLMResponse } from "./base";
|
||||
|
||||
/**
|
||||
* Sarvam AI LLM provider.
|
||||
*
|
||||
* Sarvam's API is OpenAI-compatible, so this simply reuses {@link OpenAILLM}
|
||||
* and overrides the connection defaults — mirroring `mem0/llms/sarvam.py` in the
|
||||
* Python SDK. The API key resolves from `config.apiKey` or the `SARVAM_API_KEY`
|
||||
* env var, and the base URL from `config.baseURL`, `SARVAM_API_BASE`, else
|
||||
* `https://api.sarvam.ai/v1`.
|
||||
*/
|
||||
export class SarvamLLM extends OpenAILLM {
|
||||
constructor(config: LLMConfig) {
|
||||
const apiKey = config.apiKey || process.env.SARVAM_API_KEY;
|
||||
if (!apiKey) {
|
||||
throw new Error("Sarvam API key is required");
|
||||
}
|
||||
super({
|
||||
...config,
|
||||
apiKey,
|
||||
baseURL:
|
||||
config.baseURL ||
|
||||
process.env.SARVAM_API_BASE ||
|
||||
"https://api.sarvam.ai/v1",
|
||||
model: config.model || "sarvam-m",
|
||||
});
|
||||
}
|
||||
|
||||
async generateResponse(
|
||||
messages: Message[],
|
||||
responseFormat?: { type: string },
|
||||
tools?: any[],
|
||||
): Promise<string | LLMResponse> {
|
||||
try {
|
||||
return await super.generateResponse(messages, responseFormat, tools);
|
||||
} catch (err) {
|
||||
const message = err instanceof Error ? err.message : String(err);
|
||||
throw new Error(`Sarvam LLM failed: ${message}`);
|
||||
}
|
||||
}
|
||||
|
||||
async generateChat(messages: Message[]): Promise<LLMResponse> {
|
||||
try {
|
||||
return await super.generateChat(messages);
|
||||
} catch (err) {
|
||||
const message = err instanceof Error ? err.message : String(err);
|
||||
throw new Error(`Sarvam LLM failed: ${message}`);
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -7,12 +7,14 @@ import {
|
||||
Message,
|
||||
SearchFilters,
|
||||
SearchResult,
|
||||
VectorStoreConfig,
|
||||
} from "../types";
|
||||
import {
|
||||
EmbedderFactory,
|
||||
LLMFactory,
|
||||
VectorStoreFactory,
|
||||
HistoryManagerFactory,
|
||||
RerankerFactory,
|
||||
} from "../utils/factory";
|
||||
import {
|
||||
FactRetrievalSchema,
|
||||
@@ -28,6 +30,7 @@ import {
|
||||
import { DummyHistoryManager } from "../storage/DummyHistoryManager";
|
||||
import { Embedder } from "../embeddings/base";
|
||||
import { LLM } from "../llms/base";
|
||||
import { Reranker } from "../rerankers/base";
|
||||
import { VectorStore } from "../vector_stores/base";
|
||||
import { ConfigManager } from "../config/manager";
|
||||
|
||||
@@ -36,6 +39,7 @@ import {
|
||||
SearchMemoryOptions,
|
||||
DeleteAllMemoryOptions,
|
||||
GetAllMemoryOptions,
|
||||
UpdateMemoryOptions,
|
||||
UpdateProjectOptions,
|
||||
} from "./memory.types";
|
||||
import { parse_vision_messages } from "../utils/memory";
|
||||
@@ -72,9 +76,22 @@ import {
|
||||
ScoredResult,
|
||||
} from "../utils/scoring";
|
||||
import { getDefaultVectorStoreDbPath } from "../utils/sqlite";
|
||||
import { logger } from "../utils/logger";
|
||||
import { normalizeExpirationDate, payloadIsExpired } from "../utils/expiration";
|
||||
import { getOrCreateMem0UserId } from "../../../client/config";
|
||||
|
||||
// Entity params that must be passed via filters - check both snake_case and camelCase
|
||||
export class LLMError extends Error {
|
||||
readonly cause?: unknown;
|
||||
|
||||
constructor(message: string, options: { cause?: unknown } = {}) {
|
||||
super(message);
|
||||
this.name = "LLMError";
|
||||
this.cause = options.cause;
|
||||
Object.setPrototypeOf(this, new.target.prototype);
|
||||
}
|
||||
}
|
||||
|
||||
// Entity params that must be passed via filters check both snake_case and camelCase
|
||||
const ENTITY_PARAMS = [
|
||||
"user_id",
|
||||
"agent_id",
|
||||
@@ -161,6 +178,7 @@ export class Memory {
|
||||
private embedder: Embedder;
|
||||
private vectorStore!: VectorStore;
|
||||
private llm: LLM;
|
||||
private reranker: Reranker | null = null;
|
||||
private db: HistoryManager;
|
||||
private collectionName: string | undefined;
|
||||
private apiVersion: string;
|
||||
@@ -185,6 +203,12 @@ export class Memory {
|
||||
this.config.llm.provider,
|
||||
this.config.llm.config,
|
||||
);
|
||||
if (this.config.reranker) {
|
||||
this.reranker = RerankerFactory.create(
|
||||
this.config.reranker.provider,
|
||||
this.config.reranker.config,
|
||||
);
|
||||
}
|
||||
if (this.config.disableHistory) {
|
||||
this.db = new DummyHistoryManager();
|
||||
} else {
|
||||
@@ -263,18 +287,24 @@ export class Memory {
|
||||
|
||||
private async getEntityStore(): Promise<VectorStore> {
|
||||
if (!this._entityStore) {
|
||||
const entityProvider = this.config.vectorStore.provider;
|
||||
const entityCollectionName = `${this.collectionName}_entities`;
|
||||
const entityConfig = {
|
||||
const entityConfig: VectorStoreConfig = {
|
||||
...this.config.vectorStore.config,
|
||||
collectionName: entityCollectionName,
|
||||
};
|
||||
// For file-based stores (memory/SQLite), always use a separate DB for entities
|
||||
if (this.config.vectorStore.provider === "memory") {
|
||||
if (entityProvider === "memory") {
|
||||
const basePath = entityConfig.dbPath || getDefaultVectorStoreDbPath();
|
||||
entityConfig.dbPath = basePath.replace(/\.db$/, "_entities.db");
|
||||
}
|
||||
if (entityProvider === "databricks") {
|
||||
entityConfig.tableName = entityConfig.tableName
|
||||
? `${entityConfig.tableName}_entities`
|
||||
: entityCollectionName;
|
||||
}
|
||||
this._entityStore = VectorStoreFactory.create(
|
||||
this.config.vectorStore.provider,
|
||||
entityProvider,
|
||||
entityConfig,
|
||||
);
|
||||
await this._entityStore.initialize();
|
||||
@@ -396,7 +426,7 @@ export class Memory {
|
||||
}
|
||||
let vec: number[];
|
||||
try {
|
||||
vec = await this.embedder.embed(entityText);
|
||||
vec = await this.embedder.embed(entityText, "update");
|
||||
} catch (e) {
|
||||
console.debug(`Entity re-embed failed for '${entityText}': ${e}`);
|
||||
continue;
|
||||
@@ -440,7 +470,7 @@ export class Memory {
|
||||
try {
|
||||
let entityVec: number[];
|
||||
try {
|
||||
entityVec = await this.embedder.embed(entity.text);
|
||||
entityVec = await this.embedder.embed(entity.text, "add");
|
||||
} catch (e) {
|
||||
console.debug(`Entity embed failed for '${entity.text}': ${e}`);
|
||||
continue;
|
||||
@@ -713,6 +743,11 @@ export class Memory {
|
||||
if (agentId) filters.agent_id = metadata.agent_id = agentId;
|
||||
if (runId) filters.run_id = metadata.run_id = runId;
|
||||
|
||||
// Normalize expiration date into the stored metadata (round-trips via get()).
|
||||
if (config.expirationDate != null) {
|
||||
metadata.expiration_date = normalizeExpirationDate(config.expirationDate);
|
||||
}
|
||||
|
||||
if (!filters.user_id && !filters.agent_id && !filters.run_id) {
|
||||
throw new Error(
|
||||
"One of the filters: userId, agentId or runId is required!",
|
||||
@@ -810,7 +845,7 @@ export class Memory {
|
||||
.join("\n");
|
||||
|
||||
// Phase 1: Existing memory retrieval
|
||||
const queryEmbedding = await this.embedder.embed(parsedMessages);
|
||||
const queryEmbedding = await this.embedder.embed(parsedMessages, "search");
|
||||
const existingResults = await this.vectorStore.search(
|
||||
queryEmbedding,
|
||||
10,
|
||||
@@ -854,7 +889,7 @@ export class Memory {
|
||||
)) as string;
|
||||
} catch (e) {
|
||||
console.error("LLM extraction failed:", e);
|
||||
return [];
|
||||
throw new LLMError(`LLM extraction failed: ${e}`, { cause: e });
|
||||
}
|
||||
|
||||
// Parse response
|
||||
@@ -904,7 +939,7 @@ export class Memory {
|
||||
.filter((t) => t.length > 0);
|
||||
let embedMap: Record<string, number[]> = {};
|
||||
try {
|
||||
const memEmbeddingsList = await this.embedder.embedBatch(memTexts);
|
||||
const memEmbeddingsList = await this.embedder.embedBatch(memTexts, "add");
|
||||
for (let i = 0; i < memTexts.length; i++) {
|
||||
embedMap[memTexts[i]] = memEmbeddingsList[i];
|
||||
}
|
||||
@@ -912,7 +947,7 @@ export class Memory {
|
||||
// Fallback: embed individually
|
||||
for (const text of memTexts) {
|
||||
try {
|
||||
embedMap[text] = await this.embedder.embed(text);
|
||||
embedMap[text] = await this.embedder.embed(text, "add");
|
||||
} catch (e) {
|
||||
console.warn(`Failed to embed memory text: ${e}`);
|
||||
}
|
||||
@@ -1090,13 +1125,13 @@ export class Memory {
|
||||
// 7b: Single batch embed for all unique entities
|
||||
let entityEmbeddings: (number[] | null)[];
|
||||
try {
|
||||
entityEmbeddings = await this.embedder.embedBatch(entityTexts);
|
||||
entityEmbeddings = await this.embedder.embedBatch(entityTexts, "add");
|
||||
} catch {
|
||||
// Fallback: embed individually
|
||||
entityEmbeddings = [];
|
||||
for (const t of entityTexts) {
|
||||
try {
|
||||
entityEmbeddings.push(await this.embedder.embed(t));
|
||||
entityEmbeddings.push(await this.embedder.embed(t, "add"));
|
||||
} catch {
|
||||
entityEmbeddings.push(null);
|
||||
}
|
||||
@@ -1307,7 +1342,12 @@ export class Memory {
|
||||
: {};
|
||||
|
||||
await this._ensureInitialized();
|
||||
const { topK = 20, threshold = 0.1, explain = false } = config;
|
||||
const {
|
||||
topK = 20,
|
||||
threshold = 0.1,
|
||||
explain = false,
|
||||
showExpired = false,
|
||||
} = config;
|
||||
|
||||
await this._captureEvent("search", {
|
||||
query_length: query.length,
|
||||
@@ -1355,7 +1395,7 @@ export class Memory {
|
||||
const queryEntities = extractEntities(query);
|
||||
|
||||
// Step 2: Embed query
|
||||
const queryEmbedding = await this.embedder.embed(query);
|
||||
const queryEmbedding = await this.embedder.embed(query, "search");
|
||||
|
||||
// Step 3: Semantic search (over-fetch for scoring pool)
|
||||
const internalLimit = Math.max(topK * 4, 60);
|
||||
@@ -1420,7 +1460,10 @@ export class Memory {
|
||||
entitySearchFilters[k] = effectiveFilters[k];
|
||||
}
|
||||
const entityTexts = deduped.map((e) => e.text);
|
||||
const embeddings = await this.embedder.embedBatch(entityTexts);
|
||||
const embeddings = await this.embedder.embedBatch(
|
||||
entityTexts,
|
||||
"search",
|
||||
);
|
||||
|
||||
if (embeddings.length !== entityTexts.length) {
|
||||
console.warn(
|
||||
@@ -1475,11 +1518,13 @@ export class Memory {
|
||||
}
|
||||
|
||||
// Step 7: Build candidate set from semantic results
|
||||
const candidates = semanticResults.map((mem) => ({
|
||||
id: String(mem.id),
|
||||
score: mem.score ?? 0,
|
||||
payload: mem.payload || {},
|
||||
}));
|
||||
const candidates = semanticResults
|
||||
.filter((mem) => showExpired || !payloadIsExpired(mem.payload))
|
||||
.map((mem) => ({
|
||||
id: String(mem.id),
|
||||
score: mem.score ?? 0,
|
||||
payload: mem.payload || {},
|
||||
}));
|
||||
|
||||
// Step 8: Score and rank
|
||||
const scoredResults = scoreAndRank(
|
||||
@@ -1526,8 +1571,30 @@ export class Memory {
|
||||
};
|
||||
});
|
||||
|
||||
// Step 10: Optionally re-rank with the configured reranker. Opt-in per
|
||||
// search via `rerank: true`; a no-op when no reranker is configured.
|
||||
const invokeReranker = Boolean(
|
||||
config.rerank && this.reranker && results.length > 0,
|
||||
);
|
||||
let finalResults = results;
|
||||
if (invokeReranker) {
|
||||
try {
|
||||
const ranked = await this.reranker!.rerank(
|
||||
query,
|
||||
results.map((r) => r.memory),
|
||||
topK,
|
||||
);
|
||||
finalResults = ranked.map((r) => ({
|
||||
...results[r.index],
|
||||
rerankScore: r.rerankScore,
|
||||
}));
|
||||
} catch (e) {
|
||||
console.warn(`Reranking failed, using original results: ${e}`);
|
||||
}
|
||||
}
|
||||
|
||||
const result = {
|
||||
results,
|
||||
results: finalResults,
|
||||
};
|
||||
const searchElapsedMs = Date.now() - searchStartMs;
|
||||
if (temporalUsageNotice) {
|
||||
@@ -1563,11 +1630,49 @@ export class Memory {
|
||||
return result;
|
||||
}
|
||||
|
||||
async update(memoryId: string, data: string): Promise<{ message: string }> {
|
||||
async update(
|
||||
memoryId: string,
|
||||
config: string | UpdateMemoryOptions,
|
||||
): Promise<{ message: string }> {
|
||||
await this._ensureInitialized();
|
||||
await this._captureEvent("update", { memory_id: memoryId });
|
||||
const embedding = await this.embedder.embed(data);
|
||||
await this.updateMemory(memoryId, data, { [data]: embedding });
|
||||
|
||||
const options: UpdateMemoryOptions =
|
||||
typeof config === "string" ? { text: config } : config;
|
||||
|
||||
const { data, metadata, expirationDate } = options;
|
||||
let text = options.text;
|
||||
|
||||
if (data != null) {
|
||||
logger.warn(
|
||||
"The `data` option of update() is deprecated and will be removed in " +
|
||||
"the next major release. Use `text` instead.",
|
||||
);
|
||||
if (text == null) {
|
||||
text = data;
|
||||
}
|
||||
}
|
||||
|
||||
if (text == null && metadata == null && expirationDate === undefined) {
|
||||
throw new Error(
|
||||
"At least one of text, metadata, or expirationDate must be provided.",
|
||||
);
|
||||
}
|
||||
|
||||
const updateMetadata: Record<string, any> = { ...metadata };
|
||||
if (expirationDate !== undefined) {
|
||||
updateMetadata.expiration_date =
|
||||
expirationDate === null
|
||||
? null
|
||||
: normalizeExpirationDate(expirationDate);
|
||||
}
|
||||
|
||||
const existingEmbeddings: Record<string, number[]> = {};
|
||||
if (text != null) {
|
||||
existingEmbeddings[text] = await this.embedder.embed(text, "update");
|
||||
}
|
||||
|
||||
await this.updateMemory(memoryId, text, existingEmbeddings, updateMetadata);
|
||||
const result = { message: "Memory updated successfully!" };
|
||||
await this._displayFirstRunNotice("update");
|
||||
return result;
|
||||
@@ -1649,7 +1754,7 @@ export class Memory {
|
||||
await this.db.reset();
|
||||
|
||||
// Check provider before attempting deleteCol
|
||||
if (this.config.vectorStore.provider.toLowerCase() !== "langchain") {
|
||||
if (this.config.vectorStore.provider !== "langchain") {
|
||||
try {
|
||||
await this.vectorStore.deleteCol();
|
||||
} catch (e) {
|
||||
@@ -1704,7 +1809,7 @@ export class Memory {
|
||||
|
||||
await this._ensureInitialized();
|
||||
|
||||
const { topK = 20 } = config;
|
||||
const { topK = 20, showExpired = false } = config;
|
||||
|
||||
// Validate and trim entity IDs in filters. Drop keys that resolve to
|
||||
// undefined so downstream vector stores don't receive
|
||||
@@ -1733,7 +1838,13 @@ export class Memory {
|
||||
);
|
||||
}
|
||||
|
||||
const [memories] = await this.vectorStore.list(filters, topK);
|
||||
// Over-fetch so expired memories dropped below still leave topK survivors.
|
||||
const fetchLimit = showExpired ? topK : Math.max(topK * 4, 60);
|
||||
const [memories] = await this.vectorStore.list(filters, fetchLimit);
|
||||
|
||||
const visibleMemories = showExpired
|
||||
? memories
|
||||
: memories.filter((mem) => !payloadIsExpired(mem.payload));
|
||||
|
||||
const excludedKeys = new Set([
|
||||
"user_id",
|
||||
@@ -1746,7 +1857,7 @@ export class Memory {
|
||||
"textLemmatized",
|
||||
"attributedTo",
|
||||
]);
|
||||
const results = memories.map((mem) => ({
|
||||
const results = visibleMemories.slice(0, topK).map((mem) => ({
|
||||
id: mem.id,
|
||||
memory: mem.payload.data,
|
||||
hash: mem.payload.hash,
|
||||
@@ -1783,7 +1894,7 @@ export class Memory {
|
||||
): Promise<string> {
|
||||
const memoryId = uuidv4();
|
||||
const embedding =
|
||||
existingEmbeddings[data] || (await this.embedder.embed(data));
|
||||
existingEmbeddings[data] || (await this.embedder.embed(data, "add"));
|
||||
|
||||
const memoryMetadata = {
|
||||
...metadata,
|
||||
@@ -1807,7 +1918,7 @@ export class Memory {
|
||||
|
||||
private async updateMemory(
|
||||
memoryId: string,
|
||||
data: string,
|
||||
data: string | undefined,
|
||||
existingEmbeddings: Record<string, number[]>,
|
||||
metadata: Record<string, any> = {},
|
||||
): Promise<string> {
|
||||
@@ -1817,15 +1928,25 @@ export class Memory {
|
||||
}
|
||||
|
||||
const prevValue = existingMemory.payload.data;
|
||||
// Metadata-only update: fall back to the stored text so we can re-index it.
|
||||
const newData = data ?? prevValue;
|
||||
if (typeof newData !== "string") {
|
||||
throw new Error(
|
||||
`Memory with ID ${memoryId} does not have text content to update`,
|
||||
);
|
||||
}
|
||||
const textChanged = newData !== prevValue;
|
||||
|
||||
const embedding =
|
||||
existingEmbeddings[data] || (await this.embedder.embed(data));
|
||||
existingEmbeddings[newData] ||
|
||||
(await this.embedder.embed(newData, "update"));
|
||||
|
||||
const newMetadata = {
|
||||
...existingMemory.payload,
|
||||
...metadata,
|
||||
data,
|
||||
hash: createHash("md5").update(data).digest("hex"),
|
||||
textLemmatized: lemmatizeForBm25(data),
|
||||
data: newData,
|
||||
hash: createHash("md5").update(newData).digest("hex"),
|
||||
textLemmatized: lemmatizeForBm25(newData),
|
||||
createdAt: existingMemory.payload.createdAt,
|
||||
updatedAt: new Date().toISOString(),
|
||||
};
|
||||
@@ -1834,20 +1955,22 @@ export class Memory {
|
||||
await this.db.addHistory(
|
||||
memoryId,
|
||||
prevValue,
|
||||
data,
|
||||
newData,
|
||||
"UPDATE",
|
||||
newMetadata.createdAt,
|
||||
newMetadata.updatedAt,
|
||||
);
|
||||
|
||||
// Entity-store cleanup: strip this memory's id from old-text entities,
|
||||
// then re-extract entities from the new text and link them back.
|
||||
try {
|
||||
const sessionFilters = this._sessionFiltersFromPayload(newMetadata);
|
||||
await this._removeMemoryFromEntityStore(memoryId, sessionFilters);
|
||||
await this._linkEntitiesForMemory(memoryId, data, sessionFilters);
|
||||
} catch (e) {
|
||||
console.warn(`Entity store cleanup/link failed during update: ${e}`);
|
||||
// Entity-store cleanup only when the text changed: strip this memory's id
|
||||
// from old-text entities, then re-extract from the new text and link back.
|
||||
if (textChanged) {
|
||||
try {
|
||||
const sessionFilters = this._sessionFiltersFromPayload(newMetadata);
|
||||
await this._removeMemoryFromEntityStore(memoryId, sessionFilters);
|
||||
await this._linkEntitiesForMemory(memoryId, newData, sessionFilters);
|
||||
} catch (e) {
|
||||
console.warn(`Entity store cleanup/link failed during update: ${e}`);
|
||||
}
|
||||
}
|
||||
|
||||
return memoryId;
|
||||
|
||||
@@ -12,6 +12,22 @@ export interface AddMemoryOptions extends Entity {
|
||||
filters?: SearchFilters;
|
||||
infer?: boolean;
|
||||
timestamp?: number | string | Date | null;
|
||||
/** Date (YYYY-MM-DD) after which the memory is considered expired. */
|
||||
expirationDate?: string | null;
|
||||
}
|
||||
|
||||
export interface UpdateMemoryOptions {
|
||||
/** New content to update the memory with. */
|
||||
text?: string;
|
||||
/**
|
||||
* New content to update the memory with.
|
||||
* @deprecated Use `text` instead. Will be removed in the next major release.
|
||||
*/
|
||||
data?: string;
|
||||
/** Metadata merged into the memory's existing metadata. */
|
||||
metadata?: Record<string, any>;
|
||||
/** Date (YYYY-MM-DD) after which the memory expires, or `null` to clear it. */
|
||||
expirationDate?: string | null;
|
||||
}
|
||||
|
||||
export interface SearchMemoryOptions {
|
||||
@@ -20,11 +36,20 @@ export interface SearchMemoryOptions {
|
||||
threshold?: number;
|
||||
explain?: boolean;
|
||||
referenceDate?: number | string | Date | null;
|
||||
/**
|
||||
* Re-rank the results with the configured reranker before returning. No-op
|
||||
* when no `reranker` is configured on the Memory.
|
||||
*/
|
||||
rerank?: boolean;
|
||||
/** Include expired memories in the results. Defaults to false. */
|
||||
showExpired?: boolean;
|
||||
}
|
||||
|
||||
export interface GetAllMemoryOptions {
|
||||
topK?: number;
|
||||
filters?: SearchFilters;
|
||||
/** Include expired memories in the results. Defaults to false. */
|
||||
showExpired?: boolean;
|
||||
}
|
||||
|
||||
export interface DeleteAllMemoryOptions extends Entity {}
|
||||
|
||||
@@ -0,0 +1,21 @@
|
||||
export interface RerankResult {
|
||||
/** Index into the input `documents` array. */
|
||||
index: number;
|
||||
/** Relevance of the document to the query, 0..1, higher = more relevant. */
|
||||
rerankScore: number;
|
||||
}
|
||||
|
||||
export interface Reranker {
|
||||
/**
|
||||
* Rank `documents` by relevance to `query`.
|
||||
*
|
||||
* Returns results sorted by descending relevance. When `topK` is given, at
|
||||
* most that many results are returned. Each result's `index` points back into
|
||||
* the input `documents` array so callers can recover the original item.
|
||||
*/
|
||||
rerank(
|
||||
query: string,
|
||||
documents: string[],
|
||||
topK?: number,
|
||||
): Promise<RerankResult[]>;
|
||||
}
|
||||
@@ -0,0 +1,140 @@
|
||||
const mockRerank = jest.fn();
|
||||
|
||||
jest.mock("cohere-ai", () => ({
|
||||
CohereClient: jest.fn().mockImplementation(() => ({
|
||||
rerank: mockRerank,
|
||||
})),
|
||||
}));
|
||||
|
||||
import { CohereClient } from "cohere-ai";
|
||||
import { CohereReranker } from "./cohere";
|
||||
|
||||
describe("CohereReranker", () => {
|
||||
beforeEach(() => {
|
||||
mockRerank.mockReset();
|
||||
(CohereClient as unknown as jest.Mock).mockClear();
|
||||
});
|
||||
|
||||
it("throws when no API key is provided or configured", () => {
|
||||
const originalEnv = process.env.COHERE_API_KEY;
|
||||
delete process.env.COHERE_API_KEY;
|
||||
|
||||
expect(() => new CohereReranker({})).toThrow(/Cohere API key is required/);
|
||||
|
||||
if (originalEnv !== undefined) process.env.COHERE_API_KEY = originalEnv;
|
||||
});
|
||||
|
||||
it("sends the query, documents, topN, and default model to Cohere", async () => {
|
||||
mockRerank.mockResolvedValue({ results: [] });
|
||||
const reranker = new CohereReranker({ apiKey: "key" });
|
||||
|
||||
await reranker.rerank("capital of US?", ["a", "b", "c"], 2);
|
||||
|
||||
expect(mockRerank).toHaveBeenCalledWith({
|
||||
model: "rerank-v3.5",
|
||||
query: "capital of US?",
|
||||
documents: ["a", "b", "c"],
|
||||
topN: 2,
|
||||
returnDocuments: false,
|
||||
maxChunksPerDoc: undefined,
|
||||
});
|
||||
});
|
||||
|
||||
it("defaults topN to documents.length when neither the call nor config sets a top_k", async () => {
|
||||
mockRerank.mockResolvedValue({ results: [] });
|
||||
const reranker = new CohereReranker({ apiKey: "key" });
|
||||
|
||||
await reranker.rerank("q", ["a", "b", "c"]);
|
||||
|
||||
expect(mockRerank).toHaveBeenCalledWith(
|
||||
expect.objectContaining({ topN: 3 }),
|
||||
);
|
||||
});
|
||||
|
||||
it("forwards returnDocuments and maxChunksPerDoc from config", async () => {
|
||||
mockRerank.mockResolvedValue({ results: [] });
|
||||
const reranker = new CohereReranker({
|
||||
apiKey: "key",
|
||||
returnDocuments: true,
|
||||
maxChunksPerDoc: 5,
|
||||
});
|
||||
|
||||
await reranker.rerank("q", ["a"]);
|
||||
|
||||
expect(mockRerank).toHaveBeenCalledWith(
|
||||
expect.objectContaining({ returnDocuments: true, maxChunksPerDoc: 5 }),
|
||||
);
|
||||
});
|
||||
|
||||
it("returns Cohere's ranked results as {index, rerankScore}", async () => {
|
||||
mockRerank.mockResolvedValue({
|
||||
results: [
|
||||
{ index: 2, relevanceScore: 0.9 },
|
||||
{ index: 0, relevanceScore: 0.31 },
|
||||
],
|
||||
});
|
||||
const reranker = new CohereReranker({ apiKey: "key" });
|
||||
|
||||
const results = await reranker.rerank("q", ["x", "y", "z"]);
|
||||
|
||||
expect(results).toEqual([
|
||||
{ index: 2, rerankScore: 0.9 },
|
||||
{ index: 0, rerankScore: 0.31 },
|
||||
]);
|
||||
});
|
||||
|
||||
it("uses a custom model when provided", async () => {
|
||||
mockRerank.mockResolvedValue({ results: [] });
|
||||
const reranker = new CohereReranker({
|
||||
apiKey: "key",
|
||||
model: "rerank-v4.0-pro",
|
||||
});
|
||||
|
||||
await reranker.rerank("q", ["a"]);
|
||||
|
||||
expect(mockRerank).toHaveBeenCalledWith(
|
||||
expect.objectContaining({ model: "rerank-v4.0-pro" }),
|
||||
);
|
||||
});
|
||||
|
||||
it("returns an empty array without calling Cohere when there are no documents", async () => {
|
||||
const reranker = new CohereReranker({ apiKey: "key" });
|
||||
|
||||
const results = await reranker.rerank("q", []);
|
||||
|
||||
expect(results).toEqual([]);
|
||||
expect(mockRerank).not.toHaveBeenCalled();
|
||||
});
|
||||
|
||||
it("falls back to the original order with rerankScore 0.0 when the Cohere API call fails", async () => {
|
||||
mockRerank.mockRejectedValue(new Error("cohere is down"));
|
||||
const warnSpy = jest.spyOn(console, "warn").mockImplementation(() => {});
|
||||
const reranker = new CohereReranker({ apiKey: "key" });
|
||||
|
||||
const results = await reranker.rerank("q", ["a", "b", "c"]);
|
||||
|
||||
expect(results).toEqual([
|
||||
{ index: 0, rerankScore: 0.0 },
|
||||
{ index: 1, rerankScore: 0.0 },
|
||||
{ index: 2, rerankScore: 0.0 },
|
||||
]);
|
||||
expect(warnSpy).toHaveBeenCalled();
|
||||
|
||||
warnSpy.mockRestore();
|
||||
});
|
||||
|
||||
it("slices the fallback results by topK when the Cohere API call fails", async () => {
|
||||
mockRerank.mockRejectedValue(new Error("cohere is down"));
|
||||
jest.spyOn(console, "warn").mockImplementation(() => {});
|
||||
const reranker = new CohereReranker({ apiKey: "key", topK: 2 });
|
||||
|
||||
const results = await reranker.rerank("q", ["a", "b", "c"]);
|
||||
|
||||
expect(results).toEqual([
|
||||
{ index: 0, rerankScore: 0.0 },
|
||||
{ index: 1, rerankScore: 0.0 },
|
||||
]);
|
||||
|
||||
(console.warn as jest.Mock).mockRestore();
|
||||
});
|
||||
});
|
||||
@@ -0,0 +1,89 @@
|
||||
import { RerankerConfig } from "../types";
|
||||
import { Reranker, RerankResult } from "./base";
|
||||
|
||||
const DEFAULT_MODEL = "rerank-v3.5";
|
||||
|
||||
export class CohereReranker implements Reranker {
|
||||
private clientInstance?: any;
|
||||
private clientPromise?: Promise<any>;
|
||||
private readonly apiKey: string;
|
||||
private model: string;
|
||||
private topK?: number;
|
||||
private returnDocuments: boolean;
|
||||
private maxChunksPerDoc?: number;
|
||||
|
||||
constructor(config: RerankerConfig) {
|
||||
const apiKey = config.apiKey || process.env.COHERE_API_KEY;
|
||||
if (!apiKey) {
|
||||
throw new Error(
|
||||
"Cohere API key is required. Set COHERE_API_KEY environment variable or pass apiKey in config.",
|
||||
);
|
||||
}
|
||||
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[],
|
||||
topK?: number,
|
||||
): Promise<RerankResult[]> {
|
||||
if (documents.length === 0) return [];
|
||||
|
||||
try {
|
||||
const client = await this.getClient();
|
||||
const response = await client.rerank({
|
||||
model: this.model,
|
||||
query,
|
||||
documents,
|
||||
topN: topK || this.topK || documents.length,
|
||||
returnDocuments: this.returnDocuments,
|
||||
maxChunksPerDoc: this.maxChunksPerDoc,
|
||||
});
|
||||
|
||||
return response.results.map((result: any) => ({
|
||||
index: result.index,
|
||||
rerankScore: result.relevanceScore,
|
||||
}));
|
||||
} catch (e) {
|
||||
console.warn(
|
||||
`Cohere reranking failed, falling back to original order: ${e}`,
|
||||
);
|
||||
const scored = documents.map((_, index) => ({
|
||||
index,
|
||||
rerankScore: 0.0,
|
||||
}));
|
||||
const finalTopK = topK || this.topK;
|
||||
return finalTopK ? scored.slice(0, finalTopK) : scored;
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,189 @@
|
||||
const mockModelFromPretrained = jest.fn();
|
||||
const mockTokenizerFromPretrained = jest.fn();
|
||||
|
||||
jest.mock("@huggingface/transformers", () => ({
|
||||
AutoModelForSequenceClassification: {
|
||||
from_pretrained: mockModelFromPretrained,
|
||||
},
|
||||
AutoTokenizer: { from_pretrained: mockTokenizerFromPretrained },
|
||||
}));
|
||||
|
||||
import { CrossEncoderReranker } from "./cross_encoder";
|
||||
|
||||
const sigmoid = (x: number) => 1 / (1 + Math.exp(-x));
|
||||
|
||||
/** Wire the mocked tokenizer + model so the model returns `logits` for a call. */
|
||||
function setupModel(logits: number[][]) {
|
||||
const tokenizer = jest.fn().mockReturnValue({ input_ids: [] });
|
||||
mockTokenizerFromPretrained.mockResolvedValue(tokenizer);
|
||||
const model = jest
|
||||
.fn()
|
||||
.mockResolvedValue({ logits: { tolist: () => logits } });
|
||||
mockModelFromPretrained.mockResolvedValue(model);
|
||||
return { tokenizer, model };
|
||||
}
|
||||
|
||||
describe("CrossEncoderReranker", () => {
|
||||
beforeEach(() => {
|
||||
mockModelFromPretrained.mockReset();
|
||||
mockTokenizerFromPretrained.mockReset();
|
||||
});
|
||||
|
||||
it("scores each document and returns them sorted by relevance, sigmoid-normalized to [0,1]", async () => {
|
||||
setupModel([[0.0], [2.0], [-1.0]]);
|
||||
const reranker = new CrossEncoderReranker({}, "default-model");
|
||||
|
||||
const results = await reranker.rerank("q", ["a", "b", "c"]);
|
||||
|
||||
// sigmoid: b(2.0)=0.88 > a(0.0)=0.5 > c(-1.0)=0.27
|
||||
expect(results.map((r) => r.index)).toEqual([1, 0, 2]);
|
||||
expect(results[0].rerankScore).toBeCloseTo(sigmoid(2.0), 5);
|
||||
expect(results[1].rerankScore).toBeCloseTo(sigmoid(0.0), 5);
|
||||
expect(results[2].rerankScore).toBeCloseTo(sigmoid(-1.0), 5);
|
||||
});
|
||||
|
||||
it("pairs the query with each document via text_pair when tokenizing", async () => {
|
||||
const { tokenizer } = setupModel([[0.1], [0.2]]);
|
||||
const reranker = new CrossEncoderReranker(
|
||||
{ maxLength: 128 },
|
||||
"default-model",
|
||||
);
|
||||
|
||||
await reranker.rerank("what is x", ["doc one", "doc two"]);
|
||||
|
||||
expect(tokenizer).toHaveBeenCalledWith(
|
||||
["what is x", "what is x"],
|
||||
expect.objectContaining({
|
||||
text_pair: ["doc one", "doc two"],
|
||||
padding: true,
|
||||
truncation: true,
|
||||
max_length: 128,
|
||||
}),
|
||||
);
|
||||
});
|
||||
|
||||
it("applies the topK limit", async () => {
|
||||
setupModel([[0.0], [2.0], [-1.0]]);
|
||||
const reranker = new CrossEncoderReranker({}, "default-model");
|
||||
|
||||
const results = await reranker.rerank("q", ["a", "b", "c"], 2);
|
||||
|
||||
expect(results).toHaveLength(2);
|
||||
expect(results.map((r) => r.index)).toEqual([1, 0]);
|
||||
});
|
||||
|
||||
it("falls back to config.topK when the rerank() call omits one", async () => {
|
||||
setupModel([[0.0], [2.0], [-1.0]]);
|
||||
const reranker = new CrossEncoderReranker({ topK: 1 }, "default-model");
|
||||
|
||||
const results = await reranker.rerank("q", ["a", "b", "c"]);
|
||||
|
||||
expect(results).toHaveLength(1);
|
||||
expect(results.map((r) => r.index)).toEqual([1]);
|
||||
});
|
||||
|
||||
it("returns [] without loading the model when there are no documents", async () => {
|
||||
const reranker = new CrossEncoderReranker({}, "default-model");
|
||||
|
||||
const results = await reranker.rerank("q", []);
|
||||
|
||||
expect(results).toEqual([]);
|
||||
expect(mockModelFromPretrained).not.toHaveBeenCalled();
|
||||
expect(mockTokenizerFromPretrained).not.toHaveBeenCalled();
|
||||
});
|
||||
|
||||
it("returns raw logits as scores when normalize is false", async () => {
|
||||
setupModel([[2.0], [0.0]]);
|
||||
const reranker = new CrossEncoderReranker(
|
||||
{ normalize: false },
|
||||
"default-model",
|
||||
);
|
||||
|
||||
const results = await reranker.rerank("q", ["a", "b"]);
|
||||
|
||||
expect(results.map((r) => r.index)).toEqual([0, 1]);
|
||||
expect(results[0].rerankScore).toBe(2.0);
|
||||
expect(results[1].rerankScore).toBe(0.0);
|
||||
});
|
||||
|
||||
it("loads the model and tokenizer only once across multiple rerank calls", async () => {
|
||||
setupModel([[0.5]]);
|
||||
const reranker = new CrossEncoderReranker({}, "default-model");
|
||||
|
||||
await reranker.rerank("q", ["a"]);
|
||||
await reranker.rerank("q2", ["b"]);
|
||||
|
||||
expect(mockModelFromPretrained).toHaveBeenCalledTimes(1);
|
||||
expect(mockTokenizerFromPretrained).toHaveBeenCalledTimes(1);
|
||||
});
|
||||
|
||||
it("loads the default model, or the configured model when provided", async () => {
|
||||
setupModel([[0.5]]);
|
||||
|
||||
await new CrossEncoderReranker({}, "the-default").rerank("q", ["a"]);
|
||||
expect(mockModelFromPretrained).toHaveBeenCalledWith(
|
||||
"the-default",
|
||||
expect.any(Object),
|
||||
);
|
||||
|
||||
mockModelFromPretrained.mockClear();
|
||||
setupModel([[0.5]]);
|
||||
await new CrossEncoderReranker(
|
||||
{ model: "custom/model" },
|
||||
"the-default",
|
||||
).rerank("q", ["a"]);
|
||||
expect(mockModelFromPretrained).toHaveBeenCalledWith(
|
||||
"custom/model",
|
||||
expect.any(Object),
|
||||
);
|
||||
});
|
||||
|
||||
it("applies a default maxLength (as the huggingface provider passes 512) when config omits one", async () => {
|
||||
const { tokenizer } = setupModel([[0.5]]);
|
||||
const reranker = new CrossEncoderReranker({}, "default-model", 512);
|
||||
|
||||
await reranker.rerank("q", ["a"]);
|
||||
|
||||
expect(tokenizer).toHaveBeenCalledWith(
|
||||
["q"],
|
||||
expect.objectContaining({ max_length: 512 }),
|
||||
);
|
||||
});
|
||||
|
||||
it("falls back to the original order with rerankScore 0.0 when the model fails to load", async () => {
|
||||
mockModelFromPretrained.mockResolvedValue(jest.fn());
|
||||
mockTokenizerFromPretrained.mockRejectedValue(
|
||||
new Error("model download failed"),
|
||||
);
|
||||
const warnSpy = jest.spyOn(console, "warn").mockImplementation(() => {});
|
||||
const reranker = new CrossEncoderReranker({}, "default-model");
|
||||
|
||||
const results = await reranker.rerank("q", ["a", "b", "c"]);
|
||||
|
||||
expect(results).toEqual([
|
||||
{ index: 0, rerankScore: 0.0 },
|
||||
{ index: 1, rerankScore: 0.0 },
|
||||
{ index: 2, rerankScore: 0.0 },
|
||||
]);
|
||||
expect(warnSpy).toHaveBeenCalled();
|
||||
|
||||
warnSpy.mockRestore();
|
||||
});
|
||||
|
||||
it("falls back to the original order with rerankScore 0.0, sliced by topK, when scoring fails", async () => {
|
||||
const tokenizer = jest.fn().mockReturnValue({ input_ids: [] });
|
||||
mockTokenizerFromPretrained.mockResolvedValue(tokenizer);
|
||||
mockModelFromPretrained.mockResolvedValue(
|
||||
jest.fn().mockRejectedValue(new Error("forward pass failed")),
|
||||
);
|
||||
const warnSpy = jest.spyOn(console, "warn").mockImplementation(() => {});
|
||||
const reranker = new CrossEncoderReranker({}, "default-model");
|
||||
|
||||
const results = await reranker.rerank("q", ["a", "b"], 1);
|
||||
|
||||
expect(results).toEqual([{ index: 0, rerankScore: 0.0 }]);
|
||||
expect(warnSpy).toHaveBeenCalled();
|
||||
|
||||
warnSpy.mockRestore();
|
||||
});
|
||||
});
|
||||
@@ -0,0 +1,98 @@
|
||||
import { RerankerConfig } from "../types";
|
||||
import { Reranker, RerankResult } from "./base";
|
||||
|
||||
const sigmoid = (x: number) => 1 / (1 + Math.exp(-x));
|
||||
|
||||
export class CrossEncoderReranker implements Reranker {
|
||||
private modelId: string;
|
||||
private device?: string;
|
||||
private maxLength?: number;
|
||||
private normalize: boolean;
|
||||
private topK?: number;
|
||||
// ponytail: batchSize/showProgressBar are accepted for config parity with the
|
||||
// Python SDK but are no-ops here — a memory search reranks a small candidate
|
||||
// set in a single forward pass. Chunk by batchSize if that ever grows.
|
||||
private loaded?: Promise<{ model: any; tokenizer: any }>;
|
||||
|
||||
constructor(
|
||||
config: RerankerConfig,
|
||||
defaultModel: string,
|
||||
defaultMaxLength?: number,
|
||||
) {
|
||||
this.modelId = config.model || defaultModel;
|
||||
this.device = config.device;
|
||||
this.maxLength = config.maxLength ?? defaultMaxLength;
|
||||
this.normalize = config.normalize ?? true;
|
||||
this.topK = config.topK;
|
||||
}
|
||||
|
||||
private load() {
|
||||
if (!this.loaded) {
|
||||
this.loaded = (async () => {
|
||||
// Lazy-load Transformers.js (and its onnxruntime native binding) only
|
||||
// when a rerank actually runs. A static import would pull onnxruntime
|
||||
// into every `new Memory()`, colliding on Linux with fastembed's
|
||||
// separate onnxruntime version — see the merge with the FastEmbed
|
||||
// embedder. Deferring it keeps memory construction free of ONNX.
|
||||
const { AutoModelForSequenceClassification, AutoTokenizer } =
|
||||
await import("@huggingface/transformers");
|
||||
const options: any = {};
|
||||
if (this.device) options.device = this.device;
|
||||
const model = await AutoModelForSequenceClassification.from_pretrained(
|
||||
this.modelId,
|
||||
options,
|
||||
);
|
||||
const tokenizer = await AutoTokenizer.from_pretrained(this.modelId);
|
||||
return { model, tokenizer };
|
||||
})();
|
||||
}
|
||||
return this.loaded;
|
||||
}
|
||||
|
||||
async rerank(
|
||||
query: string,
|
||||
documents: string[],
|
||||
topK?: number,
|
||||
): Promise<RerankResult[]> {
|
||||
if (documents.length === 0) return [];
|
||||
|
||||
try {
|
||||
const { model, tokenizer } = await this.load();
|
||||
|
||||
const inputs = tokenizer(
|
||||
documents.map(() => query),
|
||||
{
|
||||
text_pair: documents,
|
||||
padding: true,
|
||||
truncation: true,
|
||||
...(this.maxLength ? { max_length: this.maxLength } : {}),
|
||||
},
|
||||
);
|
||||
|
||||
const { logits } = await model(inputs);
|
||||
const rows: unknown[] = logits.tolist();
|
||||
|
||||
const scored = rows.map((row, index) => {
|
||||
const logit = Array.isArray(row) ? (row[0] as number) : (row as number);
|
||||
return {
|
||||
index,
|
||||
rerankScore: this.normalize ? sigmoid(logit) : logit,
|
||||
};
|
||||
});
|
||||
|
||||
scored.sort((a, b) => b.rerankScore - a.rerankScore);
|
||||
const finalTopK = topK || this.topK;
|
||||
return finalTopK ? scored.slice(0, finalTopK) : scored;
|
||||
} catch (e) {
|
||||
console.warn(
|
||||
`Cross-encoder reranking failed, falling back to original order: ${e}`,
|
||||
);
|
||||
const scored = documents.map((_, index) => ({
|
||||
index,
|
||||
rerankScore: 0.0,
|
||||
}));
|
||||
const finalTopK = topK || this.topK;
|
||||
return finalTopK ? scored.slice(0, finalTopK) : scored;
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,174 @@
|
||||
import { LLM } from "../llms/base";
|
||||
import { LLMReranker } from "./llm";
|
||||
|
||||
// Duplicated rather than imported from ./llm.ts so this test catches drift.
|
||||
const EXPECTED_SYSTEM_PROMPT = `You are a relevance scoring assistant. Given a query and a document, score how relevant the document is to the query.
|
||||
|
||||
Score the relevance on a scale from 0.0 to 1.0, where:
|
||||
- 1.0 = Perfectly relevant and directly answers the query
|
||||
- 0.8-0.9 = Highly relevant with good information
|
||||
- 0.6-0.7 = Moderately relevant with some useful information
|
||||
- 0.4-0.5 = Slightly relevant with limited useful information
|
||||
- 0.0-0.3 = Not relevant or no useful information
|
||||
|
||||
Respond with only a single numerical score between 0.0 and 1.0. Do not include any explanation or additional text.`;
|
||||
|
||||
/**
|
||||
* Fake LLM that scores a document by looking up the document text inside the
|
||||
* prompt. Test document tokens must be distinct and must not be substrings of
|
||||
* the prompt boilerplate (e.g. avoid "a"/"b"), or the lookup resolves the
|
||||
* wrong doc.
|
||||
*/
|
||||
function makeLLM(scoreByDoc: Record<string, string>): LLM {
|
||||
return {
|
||||
generateResponse: async (
|
||||
messages: Array<{ role: string; content: string }>,
|
||||
) => {
|
||||
const prompt = messages.map((m) => m.content).join("\n");
|
||||
const doc = Object.keys(scoreByDoc).find((d) => prompt.includes(d));
|
||||
return doc ? scoreByDoc[doc] : "no number here";
|
||||
},
|
||||
generateChat: async () => ({ content: "", role: "assistant" }),
|
||||
};
|
||||
}
|
||||
|
||||
describe("LLMReranker", () => {
|
||||
it("sorts documents by descending relevance score", async () => {
|
||||
const llm = makeLLM({ cats: "0.2", dogs: "0.9", fish: "0.5" });
|
||||
const reranker = new LLMReranker({}, llm);
|
||||
|
||||
const results = await reranker.rerank("pets", ["cats", "dogs", "fish"]);
|
||||
|
||||
expect(results.map((r) => r.index)).toEqual([1, 2, 0]);
|
||||
expect(results.map((r) => r.rerankScore)).toEqual([0.9, 0.5, 0.2]);
|
||||
});
|
||||
|
||||
it("clamps scores to the [0, 1] range", async () => {
|
||||
const llm = makeLLM({ zebra: "1.5", walrus: "-0.3" });
|
||||
const reranker = new LLMReranker({}, llm);
|
||||
|
||||
const results = await reranker.rerank("q", ["zebra", "walrus"]);
|
||||
|
||||
const byIndex = new Map(results.map((r) => [r.index, r.rerankScore]));
|
||||
expect(byIndex.get(0)).toBe(1); // "zebra" 1.5 -> clamped to 1
|
||||
expect(byIndex.get(1)).toBe(0); // "walrus" -0.3 -> clamped to 0
|
||||
});
|
||||
|
||||
it("truncates results to topK", async () => {
|
||||
const llm = makeLLM({ alpha: "0.1", bravo: "0.8", charlie: "0.5" });
|
||||
const reranker = new LLMReranker({}, llm);
|
||||
|
||||
const results = await reranker.rerank(
|
||||
"q",
|
||||
["alpha", "bravo", "charlie"],
|
||||
2,
|
||||
);
|
||||
|
||||
expect(results).toHaveLength(2);
|
||||
expect(results.map((r) => r.index)).toEqual([1, 2]); // bravo(0.8), charlie(0.5)
|
||||
});
|
||||
|
||||
it("falls back to config.topK when the rerank() call omits one", async () => {
|
||||
const llm = makeLLM({ alpha: "0.1", bravo: "0.8", charlie: "0.5" });
|
||||
const reranker = new LLMReranker({ topK: 1 }, llm);
|
||||
|
||||
const results = await reranker.rerank("q", ["alpha", "bravo", "charlie"]);
|
||||
|
||||
expect(results).toHaveLength(1);
|
||||
expect(results[0].index).toBe(1); // bravo(0.8)
|
||||
});
|
||||
|
||||
it("falls back to a neutral score of 0.5 (not 0) when the LLM output has no number", async () => {
|
||||
const llm = makeLLM({ junk: "I cannot rate this" });
|
||||
const reranker = new LLMReranker({}, llm);
|
||||
|
||||
const results = await reranker.rerank("q", ["junk"]);
|
||||
|
||||
expect(results[0].rerankScore).toBe(0.5);
|
||||
});
|
||||
|
||||
it("prefers a decimal match over an integer match when extracting the score", async () => {
|
||||
const llm = makeLLM({ item: "The score is 0.73 out of 1" });
|
||||
const reranker = new LLMReranker({}, llm);
|
||||
|
||||
const results = await reranker.rerank("q", ["item"]);
|
||||
|
||||
expect(results[0].rerankScore).toBe(0.73);
|
||||
});
|
||||
|
||||
it("falls back to an integer match when no decimal is present", async () => {
|
||||
const llm = makeLLM({ item: "I'd say this is a solid 1" });
|
||||
const reranker = new LLMReranker({}, llm);
|
||||
|
||||
const results = await reranker.rerank("q", ["item"]);
|
||||
|
||||
expect(results[0].rerankScore).toBe(1);
|
||||
});
|
||||
|
||||
it("assigns a neutral 0.5 score (not 0.0) when a per-document LLM call fails, and still returns that document", async () => {
|
||||
const llm: LLM = {
|
||||
generateResponse: jest
|
||||
.fn()
|
||||
.mockResolvedValueOnce("0.9") // scores "good"
|
||||
.mockRejectedValueOnce(new Error("rate limited")), // scores "bad"
|
||||
generateChat: async () => ({ content: "", role: "assistant" }),
|
||||
};
|
||||
const warnSpy = jest.spyOn(console, "warn").mockImplementation(() => {});
|
||||
const reranker = new LLMReranker({}, llm);
|
||||
|
||||
const results = await reranker.rerank("q", ["good", "bad"]);
|
||||
|
||||
expect(results).toHaveLength(2);
|
||||
const byIndex = new Map(results.map((r) => [r.index, r.rerankScore]));
|
||||
expect(byIndex.get(0)).toBe(0.9);
|
||||
expect(byIndex.get(1)).toBe(0.5);
|
||||
expect(warnSpy).toHaveBeenCalled();
|
||||
|
||||
warnSpy.mockRestore();
|
||||
});
|
||||
|
||||
it("sends the exact system prompt and a separate user message with the query and document", async () => {
|
||||
const generateResponse = jest.fn().mockResolvedValue("0.5");
|
||||
const llm: LLM = {
|
||||
generateResponse,
|
||||
generateChat: async () => ({ content: "", role: "assistant" }),
|
||||
};
|
||||
const reranker = new LLMReranker({}, llm);
|
||||
|
||||
await reranker.rerank("what is the capital?", [
|
||||
"Paris is the capital of France.",
|
||||
]);
|
||||
|
||||
expect(generateResponse).toHaveBeenCalledWith([
|
||||
{ role: "system", content: EXPECTED_SYSTEM_PROMPT },
|
||||
{
|
||||
role: "user",
|
||||
content:
|
||||
"Query: what is the capital?\n\nDocument: Paris is the capital of France.",
|
||||
},
|
||||
]);
|
||||
});
|
||||
|
||||
it("truncates the query and document to 4000 characters before sending", async () => {
|
||||
const generateResponse = jest.fn().mockResolvedValue("0.5");
|
||||
const llm: LLM = {
|
||||
generateResponse,
|
||||
generateChat: async () => ({ content: "", role: "assistant" }),
|
||||
};
|
||||
const reranker = new LLMReranker({}, llm);
|
||||
const longQuery = "q".repeat(5000);
|
||||
const longDoc = "d".repeat(5000);
|
||||
|
||||
await reranker.rerank(longQuery, [longDoc]);
|
||||
|
||||
const userMessage = generateResponse.mock.calls[0][0][1];
|
||||
const sentQuery = userMessage.content.match(/^Query: (q+)/)[1];
|
||||
const sentDoc = userMessage.content.match(/Document: (d+)/)[1];
|
||||
expect(sentQuery).toHaveLength(4000);
|
||||
expect(sentDoc).toHaveLength(4000);
|
||||
});
|
||||
|
||||
it("throws when no LLM is provided", () => {
|
||||
expect(() => new LLMReranker({}, undefined as unknown as LLM)).toThrow();
|
||||
});
|
||||
});
|
||||
@@ -0,0 +1,87 @@
|
||||
import { RerankerConfig } from "../types";
|
||||
import { LLM, LLMResponse } from "../llms/base";
|
||||
import { Reranker, RerankResult } from "./base";
|
||||
|
||||
const SYSTEM_PROMPT = `You are a relevance scoring assistant. Given a query and a document, score how relevant the document is to the query.
|
||||
|
||||
Score the relevance on a scale from 0.0 to 1.0, where:
|
||||
- 1.0 = Perfectly relevant and directly answers the query
|
||||
- 0.8-0.9 = Highly relevant with good information
|
||||
- 0.6-0.7 = Moderately relevant with some useful information
|
||||
- 0.4-0.5 = Slightly relevant with limited useful information
|
||||
- 0.0-0.3 = Not relevant or no useful information
|
||||
|
||||
Respond with only a single numerical score between 0.0 and 1.0. Do not include any explanation or additional text.`;
|
||||
|
||||
const MAX_INPUT_LEN = 4000;
|
||||
|
||||
export class LLMReranker implements Reranker {
|
||||
private llm: LLM;
|
||||
private topK?: number;
|
||||
|
||||
constructor(config: RerankerConfig, llm: LLM) {
|
||||
if (!llm) {
|
||||
throw new Error(
|
||||
"LLMReranker requires an LLM instance; RerankerFactory should always provide one for the llm_reranker provider.",
|
||||
);
|
||||
}
|
||||
this.llm = llm;
|
||||
this.topK = config.topK;
|
||||
}
|
||||
|
||||
async rerank(
|
||||
query: string,
|
||||
documents: string[],
|
||||
topK?: number,
|
||||
): Promise<RerankResult[]> {
|
||||
if (documents.length === 0) return [];
|
||||
|
||||
const scored = await Promise.all(
|
||||
documents.map(async (document, index) => {
|
||||
try {
|
||||
const rerankScore = await this.score(query, document);
|
||||
return { index, rerankScore };
|
||||
} catch (e) {
|
||||
console.warn(
|
||||
`LLM reranking failed for a document, assigning neutral score: ${e}`,
|
||||
);
|
||||
return { index, rerankScore: 0.5 };
|
||||
}
|
||||
}),
|
||||
);
|
||||
|
||||
scored.sort((a, b) => b.rerankScore - a.rerankScore);
|
||||
const finalTopK = topK || this.topK;
|
||||
return finalTopK ? scored.slice(0, finalTopK) : scored;
|
||||
}
|
||||
|
||||
private async score(query: string, document: string): Promise<number> {
|
||||
const safeQuery = query.slice(0, MAX_INPUT_LEN);
|
||||
const safeDoc = document.slice(0, MAX_INPUT_LEN);
|
||||
const userMessage = `Query: ${safeQuery}\n\nDocument: ${safeDoc}`;
|
||||
|
||||
const response = await this.llm.generateResponse([
|
||||
{ role: "system", content: SYSTEM_PROMPT },
|
||||
{ role: "user", content: userMessage },
|
||||
]);
|
||||
|
||||
const text =
|
||||
typeof response === "string"
|
||||
? response
|
||||
: ((response as LLMResponse)?.content ?? "");
|
||||
|
||||
return this.extractScore(text);
|
||||
}
|
||||
|
||||
private extractScore(responseText: string): number {
|
||||
const matches =
|
||||
responseText.match(/-?\d+\.\d+/g) || responseText.match(/-?\d+/g);
|
||||
|
||||
if (matches && matches.length > 0) {
|
||||
const score = parseFloat(matches[0]);
|
||||
return Math.min(Math.max(score, 0.0), 1.0);
|
||||
}
|
||||
|
||||
return 0.5;
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,133 @@
|
||||
const mockRerank = jest.fn();
|
||||
|
||||
jest.mock("zeroentropy", () => ({
|
||||
ZeroEntropy: jest.fn().mockImplementation(() => ({
|
||||
models: { rerank: mockRerank },
|
||||
})),
|
||||
}));
|
||||
|
||||
import { ZeroEntropy } from "zeroentropy";
|
||||
import { ZeroEntropyReranker } from "./zeroentropy";
|
||||
|
||||
describe("ZeroEntropyReranker", () => {
|
||||
beforeEach(() => {
|
||||
mockRerank.mockReset();
|
||||
(ZeroEntropy as unknown as jest.Mock).mockClear();
|
||||
});
|
||||
|
||||
it("throws when no API key is provided or configured", () => {
|
||||
const originalEnv = process.env.ZERO_ENTROPY_API_KEY;
|
||||
delete process.env.ZERO_ENTROPY_API_KEY;
|
||||
|
||||
expect(() => new ZeroEntropyReranker({})).toThrow(
|
||||
/Zero Entropy API key is required/,
|
||||
);
|
||||
|
||||
if (originalEnv !== undefined)
|
||||
process.env.ZERO_ENTROPY_API_KEY = originalEnv;
|
||||
});
|
||||
|
||||
it("sends the query, documents, and default model to ZeroEntropy without a top_n parameter", async () => {
|
||||
mockRerank.mockResolvedValue({ results: [] });
|
||||
const reranker = new ZeroEntropyReranker({ apiKey: "key" });
|
||||
|
||||
await reranker.rerank("capital of US?", ["a", "b", "c"], 2);
|
||||
|
||||
expect(mockRerank).toHaveBeenCalledWith({
|
||||
model: "zerank-1",
|
||||
query: "capital of US?",
|
||||
documents: ["a", "b", "c"],
|
||||
});
|
||||
});
|
||||
|
||||
it("maps ZeroEntropy's results (relevance_score) to {index, rerankScore}", async () => {
|
||||
mockRerank.mockResolvedValue({
|
||||
results: [
|
||||
{ index: 2, relevance_score: 0.9 },
|
||||
{ index: 0, relevance_score: 0.31 },
|
||||
],
|
||||
});
|
||||
const reranker = new ZeroEntropyReranker({ apiKey: "key" });
|
||||
|
||||
const results = await reranker.rerank("q", ["x", "y", "z"]);
|
||||
|
||||
expect(results).toEqual([
|
||||
{ index: 2, rerankScore: 0.9 },
|
||||
{ index: 0, rerankScore: 0.31 },
|
||||
]);
|
||||
});
|
||||
|
||||
it("sorts unsorted API results by descending relevance score client-side", async () => {
|
||||
mockRerank.mockResolvedValue({
|
||||
results: [
|
||||
{ index: 0, relevance_score: 0.2 },
|
||||
{ index: 1, relevance_score: 0.9 },
|
||||
{ index: 2, relevance_score: 0.5 },
|
||||
],
|
||||
});
|
||||
const reranker = new ZeroEntropyReranker({ apiKey: "key" });
|
||||
|
||||
const results = await reranker.rerank("q", ["a", "b", "c"]);
|
||||
|
||||
expect(results.map((r) => r.index)).toEqual([1, 2, 0]);
|
||||
expect(results.map((r) => r.rerankScore)).toEqual([0.9, 0.5, 0.2]);
|
||||
});
|
||||
|
||||
it("slices to topK client-side after sorting", async () => {
|
||||
mockRerank.mockResolvedValue({
|
||||
results: [
|
||||
{ index: 0, relevance_score: 0.2 },
|
||||
{ index: 1, relevance_score: 0.9 },
|
||||
{ index: 2, relevance_score: 0.5 },
|
||||
],
|
||||
});
|
||||
const reranker = new ZeroEntropyReranker({ apiKey: "key" });
|
||||
|
||||
const results = await reranker.rerank("q", ["a", "b", "c"], 2);
|
||||
|
||||
expect(results).toEqual([
|
||||
{ index: 1, rerankScore: 0.9 },
|
||||
{ index: 2, rerankScore: 0.5 },
|
||||
]);
|
||||
});
|
||||
|
||||
it("uses a custom model when provided", async () => {
|
||||
mockRerank.mockResolvedValue({ results: [] });
|
||||
const reranker = new ZeroEntropyReranker({
|
||||
apiKey: "key",
|
||||
model: "zerank-1-small",
|
||||
});
|
||||
|
||||
await reranker.rerank("q", ["a"]);
|
||||
|
||||
expect(mockRerank).toHaveBeenCalledWith(
|
||||
expect.objectContaining({ model: "zerank-1-small" }),
|
||||
);
|
||||
});
|
||||
|
||||
it("returns an empty array without calling ZeroEntropy when there are no documents", async () => {
|
||||
const reranker = new ZeroEntropyReranker({ apiKey: "key" });
|
||||
|
||||
const results = await reranker.rerank("q", []);
|
||||
|
||||
expect(results).toEqual([]);
|
||||
expect(mockRerank).not.toHaveBeenCalled();
|
||||
});
|
||||
|
||||
it("falls back to the original order with rerankScore 0.0 when the API call fails", async () => {
|
||||
mockRerank.mockRejectedValue(new Error("zero entropy is down"));
|
||||
const warnSpy = jest.spyOn(console, "warn").mockImplementation(() => {});
|
||||
const reranker = new ZeroEntropyReranker({ apiKey: "key" });
|
||||
|
||||
const results = await reranker.rerank("q", ["a", "b", "c"]);
|
||||
|
||||
expect(results).toEqual([
|
||||
{ index: 0, rerankScore: 0.0 },
|
||||
{ index: 1, rerankScore: 0.0 },
|
||||
{ index: 2, rerankScore: 0.0 },
|
||||
]);
|
||||
expect(warnSpy).toHaveBeenCalled();
|
||||
|
||||
warnSpy.mockRestore();
|
||||
});
|
||||
});
|
||||
@@ -0,0 +1,86 @@
|
||||
import { RerankerConfig } from "../types";
|
||||
import { Reranker, RerankResult } from "./base";
|
||||
|
||||
const DEFAULT_MODEL = "zerank-1";
|
||||
|
||||
export class ZeroEntropyReranker implements Reranker {
|
||||
private clientInstance?: any;
|
||||
private clientPromise?: Promise<any>;
|
||||
private readonly apiKey: string;
|
||||
private model: string;
|
||||
private topK?: number;
|
||||
|
||||
constructor(config: RerankerConfig) {
|
||||
const apiKey = config.apiKey || process.env.ZERO_ENTROPY_API_KEY;
|
||||
if (!apiKey) {
|
||||
throw new Error(
|
||||
"Zero Entropy API key is required. Set ZERO_ENTROPY_API_KEY environment variable or pass apiKey in config.",
|
||||
);
|
||||
}
|
||||
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[],
|
||||
topK?: number,
|
||||
): Promise<RerankResult[]> {
|
||||
if (documents.length === 0) return [];
|
||||
|
||||
try {
|
||||
const client = await this.getClient();
|
||||
const response = await client.models.rerank({
|
||||
model: this.model,
|
||||
query,
|
||||
documents,
|
||||
});
|
||||
|
||||
const scored: RerankResult[] = response.results.map((result: any) => ({
|
||||
index: result.index,
|
||||
rerankScore: result.relevance_score,
|
||||
}));
|
||||
scored.sort((a, b) => b.rerankScore - a.rerankScore);
|
||||
|
||||
const finalTopK = topK || this.topK;
|
||||
return finalTopK ? scored.slice(0, finalTopK) : scored;
|
||||
} catch (e) {
|
||||
console.warn(
|
||||
`Zero Entropy reranking failed, falling back to original order: ${e}`,
|
||||
);
|
||||
const scored = documents.map((_, index) => ({
|
||||
index,
|
||||
rerankScore: 0.0,
|
||||
}));
|
||||
const finalTopK = topK || this.topK;
|
||||
return finalTopK ? scored.slice(0, finalTopK) : scored;
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,214 @@
|
||||
import { readFileSync } from "fs";
|
||||
import { join } from "path";
|
||||
|
||||
import { AWSBedrockLLM, extractProvider } from "../llms/aws_bedrock";
|
||||
|
||||
/**
|
||||
* Fake BedrockRuntimeClient capturing the Converse command input and returning
|
||||
* a scripted response — exercises request shaping + response parsing without
|
||||
* the AWS SDK or live credentials.
|
||||
*/
|
||||
class FakeBedrockClient {
|
||||
public lastInput: any = null;
|
||||
public response: any;
|
||||
constructor(response: any) {
|
||||
this.response = response;
|
||||
}
|
||||
async send(command: any) {
|
||||
// ConverseCommand stores its input on `.input` (mirrored by our fake command).
|
||||
this.lastInput = command?.input ?? command;
|
||||
return this.response;
|
||||
}
|
||||
}
|
||||
|
||||
// Counts real client constructions so we can prove the SDK stays untouched
|
||||
// until the first request. (jest only lets mock factories close over `mock*`.)
|
||||
let mockClientConstructions = 0;
|
||||
|
||||
// The provider loads the AWS SDK on first use via dynamic import(). Mock the
|
||||
// module so it resolves and ConverseCommand simply wraps its input.
|
||||
jest.mock(
|
||||
"@aws-sdk/client-bedrock-runtime",
|
||||
() => ({
|
||||
BedrockRuntimeClient: class {
|
||||
constructor(public config: any) {
|
||||
mockClientConstructions++;
|
||||
}
|
||||
},
|
||||
ConverseCommand: class {
|
||||
input: any;
|
||||
constructor(input: any) {
|
||||
this.input = input;
|
||||
}
|
||||
},
|
||||
}),
|
||||
{ virtual: true },
|
||||
);
|
||||
|
||||
function makeLLM(client: FakeBedrockClient, overrides: any = {}) {
|
||||
return new AWSBedrockLLM({
|
||||
model: "anthropic.claude-3-5-sonnet-20240620-v1:0",
|
||||
client,
|
||||
...overrides,
|
||||
});
|
||||
}
|
||||
|
||||
describe("extractProvider", () => {
|
||||
it("detects the model family from the Bedrock model id", () => {
|
||||
expect(extractProvider("anthropic.claude-3-5-sonnet-20240620-v1:0")).toBe(
|
||||
"anthropic",
|
||||
);
|
||||
expect(extractProvider("amazon.nova-pro-v1:0")).toBe("amazon");
|
||||
expect(extractProvider("meta.llama3-70b-instruct-v1:0")).toBe("meta");
|
||||
expect(extractProvider("mistral.mistral-large-2407-v1:0")).toBe("mistral");
|
||||
});
|
||||
|
||||
it("throws on an unknown provider", () => {
|
||||
expect(() => extractProvider("totally-unknown-model")).toThrow(
|
||||
/Unknown provider/,
|
||||
);
|
||||
});
|
||||
});
|
||||
|
||||
describe("AWSBedrockLLM", () => {
|
||||
const textResponse = {
|
||||
output: { message: { content: [{ text: "hello from bedrock" }] } },
|
||||
};
|
||||
|
||||
it("returns assistant text from a Converse response", async () => {
|
||||
const client = new FakeBedrockClient(textResponse);
|
||||
const llm = makeLLM(client);
|
||||
const out = await llm.generateResponse([{ role: "user", content: "hi" }]);
|
||||
expect(out).toBe("hello from bedrock");
|
||||
});
|
||||
|
||||
it("lifts system messages into the top-level system block", async () => {
|
||||
const client = new FakeBedrockClient(textResponse);
|
||||
const llm = makeLLM(client);
|
||||
await llm.generateResponse([
|
||||
{ role: "system", content: "be terse" },
|
||||
{ role: "user", content: "hi" },
|
||||
]);
|
||||
expect(client.lastInput.system).toEqual([{ text: "be terse" }]);
|
||||
expect(client.lastInput.messages).toEqual([
|
||||
{ role: "user", content: [{ text: "hi" }] },
|
||||
]);
|
||||
});
|
||||
|
||||
it("omits topP for anthropic models even when configured", async () => {
|
||||
const client = new FakeBedrockClient(textResponse);
|
||||
const llm = makeLLM(client, { topP: 0.9, temperature: 0.2 });
|
||||
await llm.generateResponse([{ role: "user", content: "hi" }]);
|
||||
expect(client.lastInput.inferenceConfig.topP).toBeUndefined();
|
||||
expect(client.lastInput.inferenceConfig.temperature).toBe(0.2);
|
||||
});
|
||||
|
||||
it("includes topP for non-anthropic models", async () => {
|
||||
const client = new FakeBedrockClient(textResponse);
|
||||
const llm = makeLLM(client, {
|
||||
model: "meta.llama3-70b-instruct-v1:0",
|
||||
topP: 0.9,
|
||||
});
|
||||
await llm.generateResponse([{ role: "user", content: "hi" }]);
|
||||
expect(client.lastInput.inferenceConfig.topP).toBe(0.9);
|
||||
});
|
||||
|
||||
it("converts OpenAI-style tools to Converse toolConfig and parses toolUse", async () => {
|
||||
const toolResponse = {
|
||||
output: {
|
||||
message: {
|
||||
content: [{ toolUse: { name: "add_memory", input: { text: "x" } } }],
|
||||
},
|
||||
},
|
||||
};
|
||||
const client = new FakeBedrockClient(toolResponse);
|
||||
const llm = makeLLM(client);
|
||||
const tools = [
|
||||
{
|
||||
type: "function",
|
||||
function: {
|
||||
name: "add_memory",
|
||||
description: "store a memory",
|
||||
parameters: { type: "object", properties: {} },
|
||||
},
|
||||
},
|
||||
];
|
||||
const out = await llm.generateResponse(
|
||||
[{ role: "user", content: "remember x" }],
|
||||
undefined,
|
||||
tools,
|
||||
);
|
||||
expect(client.lastInput.toolConfig.tools[0].toolSpec.name).toBe(
|
||||
"add_memory",
|
||||
);
|
||||
expect(typeof out).toBe("object");
|
||||
expect((out as any).toolCalls[0]).toEqual({
|
||||
name: "add_memory",
|
||||
arguments: JSON.stringify({ text: "x" }),
|
||||
});
|
||||
});
|
||||
|
||||
it("never sends an empty messages array", async () => {
|
||||
const client = new FakeBedrockClient(textResponse);
|
||||
const llm = makeLLM(client);
|
||||
await llm.generateResponse([{ role: "system", content: "only system" }]);
|
||||
expect(client.lastInput.messages).toEqual([
|
||||
{ role: "user", content: [{ text: "" }] },
|
||||
]);
|
||||
});
|
||||
|
||||
it("wraps SDK errors with a provider-tagged message", async () => {
|
||||
const client = {
|
||||
send: async () => {
|
||||
throw new Error("AccessDeniedException");
|
||||
},
|
||||
} as any;
|
||||
const llm = makeLLM(client);
|
||||
await expect(
|
||||
llm.generateResponse([{ role: "user", content: "hi" }]),
|
||||
).rejects.toThrow(/AWS Bedrock LLM failed: AccessDeniedException/);
|
||||
});
|
||||
|
||||
it("generateChat returns text content with assistant role", async () => {
|
||||
const client = new FakeBedrockClient(textResponse);
|
||||
const llm = makeLLM(client);
|
||||
const res = await llm.generateChat([{ role: "user", content: "hi" }]);
|
||||
expect(res).toEqual({ content: "hello from bedrock", role: "assistant" });
|
||||
});
|
||||
|
||||
it("does not construct the Bedrock client until the first request", () => {
|
||||
const before = mockClientConstructions;
|
||||
// No injected client: the old constructor eagerly built a real one.
|
||||
new AWSBedrockLLM({
|
||||
model: "anthropic.claude-3-5-sonnet-20240620-v1:0",
|
||||
awsRegion: "us-west-2",
|
||||
});
|
||||
expect(mockClientConstructions).toBe(before);
|
||||
});
|
||||
});
|
||||
|
||||
/**
|
||||
* ts-jest downlevels this suite to CommonJS, where `require()` works fine — so
|
||||
* no runtime test here can observe the ESM failure. Guard the source invariant
|
||||
* that causes it instead: esbuild rewrites `require()` in the published ESM
|
||||
* bundle (dist/oss/index.mjs) into a `__require` shim that throws
|
||||
* `Dynamic require of "..." is not supported`, stranding every ESM consumer.
|
||||
*/
|
||||
describe("aws_bedrock.ts ESM safety", () => {
|
||||
const source = readFileSync(
|
||||
join(__dirname, "..", "llms", "aws_bedrock.ts"),
|
||||
"utf8",
|
||||
);
|
||||
// Comments legitimately mention require(); only real code should be checked.
|
||||
const code = source
|
||||
.replace(/\/\*[\s\S]*?\*\//g, "")
|
||||
.replace(/(^|[^:])\/\/.*$/gm, "$1");
|
||||
|
||||
it("never calls require() — the ESM bundle turns it into a throwing shim", () => {
|
||||
expect(code).not.toMatch(/\brequire\s*\(/);
|
||||
});
|
||||
|
||||
it("loads the optional AWS SDK through a dynamic import()", () => {
|
||||
expect(code).toMatch(/import\("@aws-sdk\/client-bedrock-runtime"\)/);
|
||||
});
|
||||
});
|
||||
@@ -0,0 +1,443 @@
|
||||
// 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 collectionHandle = { add, query, get, update, upsert, delete: deleteFn };
|
||||
const getOrCreateCollection = jest.fn().mockResolvedValue(collectionHandle);
|
||||
const deleteCollection = jest.fn().mockResolvedValue(undefined);
|
||||
|
||||
const clientImpl = () => ({ getOrCreateCollection, deleteCollection });
|
||||
const ChromaClient = jest.fn().mockImplementation(clientImpl);
|
||||
const CloudClient = jest.fn().mockImplementation(clientImpl);
|
||||
|
||||
const __mocks__ = {
|
||||
add,
|
||||
query,
|
||||
get,
|
||||
update,
|
||||
upsert,
|
||||
deleteFn,
|
||||
getOrCreateCollection,
|
||||
deleteCollection,
|
||||
ChromaClient,
|
||||
CloudClient,
|
||||
};
|
||||
|
||||
jest.mock("chromadb", () => ({
|
||||
ChromaClient: __mocks__.ChromaClient,
|
||||
CloudClient: __mocks__.CloudClient,
|
||||
}));
|
||||
|
||||
import { ChromaDB } from "../vector_stores/chroma";
|
||||
import { VectorStoreFactory } from "../utils/factory";
|
||||
|
||||
// --- Helpers ---
|
||||
|
||||
function makeDb(overrides: Record<string, any> = {}): ChromaDB {
|
||||
return new ChromaDB({
|
||||
collectionName: "test-collection",
|
||||
host: "localhost",
|
||||
port: 8000,
|
||||
...overrides,
|
||||
} as any);
|
||||
}
|
||||
|
||||
async function initDb(overrides: Record<string, any> = {}): Promise<ChromaDB> {
|
||||
const db = makeDb(overrides);
|
||||
await db.initialize();
|
||||
return db;
|
||||
}
|
||||
|
||||
// --- Reset mocks between tests ---
|
||||
|
||||
beforeEach(() => {
|
||||
jest.clearAllMocks();
|
||||
|
||||
__mocks__.add.mockResolvedValue(undefined);
|
||||
__mocks__.query.mockResolvedValue({
|
||||
ids: [[]],
|
||||
distances: [[]],
|
||||
metadatas: [[]],
|
||||
});
|
||||
__mocks__.get.mockResolvedValue({ ids: [], metadatas: [] });
|
||||
__mocks__.update.mockResolvedValue(undefined);
|
||||
__mocks__.upsert.mockResolvedValue(undefined);
|
||||
__mocks__.deleteFn.mockResolvedValue(undefined);
|
||||
|
||||
const collectionHandle = {
|
||||
add: __mocks__.add,
|
||||
query: __mocks__.query,
|
||||
get: __mocks__.get,
|
||||
update: __mocks__.update,
|
||||
upsert: __mocks__.upsert,
|
||||
delete: __mocks__.deleteFn,
|
||||
};
|
||||
__mocks__.getOrCreateCollection.mockResolvedValue(collectionHandle);
|
||||
__mocks__.deleteCollection.mockResolvedValue(undefined);
|
||||
__mocks__.ChromaClient.mockImplementation(() => ({
|
||||
getOrCreateCollection: __mocks__.getOrCreateCollection,
|
||||
deleteCollection: __mocks__.deleteCollection,
|
||||
}));
|
||||
__mocks__.CloudClient.mockImplementation(() => ({
|
||||
getOrCreateCollection: __mocks__.getOrCreateCollection,
|
||||
deleteCollection: __mocks__.deleteCollection,
|
||||
}));
|
||||
});
|
||||
|
||||
// --- Test suites ---
|
||||
|
||||
describe("VectorStoreFactory", () => {
|
||||
it("returns a ChromaDB instance for provider 'chroma'", async () => {
|
||||
const db = VectorStoreFactory.create("chroma", {
|
||||
collectionName: "x",
|
||||
host: "localhost",
|
||||
port: 8000,
|
||||
} as any);
|
||||
expect(db).toBeInstanceOf(ChromaDB);
|
||||
await (db as any).initialize();
|
||||
});
|
||||
});
|
||||
|
||||
describe("Constructor", () => {
|
||||
// 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,
|
||||
});
|
||||
expect(__mocks__.CloudClient).not.toHaveBeenCalled();
|
||||
});
|
||||
|
||||
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,
|
||||
ssl: true,
|
||||
path: "/db",
|
||||
});
|
||||
});
|
||||
|
||||
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",
|
||||
database: "mem0",
|
||||
});
|
||||
expect(__mocks__.ChromaClient).not.toHaveBeenCalled();
|
||||
});
|
||||
|
||||
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" }),
|
||||
);
|
||||
});
|
||||
|
||||
it("accepts a pre-built client via config.client", async () => {
|
||||
const fakeClient = {
|
||||
getOrCreateCollection: __mocks__.getOrCreateCollection,
|
||||
deleteCollection: __mocks__.deleteCollection,
|
||||
};
|
||||
const db = new ChromaDB({
|
||||
collectionName: "test-collection",
|
||||
client: fakeClient,
|
||||
} as any);
|
||||
await db.initialize();
|
||||
expect(__mocks__.ChromaClient).not.toHaveBeenCalled();
|
||||
expect(__mocks__.CloudClient).not.toHaveBeenCalled();
|
||||
expect(__mocks__.getOrCreateCollection).toHaveBeenCalled();
|
||||
});
|
||||
});
|
||||
|
||||
describe("initialize", () => {
|
||||
it("creates the collection with embeddingFunction null (mem0 supplies embeddings)", async () => {
|
||||
await initDb();
|
||||
expect(__mocks__.getOrCreateCollection).toHaveBeenCalledWith({
|
||||
name: "test-collection",
|
||||
embeddingFunction: null,
|
||||
});
|
||||
});
|
||||
|
||||
it("also creates the migrations collection", async () => {
|
||||
await initDb();
|
||||
expect(__mocks__.getOrCreateCollection).toHaveBeenCalledWith({
|
||||
name: "memory_migrations",
|
||||
embeddingFunction: null,
|
||||
});
|
||||
});
|
||||
});
|
||||
|
||||
describe("insert", () => {
|
||||
it("adds records with ids, embeddings, and metadatas", async () => {
|
||||
const db = await initDb();
|
||||
await db.insert([[1, 2, 3]], ["id-1"], [{ text: "hello" }]);
|
||||
expect(__mocks__.add).toHaveBeenCalledWith({
|
||||
ids: ["id-1"],
|
||||
embeddings: [[1, 2, 3]],
|
||||
metadatas: [{ text: "hello" }],
|
||||
});
|
||||
});
|
||||
});
|
||||
|
||||
describe("search", () => {
|
||||
it("queries with embeddings, nResults, and where", async () => {
|
||||
const db = await initDb();
|
||||
await db.search([1, 2, 3], 10, { user_id: "alice" });
|
||||
expect(__mocks__.query).toHaveBeenCalledWith({
|
||||
queryEmbeddings: [[1, 2, 3]],
|
||||
nResults: 10,
|
||||
where: { user_id: { $eq: "alice" } },
|
||||
});
|
||||
});
|
||||
|
||||
it("maps a nested query response to VectorStoreResult with 1/(1+distance) scores", async () => {
|
||||
__mocks__.query.mockResolvedValue({
|
||||
ids: [["a", "b"]],
|
||||
distances: [[0, 1]],
|
||||
metadatas: [[{ k: "v1" }, { k: "v2" }]],
|
||||
});
|
||||
const db = await initDb();
|
||||
const results = await db.search([1, 2, 3]);
|
||||
expect(results).toEqual([
|
||||
{ id: "a", payload: { k: "v1" }, score: 1 },
|
||||
{ id: "b", payload: { k: "v2" }, score: 0.5 },
|
||||
]);
|
||||
});
|
||||
|
||||
it("returns [] when the query yields no matches", async () => {
|
||||
__mocks__.query.mockResolvedValue({
|
||||
ids: [[]],
|
||||
distances: [[]],
|
||||
metadatas: [[]],
|
||||
});
|
||||
const db = await initDb();
|
||||
const results = await db.search([1, 2, 3]);
|
||||
expect(results).toEqual([]);
|
||||
});
|
||||
|
||||
it("passes where undefined when there are no filters", async () => {
|
||||
const db = await initDb();
|
||||
await db.search([1, 2, 3], 5);
|
||||
expect(__mocks__.query).toHaveBeenCalledWith(
|
||||
expect.objectContaining({ where: undefined }),
|
||||
);
|
||||
});
|
||||
});
|
||||
|
||||
describe("get", () => {
|
||||
it("fetches by id and returns the first parsed result", async () => {
|
||||
__mocks__.get.mockResolvedValue({
|
||||
ids: ["vec-1"],
|
||||
metadatas: [{ text: "foo" }],
|
||||
});
|
||||
const db = await initDb();
|
||||
const result = await db.get("vec-1");
|
||||
expect(__mocks__.get).toHaveBeenCalledWith({ ids: ["vec-1"] });
|
||||
expect(result).toEqual({ id: "vec-1", payload: { text: "foo" } });
|
||||
});
|
||||
|
||||
it("returns null when the id is not found", async () => {
|
||||
__mocks__.get.mockResolvedValue({ ids: [], metadatas: [] });
|
||||
const db = await initDb();
|
||||
const result = await db.get("missing");
|
||||
expect(result).toBeNull();
|
||||
});
|
||||
});
|
||||
|
||||
describe("update", () => {
|
||||
it("updates a single record with embedding and metadata", async () => {
|
||||
const db = await initDb();
|
||||
await db.update("vec-1", [1, 2, 3], { text: "updated" });
|
||||
expect(__mocks__.update).toHaveBeenCalledWith({
|
||||
ids: ["vec-1"],
|
||||
embeddings: [[1, 2, 3]],
|
||||
metadatas: [{ text: "updated" }],
|
||||
});
|
||||
});
|
||||
});
|
||||
|
||||
describe("delete", () => {
|
||||
it("deletes by id", async () => {
|
||||
const db = await initDb();
|
||||
await db.delete("vec-1");
|
||||
expect(__mocks__.deleteFn).toHaveBeenCalledWith({ ids: ["vec-1"] });
|
||||
});
|
||||
});
|
||||
|
||||
describe("deleteCol", () => {
|
||||
it("deletes the collection and resets its cached handle", async () => {
|
||||
const db = await initDb();
|
||||
await db.deleteCol();
|
||||
expect(__mocks__.deleteCollection).toHaveBeenCalledWith({
|
||||
name: "test-collection",
|
||||
});
|
||||
// Cached collection promise is cleared, so the next call re-creates it.
|
||||
__mocks__.getOrCreateCollection.mockClear();
|
||||
await db.search([1, 2, 3]);
|
||||
expect(__mocks__.getOrCreateCollection).toHaveBeenCalledWith({
|
||||
name: "test-collection",
|
||||
embeddingFunction: null,
|
||||
});
|
||||
});
|
||||
});
|
||||
|
||||
describe("list", () => {
|
||||
it("gets with where and limit and returns [results, count]", async () => {
|
||||
__mocks__.get.mockResolvedValue({
|
||||
ids: ["a", "b"],
|
||||
metadatas: [{ k: 1 }, { k: 2 }],
|
||||
});
|
||||
const db = await initDb();
|
||||
const [results, count] = await db.list({ user_id: "alice" }, 50);
|
||||
expect(__mocks__.get).toHaveBeenCalledWith({
|
||||
where: { user_id: { $eq: "alice" } },
|
||||
limit: 50,
|
||||
});
|
||||
expect(results).toHaveLength(2);
|
||||
expect(count).toBe(2);
|
||||
});
|
||||
});
|
||||
|
||||
describe("getUserId", () => {
|
||||
it("returns an existing user_id from the migrations collection", async () => {
|
||||
__mocks__.get.mockResolvedValue({
|
||||
ids: ["mig-1"],
|
||||
metadatas: [{ user_id: "u-123" }],
|
||||
});
|
||||
const db = await initDb();
|
||||
const uid = await db.getUserId();
|
||||
expect(uid).toBe("u-123");
|
||||
expect(__mocks__.add).not.toHaveBeenCalled();
|
||||
});
|
||||
|
||||
it("generates and stores a new user_id when none exists", async () => {
|
||||
__mocks__.get.mockResolvedValue({ ids: [], metadatas: [] });
|
||||
const db = await initDb();
|
||||
const uid = await db.getUserId();
|
||||
expect(typeof uid).toBe("string");
|
||||
expect(uid.length).toBeGreaterThan(0);
|
||||
expect(__mocks__.add).toHaveBeenCalledWith(
|
||||
expect.objectContaining({
|
||||
embeddings: [[0]],
|
||||
metadatas: [{ user_id: uid }],
|
||||
}),
|
||||
);
|
||||
});
|
||||
});
|
||||
|
||||
describe("setUserId", () => {
|
||||
it("upserts the user_id onto the existing migration marker", async () => {
|
||||
__mocks__.get.mockResolvedValue({
|
||||
ids: ["mig-1"],
|
||||
metadatas: [{ user_id: "old" }],
|
||||
});
|
||||
const db = await initDb();
|
||||
await db.setUserId("u-456");
|
||||
expect(__mocks__.upsert).toHaveBeenCalledWith({
|
||||
ids: ["mig-1"],
|
||||
embeddings: [[0]],
|
||||
metadatas: [{ user_id: "u-456" }],
|
||||
});
|
||||
});
|
||||
});
|
||||
|
||||
describe("keywordSearch", () => {
|
||||
it("returns null (Chroma has no keyword search)", async () => {
|
||||
const db = await initDb();
|
||||
const result = await db.keywordSearch();
|
||||
expect(result).toBeNull();
|
||||
});
|
||||
});
|
||||
|
||||
describe("generateWhereClause", () => {
|
||||
it("returns undefined for undefined and empty filters", () => {
|
||||
expect(ChromaDB.generateWhereClause(undefined)).toBeUndefined();
|
||||
expect(ChromaDB.generateWhereClause({})).toBeUndefined();
|
||||
});
|
||||
|
||||
it("converts a single equality filter", () => {
|
||||
expect(ChromaDB.generateWhereClause({ user_id: "alice" })).toEqual({
|
||||
user_id: { $eq: "alice" },
|
||||
});
|
||||
});
|
||||
|
||||
it("combines multiple keys with $and", () => {
|
||||
expect(
|
||||
ChromaDB.generateWhereClause({ user_id: "alice", agent_id: "bot" }),
|
||||
).toEqual({
|
||||
$and: [{ user_id: { $eq: "alice" } }, { agent_id: { $eq: "bot" } }],
|
||||
});
|
||||
});
|
||||
|
||||
it("converts an array value to $in", () => {
|
||||
expect(ChromaDB.generateWhereClause({ tags: ["x", "y"] })).toEqual({
|
||||
tags: { $in: ["x", "y"] },
|
||||
});
|
||||
});
|
||||
|
||||
it("skips a wildcard filter", () => {
|
||||
expect(ChromaDB.generateWhereClause({ user_id: "*" })).toBeUndefined();
|
||||
});
|
||||
|
||||
it("maps comparison operators", () => {
|
||||
expect(ChromaDB.generateWhereClause({ age: { gte: 18 } })).toEqual({
|
||||
age: { $gte: 18 },
|
||||
});
|
||||
});
|
||||
|
||||
it("converts an $or block", () => {
|
||||
expect(
|
||||
ChromaDB.generateWhereClause({ $or: [{ a: "x" }, { b: "y" }] }),
|
||||
).toEqual({ $or: [{ a: { $eq: "x" } }, { b: { $eq: "y" } }] });
|
||||
});
|
||||
|
||||
it("collapses a single-branch $or", () => {
|
||||
expect(ChromaDB.generateWhereClause({ $or: [{ a: "x" }] })).toEqual({
|
||||
a: { $eq: "x" },
|
||||
});
|
||||
});
|
||||
|
||||
it("negates a single-field $not condition", () => {
|
||||
expect(ChromaDB.generateWhereClause({ $not: [{ a: "x" }] })).toEqual({
|
||||
a: { $ne: "x" },
|
||||
});
|
||||
});
|
||||
|
||||
it("applies De Morgan to a multi-field $not condition", () => {
|
||||
// NOT(a=x AND b=y) is (a!=x) OR (b!=y)
|
||||
expect(
|
||||
ChromaDB.generateWhereClause({ $not: [{ a: "x", b: "y" }] }),
|
||||
).toEqual({ $or: [{ a: { $ne: "x" } }, { b: { $ne: "y" } }] });
|
||||
});
|
||||
|
||||
it("negates a $not comparison operator", () => {
|
||||
expect(
|
||||
ChromaDB.generateWhereClause({ $not: [{ age: { gt: 18 } }] }),
|
||||
).toEqual({ age: { $lte: 18 } });
|
||||
});
|
||||
});
|
||||
@@ -0,0 +1,517 @@
|
||||
import { Milvus } from "../vector_stores/milvus";
|
||||
|
||||
/**
|
||||
* In-memory fake of the subset of the `@zilliz/milvus2-sdk-node` MilvusClient
|
||||
* API that the Milvus vector store uses. Lets us exercise the provider's
|
||||
* request shaping and response parsing without a live Milvus server or the SDK.
|
||||
*/
|
||||
class FakeMilvusClient {
|
||||
public calls: { method: string; args: any }[] = [];
|
||||
private collections = new Set<string>();
|
||||
// collection -> id -> row
|
||||
private store: Record<string, Record<string, any>> = {};
|
||||
// collection -> declared vector dimension (recorded from createCollection)
|
||||
private dims: Record<string, number> = {};
|
||||
// collection -> declared field names (surfaced via describeCollection)
|
||||
private fieldNames: Record<string, string[]> = {};
|
||||
|
||||
// Allow tests to script search responses.
|
||||
public searchResponse: any = { results: [] };
|
||||
|
||||
constructor(opts?: { existing?: string[]; bm25?: string[] }) {
|
||||
// Legacy dense-only collections: no text/sparse fields.
|
||||
for (const c of opts?.existing || []) {
|
||||
this.collections.add(c);
|
||||
this.store[c] = {};
|
||||
this.fieldNames[c] = ["id", "vectors", "metadata"];
|
||||
}
|
||||
// Pre-existing collections that already carry the BM25 schema.
|
||||
for (const c of opts?.bm25 || []) {
|
||||
this.collections.add(c);
|
||||
this.store[c] = {};
|
||||
this.fieldNames[c] = ["id", "vectors", "metadata", "text", "sparse"];
|
||||
}
|
||||
}
|
||||
|
||||
async hasCollection({ collection_name }: any) {
|
||||
this.calls.push({ method: "hasCollection", args: { collection_name } });
|
||||
return { value: this.collections.has(collection_name) };
|
||||
}
|
||||
|
||||
async createCollection(args: any) {
|
||||
this.calls.push({ method: "createCollection", args });
|
||||
// Mirror Milvus's real server constraint: a FloatVector field's dim must be
|
||||
// in [2, 32768]. Vector fields are the ones carrying a numeric `dim`.
|
||||
for (const f of args.fields || []) {
|
||||
if (typeof f.dim === "number" && (f.dim < 2 || f.dim > 32768)) {
|
||||
throw new Error(
|
||||
`invalid dimension: ${f.dim}. should be in range 2 ~ 32768`,
|
||||
);
|
||||
}
|
||||
}
|
||||
const vectorField = (args.fields || []).find(
|
||||
(f: any) => typeof f.dim === "number",
|
||||
);
|
||||
if (vectorField) this.dims[args.collection_name] = vectorField.dim;
|
||||
this.fieldNames[args.collection_name] = (args.fields || []).map(
|
||||
(f: any) => f.name,
|
||||
);
|
||||
this.collections.add(args.collection_name);
|
||||
this.store[args.collection_name] = this.store[args.collection_name] || {};
|
||||
}
|
||||
|
||||
async describeCollection({ collection_name }: any) {
|
||||
this.calls.push({
|
||||
method: "describeCollection",
|
||||
args: { collection_name },
|
||||
});
|
||||
const fields = (this.fieldNames[collection_name] || []).map((name) => ({
|
||||
name,
|
||||
}));
|
||||
return { schema: { fields } };
|
||||
}
|
||||
|
||||
// Reject rows whose vector length disagrees with the collection's declared
|
||||
// dim, exactly as the real server would.
|
||||
private checkDims(collection: string, data: any[]) {
|
||||
const dim = this.dims[collection];
|
||||
if (dim == null) return; // pre-seeded collection: dim not tracked
|
||||
for (const row of data || []) {
|
||||
if (Array.isArray(row.vectors) && row.vectors.length !== dim) {
|
||||
throw new Error(
|
||||
`vector dimension mismatch: expected ${dim}, got ${row.vectors.length}`,
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
async loadCollection(args: any) {
|
||||
this.calls.push({ method: "loadCollection", args });
|
||||
}
|
||||
|
||||
async dropCollection(args: any) {
|
||||
this.calls.push({ method: "dropCollection", args });
|
||||
this.collections.delete(args.collection_name);
|
||||
delete this.store[args.collection_name];
|
||||
}
|
||||
|
||||
async insert(args: any) {
|
||||
this.calls.push({ method: "insert", args });
|
||||
this.checkDims(args.collection_name, args.data);
|
||||
const col = (this.store[args.collection_name] =
|
||||
this.store[args.collection_name] || {});
|
||||
for (const row of args.data) col[String(row.id)] = row;
|
||||
}
|
||||
|
||||
async upsert(args: any) {
|
||||
this.calls.push({ method: "upsert", args });
|
||||
this.checkDims(args.collection_name, args.data);
|
||||
const col = (this.store[args.collection_name] =
|
||||
this.store[args.collection_name] || {});
|
||||
for (const row of args.data) col[String(row.id)] = row;
|
||||
}
|
||||
|
||||
async delete(args: any) {
|
||||
this.calls.push({ method: "delete", args });
|
||||
const col = this.store[args.collection_name] || {};
|
||||
for (const id of args.ids) delete col[String(id)];
|
||||
}
|
||||
|
||||
async get(args: any) {
|
||||
this.calls.push({ method: "get", args });
|
||||
const col = this.store[args.collection_name] || {};
|
||||
const data = args.ids.map((id: string) => col[String(id)]).filter(Boolean);
|
||||
return { data };
|
||||
}
|
||||
|
||||
async query(args: any) {
|
||||
this.calls.push({ method: "query", args });
|
||||
const col = this.store[args.collection_name] || {};
|
||||
return { data: Object.values(col).slice(0, args.limit ?? 100) };
|
||||
}
|
||||
|
||||
async search(args: any) {
|
||||
this.calls.push({ method: "search", args });
|
||||
return this.searchResponse;
|
||||
}
|
||||
}
|
||||
|
||||
// Suppress the constructor's fire-and-forget initialize() console noise and the
|
||||
// legacy-collection BM25 warning.
|
||||
beforeAll(() => {
|
||||
jest.spyOn(console, "error").mockImplementation(() => {});
|
||||
jest.spyOn(console, "warn").mockImplementation(() => {});
|
||||
});
|
||||
afterAll(() => {
|
||||
(console.error as jest.Mock).mockRestore?.();
|
||||
(console.warn as jest.Mock).mockRestore?.();
|
||||
});
|
||||
|
||||
describe("Milvus vector store (TS OSS SDK)", () => {
|
||||
function makeStore(client: FakeMilvusClient, overrides: any = {}) {
|
||||
return new Milvus({
|
||||
client,
|
||||
collectionName: "mem0",
|
||||
embeddingModelDims: 3,
|
||||
...overrides,
|
||||
});
|
||||
}
|
||||
|
||||
it("creates the collection on initialize when it does not exist", async () => {
|
||||
const client = new FakeMilvusClient();
|
||||
const store = makeStore(client);
|
||||
await store.initialize();
|
||||
|
||||
const created = client.calls.find((c) => c.method === "createCollection");
|
||||
expect(created).toBeDefined();
|
||||
expect(created!.args.collection_name).toBe("mem0");
|
||||
const vectorField = created!.args.fields.find(
|
||||
(f: any) => f.name === "vectors",
|
||||
);
|
||||
expect(vectorField.dim).toBe(3);
|
||||
// AUTOINDEX dense index with the default metric (L2, matching the Python provider).
|
||||
expect(created!.args.index_params[0].metric_type).toBe("L2");
|
||||
});
|
||||
|
||||
it("does not recreate an existing collection", async () => {
|
||||
const client = new FakeMilvusClient({ existing: ["mem0"] });
|
||||
const store = makeStore(client);
|
||||
await store.initialize();
|
||||
expect(client.calls.some((c) => c.method === "createCollection")).toBe(
|
||||
false,
|
||||
);
|
||||
});
|
||||
|
||||
it("inserts records mapping payloads into the metadata field", async () => {
|
||||
const client = new FakeMilvusClient({ existing: ["mem0"] });
|
||||
const store = makeStore(client);
|
||||
await store.initialize();
|
||||
|
||||
await store.insert(
|
||||
[
|
||||
[0.1, 0.2, 0.3],
|
||||
[0.4, 0.5, 0.6],
|
||||
],
|
||||
["a", "b"],
|
||||
[
|
||||
{ data: "first", user_id: "u1" },
|
||||
{ data: "second", user_id: "u1" },
|
||||
],
|
||||
);
|
||||
|
||||
const insertCall = client.calls.find((c) => c.method === "insert")!;
|
||||
expect(insertCall.args.data).toHaveLength(2);
|
||||
expect(insertCall.args.data[0]).toEqual({
|
||||
id: "a",
|
||||
vectors: [0.1, 0.2, 0.3],
|
||||
metadata: { data: "first", user_id: "u1" },
|
||||
});
|
||||
});
|
||||
|
||||
it("round-trips a stored record through get()", async () => {
|
||||
const client = new FakeMilvusClient({ existing: ["mem0"] });
|
||||
const store = makeStore(client);
|
||||
await store.initialize();
|
||||
await store.insert([[1, 0, 0]], ["x"], [{ data: "hello" }]);
|
||||
|
||||
const got = await store.get("x");
|
||||
expect(got).not.toBeNull();
|
||||
expect(got!.id).toBe("x");
|
||||
expect(got!.payload).toEqual({ data: "hello" });
|
||||
|
||||
const missing = await store.get("nope");
|
||||
expect(missing).toBeNull();
|
||||
});
|
||||
|
||||
it("updates a record in place via upsert", async () => {
|
||||
const client = new FakeMilvusClient({ existing: ["mem0"] });
|
||||
const store = makeStore(client);
|
||||
await store.initialize();
|
||||
await store.insert([[1, 0, 0]], ["a"], [{ data: "old" }]);
|
||||
|
||||
await store.update("a", [0, 1, 0], { data: "new" });
|
||||
|
||||
const upsertCall = client.calls.find((c) => c.method === "upsert")!;
|
||||
expect(upsertCall.args.data[0]).toEqual({
|
||||
id: "a",
|
||||
vectors: [0, 1, 0],
|
||||
metadata: { data: "new" },
|
||||
});
|
||||
const got = await store.get("a");
|
||||
expect(got!.payload).toEqual({ data: "new" });
|
||||
});
|
||||
|
||||
it("builds an AND-combined equality filter expression for search", async () => {
|
||||
const client = new FakeMilvusClient({ existing: ["mem0"] });
|
||||
client.searchResponse = {
|
||||
results: [{ id: "a", score: 0.9, metadata: { data: "first" } }],
|
||||
};
|
||||
// Pin a non-L2 metric so the raw score passes through unnormalised.
|
||||
const store = makeStore(client, { metricType: "COSINE" });
|
||||
await store.initialize();
|
||||
|
||||
const res = await store.search([0.1, 0.2, 0.3], 5, {
|
||||
user_id: "u1",
|
||||
agent_id: 7 as any,
|
||||
});
|
||||
|
||||
const searchCall = client.calls.find((c) => c.method === "search")!;
|
||||
expect(searchCall.args.filter).toBe(
|
||||
'(metadata["user_id"] == "u1") and (metadata["agent_id"] == 7)',
|
||||
);
|
||||
expect(searchCall.args.limit).toBe(5);
|
||||
expect(res[0]).toEqual({
|
||||
id: "a",
|
||||
payload: { data: "first" },
|
||||
score: 0.9,
|
||||
});
|
||||
});
|
||||
|
||||
it("normalises L2 distances into a 0..1 similarity score", async () => {
|
||||
const client = new FakeMilvusClient({ existing: ["mem0"] });
|
||||
client.searchResponse = {
|
||||
results: [{ id: "a", score: 3.0, metadata: {} }],
|
||||
};
|
||||
const store = makeStore(client, { metricType: "L2" });
|
||||
await store.initialize();
|
||||
|
||||
const res = await store.search([0, 0, 1], 1);
|
||||
// 1 / (1 + 3) = 0.25
|
||||
expect(res[0].score).toBeCloseTo(0.25, 6);
|
||||
});
|
||||
|
||||
it("escapes embedded quotes in string filter values", async () => {
|
||||
const client = new FakeMilvusClient({ existing: ["mem0"] });
|
||||
const store = makeStore(client);
|
||||
await store.initialize();
|
||||
await store.list({ data: 'a"b' });
|
||||
const queryCall = client.calls.filter((c) => c.method === "query").pop()!;
|
||||
expect(queryCall.args.filter).toBe('(metadata["data"] == "a\\"b")');
|
||||
});
|
||||
|
||||
it("escapes backslashes in string filter values", async () => {
|
||||
const client = new FakeMilvusClient({ existing: ["mem0"] });
|
||||
const store = makeStore(client);
|
||||
await store.initialize();
|
||||
// Value is a\b (one backslash); it must be doubled so the expression
|
||||
// stays well-formed and the backslash can't escape the closing quote.
|
||||
await store.list({ data: "a\\b" });
|
||||
const queryCall = client.calls.filter((c) => c.method === "query").pop()!;
|
||||
expect(queryCall.args.filter).toBe('(metadata["data"] == "a\\\\b")');
|
||||
});
|
||||
|
||||
it("rejects filter keys that are not safe identifiers", async () => {
|
||||
const client = new FakeMilvusClient({ existing: ["mem0"] });
|
||||
const store = makeStore(client);
|
||||
await store.initialize();
|
||||
await expect(
|
||||
store.list({ 'x"] or true or ["y': "z" } as any),
|
||||
).rejects.toThrow(/Invalid filter key/);
|
||||
});
|
||||
|
||||
it("lists records and returns the count", async () => {
|
||||
const client = new FakeMilvusClient({ existing: ["mem0"] });
|
||||
const store = makeStore(client);
|
||||
await store.initialize();
|
||||
await store.insert(
|
||||
[
|
||||
[1, 0, 0],
|
||||
[0, 1, 0],
|
||||
],
|
||||
["a", "b"],
|
||||
[{ data: "x" }, { data: "y" }],
|
||||
);
|
||||
|
||||
const [results, count] = await store.list(undefined, 100);
|
||||
expect(count).toBe(2);
|
||||
expect(results.map((r) => r.id).sort()).toEqual(["a", "b"]);
|
||||
});
|
||||
|
||||
it("deletes a record by id", async () => {
|
||||
const client = new FakeMilvusClient({ existing: ["mem0"] });
|
||||
const store = makeStore(client);
|
||||
await store.initialize();
|
||||
await store.insert([[1, 0, 0]], ["a"], [{ data: "x" }]);
|
||||
await store.delete("a");
|
||||
const got = await store.get("a");
|
||||
expect(got).toBeNull();
|
||||
});
|
||||
|
||||
it("generates and persists a user id, then reads it back", async () => {
|
||||
const client = new FakeMilvusClient({ existing: ["mem0"] });
|
||||
const store = makeStore(client);
|
||||
await store.initialize();
|
||||
|
||||
const created = await store.getUserId();
|
||||
expect(typeof created).toBe("string");
|
||||
expect(created.length).toBeGreaterThan(0);
|
||||
|
||||
const readBack = await store.getUserId();
|
||||
expect(readBack).toBe(created);
|
||||
});
|
||||
|
||||
it("setUserId overwrites in place instead of appending rows", async () => {
|
||||
const client = new FakeMilvusClient();
|
||||
const store = makeStore(client);
|
||||
await store.initialize();
|
||||
|
||||
await store.setUserId("user-1");
|
||||
await store.setUserId("user-2");
|
||||
|
||||
// A fresh insert per call would leave two rows; the reuse-id upsert keeps one.
|
||||
const all = await client.query({
|
||||
collection_name: "memory_migrations",
|
||||
limit: 100,
|
||||
});
|
||||
expect(all.data).toHaveLength(1);
|
||||
expect(all.data[0].user_id).toBe("user-2");
|
||||
expect(await store.getUserId()).toBe("user-2");
|
||||
});
|
||||
|
||||
it("creates the telemetry collection with a Milvus-valid vector dim (>= 2)", async () => {
|
||||
// Regression: the helper collection previously used dim 1, which a real
|
||||
// Milvus server rejects (valid range is 2~32768), silently breaking
|
||||
// getUserId/setUserId. The fake now enforces the same bound, so a dim < 2
|
||||
// would throw here instead of passing as it did against the old mock.
|
||||
const client = new FakeMilvusClient({ existing: ["mem0"] });
|
||||
const store = makeStore(client);
|
||||
await store.initialize();
|
||||
await store.getUserId();
|
||||
|
||||
const migCreate = client.calls.find(
|
||||
(c) =>
|
||||
c.method === "createCollection" &&
|
||||
c.args.collection_name === "memory_migrations",
|
||||
)!;
|
||||
expect(migCreate).toBeDefined();
|
||||
const vectorField = migCreate.args.fields.find(
|
||||
(f: any) => f.name === "vectors",
|
||||
);
|
||||
expect(vectorField.dim).toBeGreaterThanOrEqual(2);
|
||||
});
|
||||
|
||||
it("keywordSearch returns null on a collection without the BM25 schema", async () => {
|
||||
const client = new FakeMilvusClient({ existing: ["mem0"] });
|
||||
const store = makeStore(client);
|
||||
await store.initialize();
|
||||
expect(await store.keywordSearch("hello")).toBeNull();
|
||||
});
|
||||
|
||||
it("creates a BM25 hybrid schema on a fresh collection", async () => {
|
||||
const client = new FakeMilvusClient();
|
||||
const store = makeStore(client);
|
||||
await store.initialize();
|
||||
|
||||
const created = client.calls.find((c) => c.method === "createCollection")!;
|
||||
const fieldNames = created.args.fields.map((f: any) => f.name);
|
||||
expect(fieldNames).toEqual(
|
||||
expect.arrayContaining(["id", "vectors", "metadata", "text", "sparse"]),
|
||||
);
|
||||
// The text field must have the analyzer enabled for BM25 tokenization.
|
||||
const textField = created.args.fields.find((f: any) => f.name === "text");
|
||||
expect(textField.enable_analyzer).toBe(true);
|
||||
// BM25 function maps text -> sparse.
|
||||
expect(created.args.functions[0].input_field_names).toEqual(["text"]);
|
||||
expect(created.args.functions[0].output_field_names).toEqual(["sparse"]);
|
||||
// Sparse BM25 index sits alongside the dense vector index.
|
||||
const sparseIdx = created.args.index_params.find(
|
||||
(i: any) => i.field_name === "sparse",
|
||||
);
|
||||
expect(sparseIdx.index_type).toBe("SPARSE_INVERTED_INDEX");
|
||||
expect(sparseIdx.metric_type).toBe("BM25");
|
||||
});
|
||||
|
||||
it("populates the BM25 text field from payload on a fresh collection", async () => {
|
||||
const client = new FakeMilvusClient();
|
||||
const store = makeStore(client);
|
||||
await store.initialize();
|
||||
|
||||
// Prefers the lemmatized text when present.
|
||||
await store.insert(
|
||||
[[0.1, 0.2, 0.3]],
|
||||
["a"],
|
||||
[{ data: "hello world", text_lemmatized: "hello world lemma" }],
|
||||
);
|
||||
// Falls back to raw data when there is no lemmatized text.
|
||||
await store.insert([[0.4, 0.5, 0.6]], ["b"], [{ data: "just data" }]);
|
||||
|
||||
const insertCalls = client.calls.filter((c) => c.method === "insert");
|
||||
expect(insertCalls[0].args.data[0].text).toBe("hello world lemma");
|
||||
expect(insertCalls[1].args.data[0].text).toBe("just data");
|
||||
});
|
||||
|
||||
it("writes the BM25 text field on update for a BM25 collection", async () => {
|
||||
const client = new FakeMilvusClient();
|
||||
const store = makeStore(client);
|
||||
await store.initialize();
|
||||
await store.insert([[1, 0, 0]], ["a"], [{ data: "old" }]);
|
||||
|
||||
await store.update("a", [0, 1, 0], { data: "new" });
|
||||
|
||||
const upsertCall = client.calls.filter((c) => c.method === "upsert").pop()!;
|
||||
expect(upsertCall.args.data[0].text).toBe("new");
|
||||
expect(upsertCall.args.data[0].metadata).toEqual({ data: "new" });
|
||||
});
|
||||
|
||||
it("names the dense anns_field when searching a BM25 collection", async () => {
|
||||
const client = new FakeMilvusClient();
|
||||
client.searchResponse = {
|
||||
results: [{ id: "a", score: 0.5, metadata: {} }],
|
||||
};
|
||||
const store = makeStore(client, { metricType: "COSINE" });
|
||||
await store.initialize();
|
||||
|
||||
await store.search([0.1, 0.2, 0.3], 3);
|
||||
const searchCall = client.calls.find((c) => c.method === "search")!;
|
||||
expect(searchCall.args.anns_field).toBe("vectors");
|
||||
});
|
||||
|
||||
it("omits anns_field when searching a legacy dense-only collection", async () => {
|
||||
const client = new FakeMilvusClient({ existing: ["mem0"] });
|
||||
client.searchResponse = { results: [] };
|
||||
const store = makeStore(client, { metricType: "COSINE" });
|
||||
await store.initialize();
|
||||
|
||||
await store.search([0.1, 0.2, 0.3], 3);
|
||||
const searchCall = client.calls.find((c) => c.method === "search")!;
|
||||
expect(searchCall.args.anns_field).toBeUndefined();
|
||||
});
|
||||
|
||||
it("runs a BM25 keyword search over the sparse field and parses hits", async () => {
|
||||
const client = new FakeMilvusClient();
|
||||
client.searchResponse = {
|
||||
results: [{ id: "a", score: 4.2, metadata: { data: "kw hit" } }],
|
||||
};
|
||||
// COSINE so the BM25 score passes through unnormalised for a clean assert.
|
||||
const store = makeStore(client, { metricType: "COSINE" });
|
||||
await store.initialize();
|
||||
|
||||
const res = await store.keywordSearch("hello", 7, { user_id: "u1" });
|
||||
|
||||
const searchCall = client.calls.filter((c) => c.method === "search").pop()!;
|
||||
expect(searchCall.args.data).toEqual(["hello"]); // raw text, not a vector
|
||||
expect(searchCall.args.anns_field).toBe("sparse");
|
||||
expect(searchCall.args.limit).toBe(7);
|
||||
expect(searchCall.args.filter).toBe('(metadata["user_id"] == "u1")');
|
||||
expect(res).not.toBeNull();
|
||||
expect(res![0]).toEqual({
|
||||
id: "a",
|
||||
payload: { data: "kw hit" },
|
||||
score: 4.2,
|
||||
});
|
||||
});
|
||||
|
||||
it("detects the BM25 schema on a pre-existing collection via describeCollection", async () => {
|
||||
const client = new FakeMilvusClient({ bm25: ["mem0"] });
|
||||
client.searchResponse = { results: [] };
|
||||
const store = makeStore(client, { metricType: "COSINE" });
|
||||
await store.initialize();
|
||||
|
||||
// Detected as BM25: keywordSearch runs (returns []) instead of null, and
|
||||
// insert writes the text field.
|
||||
expect(await store.keywordSearch("q")).not.toBeNull();
|
||||
await store.insert([[1, 2, 3]], ["a"], [{ data: "d" }]);
|
||||
const insertCall = client.calls.find((c) => c.method === "insert")!;
|
||||
expect(insertCall.args.data[0].text).toBe("d");
|
||||
});
|
||||
});
|
||||
@@ -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";
|
||||
|
||||
@@ -19,8 +19,19 @@ export interface EmbeddingConfig {
|
||||
url?: string;
|
||||
embeddingDims?: number;
|
||||
modelProperties?: Record<string, any>;
|
||||
// HuggingFace TEI / OpenAI-compatible inference endpoint base URL.
|
||||
huggingfaceBaseUrl?: string;
|
||||
}
|
||||
|
||||
export interface VertexAIConfig extends EmbeddingConfig {
|
||||
vertexCredentialsJson?: string;
|
||||
googleServiceAccountJson?: string | Record<string, any>;
|
||||
googleProjectId?: string;
|
||||
location?: string;
|
||||
memoryAddEmbeddingType?: string;
|
||||
memoryUpdateEmbeddingType?: string;
|
||||
memorySearchEmbeddingType?: string;
|
||||
}
|
||||
export type { ValkeyConfig } from "./valkey";
|
||||
|
||||
export interface VectorStoreConfig {
|
||||
@@ -56,6 +67,61 @@ export interface LLMConfig {
|
||||
temperature?: number;
|
||||
topP?: number;
|
||||
maxTokens?: number;
|
||||
// AWS Bedrock provider config (used when provider === "aws_bedrock").
|
||||
// Credentials otherwise resolve via the standard AWS credential chain.
|
||||
awsRegion?: string;
|
||||
awsAccessKeyId?: string;
|
||||
awsSecretAccessKey?: string;
|
||||
awsSessionToken?: string;
|
||||
// Optional pre-constructed client (e.g. BedrockRuntimeClient) for DI/testing.
|
||||
client?: any;
|
||||
}
|
||||
|
||||
export interface RerankerConfig {
|
||||
apiKey?: string;
|
||||
/** The reranker model to use. Default varies by provider. */
|
||||
model?: string;
|
||||
/** Maximum number of documents to return after reranking. Default: unset (return all). */
|
||||
topK?: number;
|
||||
/** `cohere` only. Return document texts in the response. Default: `false`. */
|
||||
returnDocuments?: boolean;
|
||||
/** `cohere` only. Maximum number of chunks per document. Default: unset. */
|
||||
maxChunksPerDoc?: number;
|
||||
/**
|
||||
* `sentence_transformer` / `huggingface` only. Transformers.js device, e.g.
|
||||
* `"cpu"`, `"wasm"`, `"webgpu"`. Default: unset (auto-detect).
|
||||
*/
|
||||
device?: string;
|
||||
/** `huggingface` only. Max token length per query-document pair. Default: `512`. */
|
||||
maxLength?: number;
|
||||
/**
|
||||
* `sentence_transformer` / `huggingface` only. Sigmoid-normalize raw logits
|
||||
* to `[0, 1]`. Default: `true`; set `false` to surface raw logits.
|
||||
*/
|
||||
normalize?: boolean;
|
||||
/** No-op: a search reranks a small candidate set in one forward pass. */
|
||||
batchSize?: number;
|
||||
/** No-op in this runtime. */
|
||||
showProgressBar?: boolean;
|
||||
/**
|
||||
* `llm_reranker` only. LLM provider used to build the scoring LLM when
|
||||
* `llm` is not set. Default: `"openai"`.
|
||||
*/
|
||||
provider?: string;
|
||||
/** `llm_reranker` only. Temperature for LLM generation. Default: `0.0`. */
|
||||
temperature?: number;
|
||||
/** `llm_reranker` only. Maximum tokens for the LLM response. Default: `100`. */
|
||||
maxTokens?: number;
|
||||
/**
|
||||
* `llm_reranker` only. Nested LLM configuration. When set, it overrides the
|
||||
* top-level `provider`/`model`/`temperature`/`maxTokens`/`apiKey`, which
|
||||
* then only act as defaults for fields missing from `llm.config`.
|
||||
*/
|
||||
llm?: {
|
||||
provider: string;
|
||||
config: LLMConfig;
|
||||
};
|
||||
[key: string]: any;
|
||||
}
|
||||
|
||||
export interface MemoryConfig {
|
||||
@@ -72,6 +138,10 @@ export interface MemoryConfig {
|
||||
provider: string;
|
||||
config: LLMConfig;
|
||||
};
|
||||
reranker?: {
|
||||
provider: string;
|
||||
config: RerankerConfig;
|
||||
};
|
||||
historyStore?: HistoryStoreConfig;
|
||||
disableHistory?: boolean;
|
||||
historyDbPath?: string;
|
||||
@@ -85,6 +155,8 @@ export interface MemoryItem {
|
||||
createdAt?: string;
|
||||
updatedAt?: string;
|
||||
score?: number;
|
||||
/** Relevance score added by the reranker, alongside (not replacing) `score`. */
|
||||
rerankScore?: number;
|
||||
metadata?: Record<string, any>;
|
||||
attributedTo?: string;
|
||||
}
|
||||
@@ -117,6 +189,15 @@ export const MemoryConfigSchema = z.object({
|
||||
baseURL: z.string().optional(),
|
||||
embeddingDims: z.number().optional(),
|
||||
url: z.string().optional(),
|
||||
vertexCredentialsJson: z.string().optional(),
|
||||
googleServiceAccountJson: z
|
||||
.union([z.string(), z.record(z.string(), z.any())])
|
||||
.optional(),
|
||||
googleProjectId: z.string().optional(),
|
||||
location: z.string().optional(),
|
||||
memoryAddEmbeddingType: z.string().optional(),
|
||||
memoryUpdateEmbeddingType: z.string().optional(),
|
||||
memorySearchEmbeddingType: z.string().optional(),
|
||||
}),
|
||||
}),
|
||||
vectorStore: z.object({
|
||||
@@ -130,21 +211,29 @@ export const MemoryConfigSchema = z.object({
|
||||
})
|
||||
.passthrough(),
|
||||
}),
|
||||
|
||||
llm: z.object({
|
||||
provider: z.string(),
|
||||
config: z.object({
|
||||
apiKey: z.string().optional(),
|
||||
model: z.union([z.string(), z.any()]).optional(),
|
||||
modelProperties: z.record(z.string(), z.any()).optional(),
|
||||
baseURL: z.string().optional(),
|
||||
vllmBaseURL: z.string().optional(),
|
||||
vllm_base_url: z.string().optional(),
|
||||
url: z.string().optional(),
|
||||
timeout: z.number().optional(),
|
||||
temperature: z.number().optional(),
|
||||
topP: z.number().optional(),
|
||||
maxTokens: z.number().optional(),
|
||||
}),
|
||||
config: z
|
||||
.object({
|
||||
apiKey: z.string().optional(),
|
||||
model: z.union([z.string(), z.any()]).optional(),
|
||||
modelProperties: z.record(z.string(), z.any()).optional(),
|
||||
baseURL: z.string().optional(),
|
||||
vllmBaseURL: z.string().optional(),
|
||||
vllm_base_url: z.string().optional(),
|
||||
url: z.string().optional(),
|
||||
timeout: z.number().optional(),
|
||||
temperature: z.number().optional(),
|
||||
topP: z.number().optional(),
|
||||
maxTokens: z.number().optional(),
|
||||
awsRegion: z.string().optional(),
|
||||
awsAccessKeyId: z.string().optional(),
|
||||
awsSecretAccessKey: z.string().optional(),
|
||||
awsSessionToken: z.string().optional(),
|
||||
client: z.any().optional(),
|
||||
})
|
||||
.passthrough(),
|
||||
}),
|
||||
historyDbPath: z.string().optional(),
|
||||
customInstructions: z.string().optional(),
|
||||
@@ -154,5 +243,11 @@ export const MemoryConfigSchema = z.object({
|
||||
config: z.record(z.string(), z.any()),
|
||||
})
|
||||
.optional(),
|
||||
reranker: z
|
||||
.object({
|
||||
provider: z.string(),
|
||||
config: z.record(z.string(), z.any()),
|
||||
})
|
||||
.optional(),
|
||||
disableHistory: z.boolean().optional(),
|
||||
});
|
||||
|
||||
@@ -0,0 +1,52 @@
|
||||
/**
|
||||
* Expiration date handling for memories.
|
||||
*
|
||||
* Memories may carry an `expiration_date` (YYYY-MM-DD) after which they are
|
||||
* hidden from `getAll()` and `search()` unless `showExpired` is set.
|
||||
*/
|
||||
|
||||
const EXPIRATION_DATE_PATTERN = /^(\d{4})-(\d{2})-(\d{2})$/;
|
||||
|
||||
function todayUtc(): string {
|
||||
return new Date().toISOString().slice(0, 10);
|
||||
}
|
||||
|
||||
/**
|
||||
* Normalize a user-supplied expiration date to a YYYY-MM-DD string.
|
||||
*
|
||||
* Deliberately stricter than `new Date(value)`, which accepts formats the
|
||||
* Python SDK rejects ("12/31/2099", "2099") and resolves them against the
|
||||
* local timezone, shifting the calendar day. It also silently rolls invalid
|
||||
* dates over — `new Date("2099-02-30T00:00:00Z")` yields March 2nd.
|
||||
*/
|
||||
export function normalizeExpirationDate(value: string): string {
|
||||
const match = EXPIRATION_DATE_PATTERN.exec(value);
|
||||
if (match) {
|
||||
const [, year, month, day] = match;
|
||||
const parsed = new Date(`${value}T00:00:00Z`);
|
||||
if (
|
||||
!Number.isNaN(parsed.getTime()) &&
|
||||
parsed.getUTCFullYear() === Number(year) &&
|
||||
parsed.getUTCMonth() === Number(month) - 1 &&
|
||||
parsed.getUTCDate() === Number(day)
|
||||
) {
|
||||
return value;
|
||||
}
|
||||
}
|
||||
throw new Error("expirationDate must be a valid date in YYYY-MM-DD format.");
|
||||
}
|
||||
|
||||
/** True when the payload carries an expiration date strictly before today (UTC). */
|
||||
export function payloadIsExpired(
|
||||
payload: Record<string, any> | null | undefined,
|
||||
) {
|
||||
const raw = payload?.expiration_date;
|
||||
if (!raw) return false;
|
||||
try {
|
||||
// YYYY-MM-DD sorts lexicographically the same way it sorts chronologically.
|
||||
return normalizeExpirationDate(String(raw)) < todayUtc();
|
||||
} catch {
|
||||
// Unparseable stored value: treat as non-expiring rather than hiding data.
|
||||
return false;
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,86 @@
|
||||
jest.mock("zeroentropy", () => ({ ZeroEntropy: jest.fn() }));
|
||||
jest.mock("@huggingface/transformers", () => ({
|
||||
AutoModelForSequenceClassification: { from_pretrained: jest.fn() },
|
||||
AutoTokenizer: { from_pretrained: jest.fn() },
|
||||
}));
|
||||
|
||||
import { RerankerFactory } from "./factory";
|
||||
import { CohereReranker } from "../rerankers/cohere";
|
||||
import { LLMReranker } from "../rerankers/llm";
|
||||
import { ZeroEntropyReranker } from "../rerankers/zeroentropy";
|
||||
import { CrossEncoderReranker } from "../rerankers/cross_encoder";
|
||||
|
||||
describe("RerankerFactory", () => {
|
||||
const originalOpenAiKey = process.env.OPENAI_API_KEY;
|
||||
|
||||
afterEach(() => {
|
||||
if (originalOpenAiKey === undefined) {
|
||||
delete process.env.OPENAI_API_KEY;
|
||||
} else {
|
||||
process.env.OPENAI_API_KEY = originalOpenAiKey;
|
||||
}
|
||||
});
|
||||
|
||||
it("creates a CohereReranker for provider 'cohere'", () => {
|
||||
const reranker = RerankerFactory.create("cohere", { apiKey: "key" });
|
||||
expect(reranker).toBeInstanceOf(CohereReranker);
|
||||
});
|
||||
|
||||
it("matches the provider name case-insensitively", () => {
|
||||
const reranker = RerankerFactory.create("Cohere", { apiKey: "key" });
|
||||
expect(reranker).toBeInstanceOf(CohereReranker);
|
||||
});
|
||||
|
||||
it("creates a ZeroEntropyReranker for provider 'zero_entropy'", () => {
|
||||
const reranker = RerankerFactory.create("zero_entropy", {
|
||||
apiKey: "key",
|
||||
});
|
||||
expect(reranker).toBeInstanceOf(ZeroEntropyReranker);
|
||||
});
|
||||
|
||||
it("creates a CrossEncoderReranker for provider 'sentence_transformer'", () => {
|
||||
const reranker = RerankerFactory.create("sentence_transformer", {});
|
||||
expect(reranker).toBeInstanceOf(CrossEncoderReranker);
|
||||
});
|
||||
|
||||
it("creates a CrossEncoderReranker for provider 'huggingface'", () => {
|
||||
const reranker = RerankerFactory.create("huggingface", {});
|
||||
expect(reranker).toBeInstanceOf(CrossEncoderReranker);
|
||||
});
|
||||
|
||||
it("creates an LLMReranker for provider 'llm_reranker', building a default openai LLM from top-level config", () => {
|
||||
const reranker = RerankerFactory.create("llm_reranker", {
|
||||
apiKey: "key",
|
||||
});
|
||||
expect(reranker).toBeInstanceOf(LLMReranker);
|
||||
});
|
||||
|
||||
it("creates an LLMReranker that builds its own LLM from a nested config.llm", () => {
|
||||
const reranker = RerankerFactory.create("llm_reranker", {
|
||||
llm: { provider: "openai", config: { apiKey: "x" } },
|
||||
});
|
||||
expect(reranker).toBeInstanceOf(LLMReranker);
|
||||
});
|
||||
|
||||
it("prefers the nested llm.provider over the top-level provider when building the llm_reranker's LLM", () => {
|
||||
// If the top-level `provider` were used instead of the nested one, this
|
||||
// would throw ("Unsupported LLM provider: not-a-real-provider").
|
||||
const reranker = RerankerFactory.create("llm_reranker", {
|
||||
provider: "not-a-real-provider",
|
||||
llm: { provider: "openai", config: { apiKey: "key" } },
|
||||
});
|
||||
expect(reranker).toBeInstanceOf(LLMReranker);
|
||||
});
|
||||
|
||||
it("throws for the 'llm_reranker' provider when the default LLM has no API key available", () => {
|
||||
delete process.env.OPENAI_API_KEY;
|
||||
|
||||
expect(() => RerankerFactory.create("llm_reranker", {})).toThrow();
|
||||
});
|
||||
|
||||
it("throws for an unsupported provider", () => {
|
||||
expect(() => RerankerFactory.create("banana", {})).toThrow(
|
||||
/unsupported reranker provider/i,
|
||||
);
|
||||
});
|
||||
});
|
||||
@@ -12,12 +12,20 @@ import {
|
||||
EmbeddingConfig,
|
||||
HistoryStoreConfig,
|
||||
LLMConfig,
|
||||
RerankerConfig,
|
||||
VectorStoreConfig,
|
||||
} from "../types";
|
||||
import { Reranker } from "../rerankers/base";
|
||||
import { CohereReranker } from "../rerankers/cohere";
|
||||
import { LLMReranker } from "../rerankers/llm";
|
||||
import { ZeroEntropyReranker } from "../rerankers/zeroentropy";
|
||||
import { CrossEncoderReranker } from "../rerankers/cross_encoder";
|
||||
import { Embedder } from "../embeddings/base";
|
||||
import { LLM } from "../llms/base";
|
||||
import { VectorStore } from "../vector_stores/base";
|
||||
import { BaiduDB } from "../vector_stores/baidu";
|
||||
import { Qdrant } from "../vector_stores/qdrant";
|
||||
import { ChromaDB } from "../vector_stores/chroma";
|
||||
import { VectorizeDB } from "../vector_stores/vectorize";
|
||||
import { RedisDB } from "../vector_stores/redis";
|
||||
import { ValkeyDB } from "../vector_stores/valkey";
|
||||
@@ -25,6 +33,8 @@ import { OllamaLLM } from "../llms/ollama";
|
||||
import { LMStudioLLM } from "../llms/lmstudio";
|
||||
import { DeepSeekLLM } from "../llms/deepseek";
|
||||
import { XAILLM } from "../llms/xai";
|
||||
import { SarvamLLM } from "../llms/sarvam";
|
||||
import { AWSBedrockLLM } from "../llms/aws_bedrock";
|
||||
import { LiteLLM } from "../llms/litellm";
|
||||
import { MiniMaxLLM } from "../llms/minimax";
|
||||
import { TogetherLLM } from "../llms/together";
|
||||
@@ -41,9 +51,13 @@ import { AzureOpenAIEmbedder } from "../embeddings/azure";
|
||||
import { FastEmbedEmbedder } from "../embeddings/fastembed";
|
||||
import { LangchainLLM } from "../llms/langchain";
|
||||
import { LangchainEmbedder } from "../embeddings/langchain";
|
||||
import { HuggingFaceEmbedder } from "../embeddings/huggingface";
|
||||
import { LangchainVectorStore } from "../vector_stores/langchain";
|
||||
import { AzureAISearch } from "../vector_stores/azure_ai_search";
|
||||
import { PGVector } from "../vector_stores/pgvector";
|
||||
import { DatabricksVectorStore } from "../vector_stores/databricks";
|
||||
import { NeptuneAnalyticsVectorStore } from "../vector_stores/neptune_analytics";
|
||||
import { VertexAIEmbedder } from "../embeddings/vertexai";
|
||||
import { ElasticsearchDB } from "../vector_stores/elasticsearch";
|
||||
import { OpenSearchDB } from "../vector_stores/opensearch";
|
||||
import { UpstashVector } from "../vector_stores/upstash_vector";
|
||||
@@ -53,7 +67,9 @@ import { CassandraDB } from "../vector_stores/cassandra";
|
||||
import { PineconeDB } from "../vector_stores/pinecone";
|
||||
import { S3Vectors } from "../vector_stores/s3_vectors";
|
||||
import { TurbopufferDB } from "../vector_stores/turbopuffer";
|
||||
import { Milvus } from "../vector_stores/milvus";
|
||||
import { MongoDB } from "../vector_stores/mongodb";
|
||||
import { WeaviateDB } from "../vector_stores/weaviate";
|
||||
|
||||
export class EmbedderFactory {
|
||||
static create(provider: string, config: EmbeddingConfig): Embedder {
|
||||
@@ -75,6 +91,10 @@ export class EmbedderFactory {
|
||||
return new FastEmbedEmbedder(config);
|
||||
case "langchain":
|
||||
return new LangchainEmbedder(config);
|
||||
case "vertexai":
|
||||
return new VertexAIEmbedder(config);
|
||||
case "huggingface":
|
||||
return new HuggingFaceEmbedder(config);
|
||||
default:
|
||||
throw new Error(`Unsupported embedder provider: ${provider}`);
|
||||
}
|
||||
@@ -109,6 +129,10 @@ export class LLMFactory {
|
||||
return new DeepSeekLLM(config);
|
||||
case "xai":
|
||||
return new XAILLM(config);
|
||||
case "sarvam":
|
||||
return new SarvamLLM(config);
|
||||
case "aws_bedrock":
|
||||
return new AWSBedrockLLM(config);
|
||||
case "litellm":
|
||||
return new LiteLLM(config);
|
||||
case "minimax":
|
||||
@@ -128,8 +152,12 @@ export class VectorStoreFactory {
|
||||
switch (provider.toLowerCase()) {
|
||||
case "memory":
|
||||
return new MemoryVectorStore(config);
|
||||
case "baidu":
|
||||
return new BaiduDB(config as any);
|
||||
case "qdrant":
|
||||
return new Qdrant(config as any);
|
||||
case "chroma":
|
||||
return new ChromaDB(config as any);
|
||||
case "redis":
|
||||
return new RedisDB(config as any);
|
||||
case "valkey":
|
||||
@@ -146,6 +174,11 @@ export class VectorStoreFactory {
|
||||
return new VertexAIVectorSearch(config as any);
|
||||
case "pgvector":
|
||||
return new PGVector(config as any);
|
||||
case "databricks":
|
||||
return new DatabricksVectorStore(config as any);
|
||||
case "neptune":
|
||||
case "neptune-analytics":
|
||||
return new NeptuneAnalyticsVectorStore(config as any);
|
||||
case "elasticsearch":
|
||||
return new ElasticsearchDB(config as any);
|
||||
case "opensearch":
|
||||
@@ -163,14 +196,81 @@ export class VectorStoreFactory {
|
||||
return new S3Vectors(config as any);
|
||||
case "turbopuffer":
|
||||
return new TurbopufferDB(config as any);
|
||||
case "milvus":
|
||||
return new Milvus(config as any);
|
||||
case "mongodb":
|
||||
return new MongoDB(config as any);
|
||||
case "weaviate":
|
||||
return new WeaviateDB(config as any);
|
||||
default:
|
||||
throw new Error(`Unsupported vector store provider: ${provider}`);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
export class RerankerFactory {
|
||||
static create(provider: string, config: RerankerConfig): Reranker {
|
||||
switch (provider.toLowerCase()) {
|
||||
case "cohere":
|
||||
return new CohereReranker(config);
|
||||
case "zero_entropy":
|
||||
return new ZeroEntropyReranker(config);
|
||||
case "sentence_transformer":
|
||||
return new CrossEncoderReranker(
|
||||
config,
|
||||
"Xenova/ms-marco-MiniLM-L-6-v2",
|
||||
);
|
||||
case "huggingface":
|
||||
return new CrossEncoderReranker(
|
||||
config,
|
||||
"Xenova/bge-reranker-base",
|
||||
512,
|
||||
);
|
||||
case "llm_reranker": {
|
||||
const llm = RerankerFactory.buildLLMRerankerLLM(config);
|
||||
return new LLMReranker(config, llm);
|
||||
}
|
||||
default:
|
||||
throw new Error(`Unsupported reranker provider: ${provider}`);
|
||||
}
|
||||
}
|
||||
|
||||
private static buildLLMRerankerLLM(config: RerankerConfig): LLM {
|
||||
const nested = config.llm;
|
||||
let llmProvider: string;
|
||||
let llmConfig: LLMConfig;
|
||||
|
||||
if (nested) {
|
||||
llmProvider = nested.provider || config.provider || "openai";
|
||||
llmConfig = { ...(nested.config || {}) };
|
||||
if (llmConfig.model === undefined) {
|
||||
llmConfig.model = config.model ?? "gpt-4o-mini";
|
||||
}
|
||||
if (llmConfig.temperature === undefined) {
|
||||
llmConfig.temperature = config.temperature ?? 0.0;
|
||||
}
|
||||
if (llmConfig.maxTokens === undefined) {
|
||||
llmConfig.maxTokens = config.maxTokens ?? 100;
|
||||
}
|
||||
if (config.apiKey && llmConfig.apiKey === undefined) {
|
||||
llmConfig.apiKey = config.apiKey;
|
||||
}
|
||||
} else {
|
||||
llmProvider = config.provider || "openai";
|
||||
llmConfig = {
|
||||
model: config.model ?? "gpt-4o-mini",
|
||||
temperature: config.temperature ?? 0.0,
|
||||
maxTokens: config.maxTokens ?? 100,
|
||||
};
|
||||
if (config.apiKey) {
|
||||
llmConfig.apiKey = config.apiKey;
|
||||
}
|
||||
}
|
||||
|
||||
return LLMFactory.create(llmProvider, llmConfig);
|
||||
}
|
||||
}
|
||||
|
||||
export class HistoryManagerFactory {
|
||||
static create(provider: string, config: HistoryStoreConfig): HistoryManager {
|
||||
switch (provider.toLowerCase()) {
|
||||
|
||||
@@ -1,4 +1,3 @@
|
||||
import { createPool } from "mysql2/promise";
|
||||
import type { Pool, RowDataPacket } from "mysql2/promise";
|
||||
import { VectorStore } from "./base";
|
||||
import { SearchFilters, VectorStoreConfig, VectorStoreResult } from "../types";
|
||||
@@ -93,6 +92,17 @@ export class AzureMySQLDB implements VectorStore {
|
||||
...(this.config.sslCa ? { ca: this.config.sslCa } : {}),
|
||||
};
|
||||
|
||||
// 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",
|
||||
);
|
||||
}
|
||||
|
||||
this.pool = createPool({
|
||||
host: this.config.host,
|
||||
port: this.config.port ?? 3306,
|
||||
|
||||
@@ -0,0 +1,613 @@
|
||||
import type {
|
||||
AutoBuildIncrementPolicy,
|
||||
CommonResponse,
|
||||
DescTableResponse,
|
||||
FieldType,
|
||||
IndexSchema,
|
||||
MochowClient,
|
||||
QueryResponse,
|
||||
SearchResponse,
|
||||
SelectResponse,
|
||||
TableSchema,
|
||||
} from "@mochow/mochow-sdk-node";
|
||||
import { VectorStore } from "./base";
|
||||
import { SearchFilters, VectorStoreConfig, VectorStoreResult } from "../types";
|
||||
|
||||
type MochowSdk = typeof import("@mochow/mochow-sdk-node");
|
||||
|
||||
export interface BaiduConfig extends VectorStoreConfig {
|
||||
endpoint: string;
|
||||
account: string;
|
||||
apiKey: string;
|
||||
databaseName: string;
|
||||
tableName: string;
|
||||
embeddingModelDims: number;
|
||||
metricType?: "L2" | "IP" | "COSINE";
|
||||
client?: MochowClient;
|
||||
}
|
||||
|
||||
const VECTOR_INDEX = "vector_idx";
|
||||
const FILTERING_INDEX = "metadata_filtering_idx";
|
||||
// Named after the column it actually indexes, and deliberately not Python's "data_bm25_idx".
|
||||
// This index holds Porter-stemmed text, but mem0/vector_stores/baidu.py's keyword_search()
|
||||
// sends a raw, unstemmed query to that name. Sharing it would let Python find an index whose
|
||||
// contents it cannot match properly, silently returning degraded hits instead of None.
|
||||
const BM25_INDEX = "text_lemmatized_bm25_idx";
|
||||
const PROJECTIONS = ["id", "data", "metadata"];
|
||||
const TABLE_POLL_INTERVAL_MS = 2000;
|
||||
const TABLE_POLL_ATTEMPTS = 60;
|
||||
|
||||
// Mochow's server accepts JSON columns, but the Node SDK's FieldType enum predates them
|
||||
// (pymochow 2.4.1 ships FieldType.JSON == "JSON"). The wire value is the bare string.
|
||||
const JSON_FIELD_TYPE = "JSON" as unknown as FieldType;
|
||||
|
||||
// Querying a primary key that isn't there answers with this code, not an empty row. The Node
|
||||
// SDK's ServerErrCode stops at 100, but its siblings against the same server name it:
|
||||
// pymochow ROW_KEY_NOT_FOUND = 101, mochow-sdk-go RowKeyNotFound = 101.
|
||||
const ROW_KEY_NOT_FOUND = 101;
|
||||
|
||||
const SAFE_FILTER_KEY = /^[a-zA-Z_][a-zA-Z0-9_]*$/;
|
||||
|
||||
function escapeFilterString(value: string): string {
|
||||
return value.replace(/\\/g, "\\\\").replace(/"/g, '\\"');
|
||||
}
|
||||
|
||||
function sleep(ms: number): Promise<void> {
|
||||
return new Promise((resolve) => setTimeout(resolve, ms));
|
||||
}
|
||||
|
||||
// Mochow resolves with a {code, msg} envelope instead of rejecting, so every call site
|
||||
// has to inspect the code. `tolerated` lets callers accept the idempotent outcomes
|
||||
// (database/table already exists, table already dropped).
|
||||
function check(
|
||||
response: CommonResponse,
|
||||
action: string,
|
||||
...tolerated: number[]
|
||||
): number {
|
||||
if (response.code !== 0 && !tolerated.includes(response.code)) {
|
||||
throw new Error(
|
||||
`Baidu Mochow ${action} failed (code ${response.code}): ${response.msg}`,
|
||||
);
|
||||
}
|
||||
return response.code;
|
||||
}
|
||||
|
||||
function lemmatizedText(payload: Record<string, any>): string {
|
||||
const data = typeof payload.data === "string" ? payload.data : "";
|
||||
return typeof payload.textLemmatized === "string" &&
|
||||
payload.textLemmatized.length > 0
|
||||
? payload.textLemmatized
|
||||
: data;
|
||||
}
|
||||
|
||||
function memoryData(payload: Record<string, any>): string {
|
||||
return typeof payload.data === "string" ? payload.data : "";
|
||||
}
|
||||
|
||||
function metadataPayload(payload: Record<string, any>): Record<string, any> {
|
||||
const { data: _data, textLemmatized: _textLemmatized, ...metadata } = payload;
|
||||
return metadata;
|
||||
}
|
||||
|
||||
function resultPayload(row: Record<string, any>): Record<string, any> {
|
||||
return {
|
||||
...(row.metadata || {}),
|
||||
...(typeof row.data === "string" ? { data: row.data } : {}),
|
||||
};
|
||||
}
|
||||
|
||||
export class BaiduDB implements VectorStore {
|
||||
private client: MochowClient | null = null;
|
||||
private sdk: MochowSdk | null = null;
|
||||
private readonly endpoint: string;
|
||||
private readonly account: string;
|
||||
private readonly apiKey: string;
|
||||
private readonly databaseName: string;
|
||||
private readonly tableName: string;
|
||||
private readonly embeddingModelDims: number;
|
||||
private readonly metricType: "L2" | "IP" | "COSINE";
|
||||
// Fails closed: keyword search stays off until an inverted index is observed.
|
||||
private supportsKeywordSearch = false;
|
||||
private storeUserId = "anonymous-baidu-user";
|
||||
private _initPromise?: Promise<void>;
|
||||
|
||||
constructor(config: BaiduConfig) {
|
||||
this.endpoint = config.endpoint;
|
||||
this.account = config.account;
|
||||
this.apiKey = config.apiKey;
|
||||
this.databaseName = config.databaseName;
|
||||
this.tableName = config.tableName;
|
||||
this.embeddingModelDims = config.embeddingModelDims;
|
||||
this.metricType = config.metricType || "L2";
|
||||
this.client = config.client || null;
|
||||
|
||||
const requiredFields: Array<
|
||||
readonly [string, string | number | undefined]
|
||||
> = [
|
||||
["databaseName", this.databaseName],
|
||||
["tableName", this.tableName],
|
||||
["embeddingModelDims", this.embeddingModelDims],
|
||||
];
|
||||
|
||||
if (!this.client) {
|
||||
requiredFields.unshift(
|
||||
["endpoint", this.endpoint],
|
||||
["account", this.account],
|
||||
["apiKey", this.apiKey],
|
||||
);
|
||||
}
|
||||
|
||||
for (const [name, value] of requiredFields) {
|
||||
if (value === undefined || value === null || value === "") {
|
||||
throw new Error(
|
||||
`Baidu vector store requires a non-empty '${name}' config value.`,
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
this.initialize().catch(console.error);
|
||||
}
|
||||
|
||||
private get ns(): { database: string; table: string } {
|
||||
return { database: this.databaseName, table: this.tableName };
|
||||
}
|
||||
|
||||
// Loaded dynamically: @mochow/mochow-sdk-node is an optional peer dependency, so a static
|
||||
// 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",
|
||||
);
|
||||
}
|
||||
this.sdk = module.default ?? module;
|
||||
}
|
||||
return this.sdk;
|
||||
}
|
||||
|
||||
private async ensureClient(): Promise<MochowClient> {
|
||||
if (!this.client) {
|
||||
const sdk = await this.loadSdk();
|
||||
this.client = new sdk.MochowClient({
|
||||
endpoint: this.endpoint,
|
||||
credential: { account: this.account, apiKey: this.apiKey },
|
||||
});
|
||||
}
|
||||
return this.client;
|
||||
}
|
||||
|
||||
private async ready(): Promise<{ client: MochowClient; sdk: MochowSdk }> {
|
||||
await this.initialize();
|
||||
return { client: await this.ensureClient(), sdk: await this.loadSdk() };
|
||||
}
|
||||
|
||||
private buildSchema(sdk: MochowSdk): TableSchema {
|
||||
const {
|
||||
AutoBuildPolicyType,
|
||||
FieldType,
|
||||
IndexType,
|
||||
InvertedIndexAnalyzer,
|
||||
InvertedIndexFieldAttribute,
|
||||
InvertedIndexParseMode,
|
||||
MetricType,
|
||||
} = sdk;
|
||||
|
||||
// sdk.AutoBuildIncrement() stamps policyType "TIMING" (bug in 2.1.5), so build the
|
||||
// increment policy by hand.
|
||||
const autoBuildPolicy: AutoBuildIncrementPolicy = {
|
||||
policyType: AutoBuildPolicyType.Increment,
|
||||
rowCountIncrement: 10000,
|
||||
};
|
||||
|
||||
const vectorIndex: IndexSchema = {
|
||||
indexName: VECTOR_INDEX,
|
||||
indexType: IndexType.HNSW,
|
||||
field: "vector",
|
||||
metricType: MetricType[this.metricType],
|
||||
params: { M: 16, efConstruction: 200 },
|
||||
autoBuild: true,
|
||||
autoBuildPolicy,
|
||||
};
|
||||
|
||||
return {
|
||||
fields: [
|
||||
{
|
||||
fieldName: "id",
|
||||
fieldType: FieldType.String,
|
||||
primaryKey: true,
|
||||
partitionKey: true,
|
||||
autoIncrement: false,
|
||||
notNull: true,
|
||||
},
|
||||
{
|
||||
fieldName: "data",
|
||||
fieldType: FieldType.Text,
|
||||
},
|
||||
{
|
||||
fieldName: "vector",
|
||||
fieldType: FieldType.FloatVector,
|
||||
notNull: true,
|
||||
dimension: this.embeddingModelDims,
|
||||
},
|
||||
// Stored outside `metadata` because Mochow cannot build an inverted index on a
|
||||
// field inside a JSON column. Memory.search() passes an already-lemmatized query,
|
||||
// so only the lemmatized form is worth indexing.
|
||||
{ fieldName: "textLemmatized", fieldType: FieldType.Text },
|
||||
{ fieldName: "metadata", fieldType: JSON_FIELD_TYPE },
|
||||
],
|
||||
indexes: [
|
||||
vectorIndex,
|
||||
{
|
||||
indexName: FILTERING_INDEX,
|
||||
indexType: IndexType.FilteringIndex,
|
||||
fields: ["metadata"],
|
||||
},
|
||||
{
|
||||
indexName: BM25_INDEX,
|
||||
indexType: IndexType.InvertedIndex,
|
||||
fields: ["textLemmatized"],
|
||||
fieldAttributes: [InvertedIndexFieldAttribute.Analyzed],
|
||||
params: {
|
||||
analyzer: InvertedIndexAnalyzer.EnglishAnalyzer,
|
||||
parseMode: InvertedIndexParseMode.FineMode,
|
||||
},
|
||||
},
|
||||
],
|
||||
};
|
||||
}
|
||||
|
||||
private buildFilter(filters: SearchFilters): string {
|
||||
const conditions: string[] = [];
|
||||
|
||||
for (const [key, value] of Object.entries(filters)) {
|
||||
if (!SAFE_FILTER_KEY.test(key)) {
|
||||
throw new Error(`Invalid filter key: ${key}`);
|
||||
}
|
||||
|
||||
if (typeof value === "string") {
|
||||
conditions.push(`metadata["${key}"] = "${escapeFilterString(value)}"`);
|
||||
continue;
|
||||
}
|
||||
|
||||
if (typeof value === "number" || typeof value === "boolean") {
|
||||
conditions.push(`metadata["${key}"] = ${value}`);
|
||||
continue;
|
||||
}
|
||||
|
||||
throw new Error(
|
||||
`Filter value for ${key} must be str, int, float, or bool, got ${Array.isArray(value) ? "array" : typeof value}`,
|
||||
);
|
||||
}
|
||||
|
||||
return conditions.join(" AND ");
|
||||
}
|
||||
|
||||
private filterOf(filters?: SearchFilters): string | undefined {
|
||||
return filters && Object.keys(filters).length > 0
|
||||
? this.buildFilter(filters)
|
||||
: undefined;
|
||||
}
|
||||
|
||||
private async pollTable(
|
||||
client: MochowClient,
|
||||
settled: (response: DescTableResponse) => boolean,
|
||||
what: string,
|
||||
): Promise<void> {
|
||||
for (let attempt = 0; attempt < TABLE_POLL_ATTEMPTS; attempt++) {
|
||||
if (settled(await client.descTable(this.databaseName, this.tableName))) {
|
||||
return;
|
||||
}
|
||||
await sleep(TABLE_POLL_INTERVAL_MS);
|
||||
}
|
||||
|
||||
throw new Error(
|
||||
`Baidu Mochow table '${this.tableName}' was not ${what} after ${(TABLE_POLL_ATTEMPTS * TABLE_POLL_INTERVAL_MS) / 1000}s.`,
|
||||
);
|
||||
}
|
||||
|
||||
private async ensureTable(): Promise<void> {
|
||||
const sdk = await this.loadSdk();
|
||||
const client = await this.ensureClient();
|
||||
const { ServerErrCode, TableState } = sdk;
|
||||
|
||||
check(
|
||||
await client.createDatabase(this.databaseName),
|
||||
`createDatabase '${this.databaseName}'`,
|
||||
ServerErrCode.DBAlreadyExist,
|
||||
);
|
||||
|
||||
const created = check(
|
||||
await client.createTable({
|
||||
...this.ns,
|
||||
description: "mem0 memories",
|
||||
replication: 3,
|
||||
partition: { partitionType: sdk.PartitionType.HASH, partitionNum: 1 },
|
||||
enableDynamicField: false,
|
||||
schema: this.buildSchema(sdk),
|
||||
}),
|
||||
`createTable '${this.tableName}'`,
|
||||
ServerErrCode.TableAlreadyExist,
|
||||
);
|
||||
|
||||
// A table is CREATING until its indexes are built; writing to it before then fails.
|
||||
let description: DescTableResponse | undefined;
|
||||
await this.pollTable(
|
||||
client,
|
||||
(response) => {
|
||||
check(response, `descTable '${this.tableName}'`);
|
||||
description = response;
|
||||
return response.table.state === TableState.Normal;
|
||||
},
|
||||
"ready",
|
||||
);
|
||||
|
||||
this.applySchema(
|
||||
created === ServerErrCode.TableAlreadyExist,
|
||||
description!.table.schema,
|
||||
);
|
||||
}
|
||||
|
||||
private applySchema(preexisting: boolean, schema: TableSchema): void {
|
||||
if (!preexisting) {
|
||||
this.supportsKeywordSearch = true;
|
||||
return;
|
||||
}
|
||||
|
||||
const fields = schema?.fields ?? [];
|
||||
const indexes = schema?.indexes ?? [];
|
||||
const field = (name: string) => fields.find((f) => f.fieldName === name);
|
||||
const typeOf = (name: string) => String(field(name)?.fieldType ?? "");
|
||||
const label = `${this.databaseName}.${this.tableName}`;
|
||||
|
||||
if (
|
||||
typeOf("id") !== "STRING" ||
|
||||
!typeOf("data").startsWith("TEXT") ||
|
||||
typeOf("vector") !== "FLOAT_VECTOR" ||
|
||||
typeOf("metadata") !== "JSON"
|
||||
) {
|
||||
throw new Error(
|
||||
`Baidu Mochow table '${label}' exists but is missing the id/data/vector/metadata schema mem0 requires. Drop it, or point 'tableName' at an unused table.`,
|
||||
);
|
||||
}
|
||||
|
||||
const dimension = field("vector")?.dimension;
|
||||
if (dimension !== undefined && dimension !== this.embeddingModelDims) {
|
||||
throw new Error(
|
||||
`Baidu Mochow table '${label}' stores ${dimension}-dimensional vectors, but 'embeddingModelDims' is ${this.embeddingModelDims}.`,
|
||||
);
|
||||
}
|
||||
|
||||
this.supportsKeywordSearch =
|
||||
typeOf("textLemmatized").startsWith("TEXT") &&
|
||||
indexes.some((index) => index.indexName === BM25_INDEX);
|
||||
|
||||
if (!this.supportsKeywordSearch) {
|
||||
console.warn(
|
||||
`Baidu Mochow table '${label}' has no '${BM25_INDEX}' inverted index. keywordSearch() will return null until the table is recreated.`,
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
async initialize(): Promise<void> {
|
||||
if (!this._initPromise) {
|
||||
this._initPromise = this.ensureTable().catch((error) => {
|
||||
this._initPromise = undefined;
|
||||
throw error;
|
||||
});
|
||||
}
|
||||
|
||||
return this._initPromise;
|
||||
}
|
||||
|
||||
async insert(
|
||||
vectors: number[][],
|
||||
ids: string[],
|
||||
payloads: Record<string, any>[],
|
||||
): Promise<void> {
|
||||
const { client } = await this.ready();
|
||||
|
||||
if (vectors.length !== ids.length || vectors.length !== payloads.length) {
|
||||
throw new Error(
|
||||
`Baidu insert requires vectors, ids, and payloads of equal length (got ${vectors.length}/${ids.length}/${payloads.length}).`,
|
||||
);
|
||||
}
|
||||
|
||||
const rows = vectors.map((vector, index) => ({
|
||||
id: ids[index],
|
||||
data: memoryData(payloads[index] || {}),
|
||||
vector,
|
||||
textLemmatized: lemmatizedText(payloads[index] || {}),
|
||||
metadata: metadataPayload(payloads[index] || {}),
|
||||
}));
|
||||
|
||||
check(await client.upsert({ ...this.ns, rows }), "upsert");
|
||||
}
|
||||
|
||||
async search(
|
||||
query: number[],
|
||||
topK = 5,
|
||||
filters?: SearchFilters,
|
||||
): Promise<VectorStoreResult[]> {
|
||||
const { client, sdk } = await this.ready();
|
||||
const filter = this.filterOf(filters);
|
||||
|
||||
const request = new sdk.VectorTopkSearchRequest(
|
||||
"vector",
|
||||
new sdk.Vector(query),
|
||||
topK,
|
||||
)
|
||||
.Projections(PROJECTIONS)
|
||||
.Config(new sdk.VectorSearchConfig().Ef(200));
|
||||
if (filter) {
|
||||
request.Filter(filter);
|
||||
}
|
||||
|
||||
const response = (await client.vectorSearch({
|
||||
...this.ns,
|
||||
request,
|
||||
})) as SearchResponse;
|
||||
check(response, "vectorSearch");
|
||||
|
||||
return (response.rows ?? []).map((result) => ({
|
||||
id: String(result.row.id),
|
||||
payload: resultPayload(result.row),
|
||||
score: result.score,
|
||||
}));
|
||||
}
|
||||
|
||||
async keywordSearch(
|
||||
query: string,
|
||||
topK = 5,
|
||||
filters?: SearchFilters,
|
||||
): Promise<VectorStoreResult[] | null> {
|
||||
const { client, sdk } = await this.ready();
|
||||
|
||||
if (!this.supportsKeywordSearch) {
|
||||
return null;
|
||||
}
|
||||
|
||||
const filter = this.filterOf(filters);
|
||||
const request = new sdk.BM25SearchRequest(BM25_INDEX, query)
|
||||
.Projections(PROJECTIONS)
|
||||
.Limit(topK);
|
||||
if (filter) {
|
||||
request.Filter(filter);
|
||||
}
|
||||
|
||||
const response = (await client.bm25Search({
|
||||
...this.ns,
|
||||
request,
|
||||
})) as SearchResponse;
|
||||
check(response, "bm25Search");
|
||||
|
||||
return (response.rows ?? []).map((result) => ({
|
||||
id: String(result.row.id),
|
||||
payload: resultPayload(result.row),
|
||||
score: result.score,
|
||||
}));
|
||||
}
|
||||
|
||||
async get(vectorId: string): Promise<VectorStoreResult | null> {
|
||||
const { client } = await this.ready();
|
||||
|
||||
const response: QueryResponse = await client.query({
|
||||
...this.ns,
|
||||
primaryKey: { id: vectorId },
|
||||
projections: PROJECTIONS,
|
||||
});
|
||||
check(response, `query '${vectorId}'`, ROW_KEY_NOT_FOUND);
|
||||
|
||||
if (!response.row || response.row.id === undefined) {
|
||||
return null;
|
||||
}
|
||||
|
||||
return {
|
||||
id: String(response.row.id),
|
||||
payload: resultPayload(response.row),
|
||||
};
|
||||
}
|
||||
|
||||
async update(
|
||||
vectorId: string,
|
||||
vector: number[],
|
||||
payload: Record<string, any>,
|
||||
): Promise<void> {
|
||||
const { client } = await this.ready();
|
||||
|
||||
check(
|
||||
await client.upsert({
|
||||
...this.ns,
|
||||
rows: [
|
||||
{
|
||||
id: vectorId,
|
||||
data: memoryData(payload),
|
||||
vector,
|
||||
textLemmatized: lemmatizedText(payload),
|
||||
metadata: metadataPayload(payload),
|
||||
},
|
||||
],
|
||||
}),
|
||||
`upsert '${vectorId}'`,
|
||||
);
|
||||
}
|
||||
|
||||
async delete(vectorId: string): Promise<void> {
|
||||
const { client } = await this.ready();
|
||||
|
||||
check(
|
||||
await client.delete({ ...this.ns, primaryKey: { id: vectorId } }),
|
||||
`delete '${vectorId}'`,
|
||||
);
|
||||
}
|
||||
|
||||
async deleteCol(): Promise<void> {
|
||||
// The constructor starts initialize() without awaiting it. Let any in-flight run land
|
||||
// first, otherwise it recreates the table after dropTable() and reset() is a no-op.
|
||||
await this._initPromise?.catch(() => undefined);
|
||||
this._initPromise = undefined;
|
||||
this.supportsKeywordSearch = false;
|
||||
|
||||
const sdk = await this.loadSdk();
|
||||
const client = await this.ensureClient();
|
||||
const { ServerErrCode } = sdk;
|
||||
|
||||
const dropped = check(
|
||||
await client.dropTable(this.databaseName, this.tableName),
|
||||
`dropTable '${this.tableName}'`,
|
||||
ServerErrCode.TableNotExist,
|
||||
);
|
||||
if (dropped === ServerErrCode.TableNotExist) {
|
||||
return;
|
||||
}
|
||||
|
||||
// Drops are asynchronous; recreating the table before it is gone fails.
|
||||
await this.pollTable(
|
||||
client,
|
||||
(response) =>
|
||||
check(
|
||||
response,
|
||||
`descTable '${this.tableName}'`,
|
||||
ServerErrCode.TableNotExist,
|
||||
) === ServerErrCode.TableNotExist,
|
||||
"dropped",
|
||||
);
|
||||
}
|
||||
|
||||
async reset(): Promise<void> {
|
||||
await this.deleteCol();
|
||||
await this.initialize();
|
||||
}
|
||||
|
||||
async list(
|
||||
filters?: SearchFilters,
|
||||
topK = 100,
|
||||
): Promise<[VectorStoreResult[], number]> {
|
||||
const { client } = await this.ready();
|
||||
|
||||
const response: SelectResponse = await client.select({
|
||||
...this.ns,
|
||||
filter: this.filterOf(filters),
|
||||
projections: PROJECTIONS,
|
||||
limit: topK,
|
||||
});
|
||||
check(response, "select");
|
||||
|
||||
const memories = (response.rows ?? []).map((row) => ({
|
||||
id: String(row.id),
|
||||
payload: resultPayload(row),
|
||||
}));
|
||||
return [memories, memories.length];
|
||||
}
|
||||
|
||||
async getUserId(): Promise<string> {
|
||||
return this.storeUserId;
|
||||
}
|
||||
|
||||
async setUserId(userId: string): Promise<void> {
|
||||
this.storeUserId = userId;
|
||||
}
|
||||
}
|
||||
@@ -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 {
|
||||
|
||||
@@ -0,0 +1,404 @@
|
||||
import type { ChromaClient, CloudClient } from "chromadb";
|
||||
import { VectorStore } from "./base";
|
||||
import { SearchFilters, VectorStoreConfig, VectorStoreResult } from "../types";
|
||||
|
||||
interface ChromaConfig extends VectorStoreConfig {
|
||||
/** Pre-configured ChromaDB client instance. */
|
||||
client?: ChromaClient | CloudClient;
|
||||
collectionName: string;
|
||||
/** Host address for a ChromaDB server (defaults to the client default). */
|
||||
host?: string;
|
||||
/** Port for a ChromaDB server. */
|
||||
port?: number;
|
||||
/** Whether to use SSL when connecting to a ChromaDB server. */
|
||||
ssl?: boolean;
|
||||
/** Path for a local ChromaDB server. */
|
||||
path?: string;
|
||||
/** ChromaDB Cloud API key. */
|
||||
apiKey?: string;
|
||||
/** ChromaDB Cloud tenant ID. */
|
||||
tenant?: string;
|
||||
/** ChromaDB Cloud database name. */
|
||||
database?: string;
|
||||
}
|
||||
|
||||
const MIGRATIONS_COLLECTION = "memory_migrations";
|
||||
|
||||
/**
|
||||
* ChromaDB vector store provider.
|
||||
*
|
||||
* Mirrors the Python SDK's `mem0.vector_stores.chroma.ChromaDB` behavior using
|
||||
* the `chromadb` v3 JavaScript client. Embeddings are always supplied by mem0,
|
||||
* so no embedding function is required on the collection.
|
||||
*/
|
||||
export class ChromaDB implements VectorStore {
|
||||
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) {
|
||||
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) {
|
||||
const client = await this.getClient();
|
||||
this.collectionPromise = client.getOrCreateCollection({
|
||||
name: this.collectionName,
|
||||
embeddingFunction: null,
|
||||
});
|
||||
}
|
||||
return this.collectionPromise;
|
||||
}
|
||||
|
||||
private async getMigrationsCollection(): Promise<any> {
|
||||
if (!this.migrationsPromise) {
|
||||
const client = await this.getClient();
|
||||
this.migrationsPromise = client.getOrCreateCollection({
|
||||
name: MIGRATIONS_COLLECTION,
|
||||
embeddingFunction: null,
|
||||
});
|
||||
}
|
||||
return this.migrationsPromise;
|
||||
}
|
||||
|
||||
private flatten(value: any): any[] {
|
||||
if (Array.isArray(value) && value.length > 0 && Array.isArray(value[0])) {
|
||||
return value[0];
|
||||
}
|
||||
return Array.isArray(value) ? value : [];
|
||||
}
|
||||
|
||||
/** Parse a ChromaDB `get`/`query` response into VectorStoreResult objects. */
|
||||
private parseOutput(data: any): VectorStoreResult[] {
|
||||
const ids = this.flatten(data?.ids);
|
||||
const distances = this.flatten(data?.distances);
|
||||
const metadatas = this.flatten(data?.metadatas);
|
||||
|
||||
const length = Math.max(ids.length, metadatas.length);
|
||||
const results: VectorStoreResult[] = [];
|
||||
|
||||
for (let i = 0; i < length; i++) {
|
||||
const rawDistance = distances[i];
|
||||
const score =
|
||||
rawDistance !== undefined && rawDistance !== null
|
||||
? 1.0 / (1.0 + rawDistance)
|
||||
: undefined;
|
||||
|
||||
results.push({
|
||||
id: String(ids[i]),
|
||||
payload: (metadatas[i] as Record<string, any>) || {},
|
||||
score,
|
||||
});
|
||||
}
|
||||
|
||||
return results;
|
||||
}
|
||||
|
||||
async insert(
|
||||
vectors: number[][],
|
||||
ids: string[],
|
||||
payloads: Record<string, any>[],
|
||||
): Promise<void> {
|
||||
const collection = await this.getCollection();
|
||||
await collection.add({
|
||||
ids,
|
||||
embeddings: vectors,
|
||||
metadatas: payloads as any,
|
||||
});
|
||||
}
|
||||
|
||||
async keywordSearch(): Promise<null> {
|
||||
return null;
|
||||
}
|
||||
|
||||
async search(
|
||||
query: number[],
|
||||
topK: number = 5,
|
||||
filters?: SearchFilters,
|
||||
): Promise<VectorStoreResult[]> {
|
||||
const collection = await this.getCollection();
|
||||
const where = ChromaDB.generateWhereClause(filters);
|
||||
const results = await collection.query({
|
||||
queryEmbeddings: [query],
|
||||
nResults: topK,
|
||||
where: where as any,
|
||||
});
|
||||
return this.parseOutput(results);
|
||||
}
|
||||
|
||||
async get(vectorId: string): Promise<VectorStoreResult | null> {
|
||||
const collection = await this.getCollection();
|
||||
const results = await collection.get({ ids: [vectorId] });
|
||||
const parsed = this.parseOutput(results);
|
||||
return parsed.length > 0 ? parsed[0] : null;
|
||||
}
|
||||
|
||||
async update(
|
||||
vectorId: string,
|
||||
vector: number[],
|
||||
payload: Record<string, any>,
|
||||
): Promise<void> {
|
||||
const collection = await this.getCollection();
|
||||
await collection.update({
|
||||
ids: [vectorId],
|
||||
embeddings: vector ? [vector] : undefined,
|
||||
metadatas: payload ? [payload] : undefined,
|
||||
} as any);
|
||||
}
|
||||
|
||||
async delete(vectorId: string): Promise<void> {
|
||||
const collection = await this.getCollection();
|
||||
await collection.delete({ ids: [vectorId] });
|
||||
}
|
||||
|
||||
async deleteCol(): Promise<void> {
|
||||
const client = await this.getClient();
|
||||
await client.deleteCollection({ name: this.collectionName });
|
||||
this.collectionPromise = undefined;
|
||||
}
|
||||
|
||||
async list(
|
||||
filters?: SearchFilters,
|
||||
topK: number = 100,
|
||||
): Promise<[VectorStoreResult[], number]> {
|
||||
const collection = await this.getCollection();
|
||||
const where = ChromaDB.generateWhereClause(filters);
|
||||
const results = await collection.get({ where: where as any, limit: topK });
|
||||
const parsed = this.parseOutput(results);
|
||||
return [parsed, parsed.length];
|
||||
}
|
||||
|
||||
private generateUUID(): string {
|
||||
return "xxxxxxxx-xxxx-4xxx-yxxx-xxxxxxxxxxxx".replace(
|
||||
/[xy]/g,
|
||||
function (c) {
|
||||
const r = (Math.random() * 16) | 0;
|
||||
const v = c === "x" ? r : (r & 0x3) | 0x8;
|
||||
return v.toString(16);
|
||||
},
|
||||
);
|
||||
}
|
||||
|
||||
async getUserId(): Promise<string> {
|
||||
const collection = await this.getMigrationsCollection();
|
||||
const result = await collection.get({ limit: 1 });
|
||||
const ids = Array.isArray(result?.ids) ? result.ids : [];
|
||||
const metadatas = Array.isArray(result?.metadatas) ? result.metadatas : [];
|
||||
|
||||
if (ids.length > 0 && metadatas[0]?.user_id) {
|
||||
return String(metadatas[0].user_id);
|
||||
}
|
||||
|
||||
const randomUserId =
|
||||
Math.random().toString(36).substring(2, 15) +
|
||||
Math.random().toString(36).substring(2, 15);
|
||||
|
||||
await collection.add({
|
||||
ids: [this.generateUUID()],
|
||||
embeddings: [[0]],
|
||||
metadatas: [{ user_id: randomUserId }] as any,
|
||||
});
|
||||
|
||||
return randomUserId;
|
||||
}
|
||||
|
||||
async setUserId(userId: string): Promise<void> {
|
||||
const collection = await this.getMigrationsCollection();
|
||||
const result = await collection.get({ limit: 1 });
|
||||
const ids = Array.isArray(result?.ids) ? result.ids : [];
|
||||
const pointId = ids.length > 0 ? String(ids[0]) : this.generateUUID();
|
||||
|
||||
await collection.upsert({
|
||||
ids: [pointId],
|
||||
embeddings: [[0]],
|
||||
metadatas: [{ user_id: userId }] as any,
|
||||
});
|
||||
}
|
||||
|
||||
async initialize(): Promise<void> {
|
||||
await this.getCollection();
|
||||
await this.getMigrationsCollection();
|
||||
}
|
||||
|
||||
/** Convert a single field filter into a ChromaDB where condition. */
|
||||
private static convertCondition(
|
||||
key: string,
|
||||
value: any,
|
||||
): Record<string, any> | null {
|
||||
// Wildcard - ChromaDB has no direct wildcard, so skip this filter.
|
||||
if (value === "*") {
|
||||
return null;
|
||||
}
|
||||
|
||||
if (Array.isArray(value)) {
|
||||
return { [key]: { $in: value } };
|
||||
}
|
||||
|
||||
if (value !== null && typeof value === "object") {
|
||||
const opMap: Record<string, string> = {
|
||||
eq: "$eq",
|
||||
ne: "$ne",
|
||||
gt: "$gt",
|
||||
gte: "$gte",
|
||||
lt: "$lt",
|
||||
lte: "$lte",
|
||||
in: "$in",
|
||||
nin: "$nin",
|
||||
};
|
||||
const condition: Record<string, any> = {};
|
||||
for (const [op, val] of Object.entries(value)) {
|
||||
if (op in opMap) {
|
||||
condition[key] = { [opMap[op]]: val };
|
||||
} else {
|
||||
// contains/icontains and unknown operators fall back to equality.
|
||||
condition[key] = { $eq: val };
|
||||
}
|
||||
}
|
||||
return condition;
|
||||
}
|
||||
|
||||
return { [key]: { $eq: value } };
|
||||
}
|
||||
|
||||
/**
|
||||
* Generate a properly formatted `where` clause for ChromaDB from mem0's
|
||||
* universal filter format. Supports comparison operators plus $or/$not.
|
||||
*/
|
||||
static generateWhereClause(
|
||||
filters?: SearchFilters,
|
||||
): Record<string, any> | undefined {
|
||||
if (!filters) {
|
||||
return undefined;
|
||||
}
|
||||
|
||||
const negateOp: Record<string, string> = {
|
||||
eq: "ne",
|
||||
ne: "eq",
|
||||
gt: "lte",
|
||||
gte: "lt",
|
||||
lt: "gte",
|
||||
lte: "gt",
|
||||
in: "nin",
|
||||
nin: "in",
|
||||
};
|
||||
|
||||
const processed: any[] = [];
|
||||
|
||||
for (const [key, value] of Object.entries(filters)) {
|
||||
if (key === "$or" || key === "OR") {
|
||||
const orConditions: any[] = [];
|
||||
for (const condition of value as any[]) {
|
||||
const built: Record<string, any> = {};
|
||||
for (const [subKey, subValue] of Object.entries(condition)) {
|
||||
const converted = ChromaDB.convertCondition(subKey, subValue);
|
||||
if (converted) Object.assign(built, converted);
|
||||
}
|
||||
if (Object.keys(built).length > 0) orConditions.push(built);
|
||||
}
|
||||
if (orConditions.length > 1) {
|
||||
processed.push({ $or: orConditions });
|
||||
} else if (orConditions.length === 1) {
|
||||
processed.push(orConditions[0]);
|
||||
}
|
||||
} else if (key === "$not" || key === "NOT") {
|
||||
// De Morgan: NOT(a AND b) is (NOT a) OR (NOT b), so the negated fields
|
||||
// within one condition are combined with $or, and separate conditions
|
||||
// are combined with $and. This mirrors the Python SDK's ChromaDB port.
|
||||
const negatedPerGroup: any[] = [];
|
||||
for (const condition of value as any[]) {
|
||||
const negatedFields: any[] = [];
|
||||
for (const [subKey, subValue] of Object.entries(condition)) {
|
||||
if (subValue !== null && typeof subValue === "object") {
|
||||
for (const [op, val] of Object.entries(subValue as any)) {
|
||||
const neg = negateOp[op];
|
||||
if (neg) {
|
||||
const converted = ChromaDB.convertCondition(subKey, {
|
||||
[neg]: val,
|
||||
});
|
||||
if (converted) negatedFields.push(converted);
|
||||
}
|
||||
}
|
||||
} else {
|
||||
const converted = ChromaDB.convertCondition(subKey, {
|
||||
ne: subValue,
|
||||
});
|
||||
if (converted) negatedFields.push(converted);
|
||||
}
|
||||
}
|
||||
if (negatedFields.length > 1) {
|
||||
negatedPerGroup.push({ $or: negatedFields });
|
||||
} else if (negatedFields.length === 1) {
|
||||
negatedPerGroup.push(negatedFields[0]);
|
||||
}
|
||||
}
|
||||
if (negatedPerGroup.length > 1) {
|
||||
processed.push({ $and: negatedPerGroup });
|
||||
} else if (negatedPerGroup.length === 1) {
|
||||
processed.push(negatedPerGroup[0]);
|
||||
}
|
||||
} else {
|
||||
const converted = ChromaDB.convertCondition(key, value);
|
||||
if (converted) processed.push(converted);
|
||||
}
|
||||
}
|
||||
|
||||
if (processed.length === 0) {
|
||||
return undefined;
|
||||
}
|
||||
if (processed.length === 1) {
|
||||
return processed[0];
|
||||
}
|
||||
return { $and: processed };
|
||||
}
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
@@ -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",
|
||||
|
||||
@@ -0,0 +1,510 @@
|
||||
import { VectorStore } from "./base";
|
||||
import { SearchFilters, VectorStoreConfig, VectorStoreResult } from "../types";
|
||||
|
||||
/**
|
||||
* Supported Milvus metric types. Mirrors the Python provider
|
||||
* (`mem0/configs/vector_stores/milvus.py`).
|
||||
*/
|
||||
export type MilvusMetricType = "L2" | "IP" | "COSINE" | "HAMMING" | "JACCARD";
|
||||
|
||||
export interface MilvusConfig extends VectorStoreConfig {
|
||||
/**
|
||||
* Full URL/address for the Milvus or Zilliz server.
|
||||
* Defaults to `http://localhost:19530`.
|
||||
*/
|
||||
url?: string;
|
||||
/** Token / API key for Zilliz Cloud. Optional for a local setup. */
|
||||
token?: string;
|
||||
/** Name of the database. Optional (Milvus default database when empty). */
|
||||
dbName?: string;
|
||||
/** Collection name. Defaults to `mem0`. */
|
||||
collectionName?: string;
|
||||
/** Embedding dimensionality. Defaults to 1536 (OpenAI). */
|
||||
embeddingModelDims?: number;
|
||||
dimension?: number;
|
||||
/** Similarity metric. Defaults to `L2` (matches the Python provider). */
|
||||
metricType?: MilvusMetricType;
|
||||
/**
|
||||
* Pre-constructed `MilvusClient` instance. When provided, `url`/`token`/`dbName`
|
||||
* are ignored. Primarily useful for dependency injection in tests.
|
||||
*/
|
||||
client?: any;
|
||||
}
|
||||
|
||||
/**
|
||||
* Milvus vector store provider for the TypeScript OSS SDK.
|
||||
*
|
||||
* Mirrors the Python provider in `mem0/vector_stores/milvus.py`: dense-vector
|
||||
* CRUD (insert / search / get / update / delete / list + user-id helpers) plus
|
||||
* BM25 hybrid keyword search. New collections are created with a `text` +
|
||||
* `sparse` field pair and a BM25 function so `keywordSearch` can run full-text
|
||||
* search; collections created before BM25 support keep working with it disabled.
|
||||
*
|
||||
* The `@zilliz/milvus2-sdk-node` dependency is lazily required so the package
|
||||
* remains optional. Importing this module never forces the SDK to be installed
|
||||
* until a Milvus store is actually constructed.
|
||||
*/
|
||||
export class Milvus implements VectorStore {
|
||||
private client: any;
|
||||
private readonly collectionName: string;
|
||||
private readonly dimension: number;
|
||||
private readonly metricType: MilvusMetricType;
|
||||
private _initPromise?: Promise<void>;
|
||||
// Lazily-resolved SDK enums (set during client construction).
|
||||
private DataType: any;
|
||||
private FunctionType: any;
|
||||
// Whether this collection has the `text` + `sparse` fields for BM25 hybrid
|
||||
// search. Collections created before BM25 support lack them, so writing a
|
||||
// top-level `text` field or passing a sparse anns_field would be rejected.
|
||||
private hasBm25Schema = false;
|
||||
|
||||
constructor(config: MilvusConfig) {
|
||||
this.collectionName = config.collectionName || "mem0";
|
||||
this.dimension = config.embeddingModelDims || config.dimension || 1536;
|
||||
this.metricType = config.metricType || "L2";
|
||||
|
||||
if (config.client) {
|
||||
this.client = config.client;
|
||||
// Best-effort SDK enum resolution for an injected client (undefined when
|
||||
// the SDK isn't installed, e.g. unit tests that pass a fake client).
|
||||
try {
|
||||
// eslint-disable-next-line @typescript-eslint/no-var-requires
|
||||
const sdk = require("@zilliz/milvus2-sdk-node");
|
||||
this.DataType = sdk.DataType;
|
||||
this.FunctionType = sdk.FunctionType;
|
||||
} catch (_) {
|
||||
this.DataType = undefined;
|
||||
this.FunctionType = undefined;
|
||||
}
|
||||
} else {
|
||||
let MilvusClient: any;
|
||||
let DataType: any;
|
||||
let FunctionType: any;
|
||||
try {
|
||||
// eslint-disable-next-line @typescript-eslint/no-var-requires
|
||||
const sdk = require("@zilliz/milvus2-sdk-node");
|
||||
MilvusClient = sdk.MilvusClient;
|
||||
DataType = sdk.DataType;
|
||||
FunctionType = sdk.FunctionType;
|
||||
} catch (_) {
|
||||
throw new Error(
|
||||
"The '@zilliz/milvus2-sdk-node' package is required to use the Milvus vector store. " +
|
||||
"Install it with: npm install @zilliz/milvus2-sdk-node",
|
||||
);
|
||||
}
|
||||
this.DataType = DataType;
|
||||
this.FunctionType = FunctionType;
|
||||
this.client = new MilvusClient({
|
||||
address: config.url || "http://localhost:19530",
|
||||
token: config.token,
|
||||
database: config.dbName || undefined,
|
||||
});
|
||||
}
|
||||
|
||||
this.initialize().catch((err) =>
|
||||
console.error("Error initializing Milvus:", err),
|
||||
);
|
||||
}
|
||||
|
||||
async initialize(): Promise<void> {
|
||||
if (!this._initPromise) {
|
||||
this._initPromise = this.createCol(this.collectionName, this.dimension);
|
||||
}
|
||||
return this._initPromise;
|
||||
}
|
||||
|
||||
/**
|
||||
* Create the collection if it does not already exist, with an AUTOINDEX
|
||||
* dense-vector index plus a BM25 `text` -> `sparse` full-text index. Idempotent
|
||||
* (mirrors the Python `create_col`). When the collection already exists, detect
|
||||
* whether the BM25 `text`/`sparse` fields are present so keyword search
|
||||
* degrades gracefully on collections created before BM25 support.
|
||||
*/
|
||||
private async createCol(
|
||||
collectionName: string,
|
||||
vectorSize: number,
|
||||
): Promise<void> {
|
||||
const has = await this.client.hasCollection({
|
||||
collection_name: collectionName,
|
||||
});
|
||||
// milvus2-sdk-node returns { value: boolean } for hasCollection.
|
||||
const exists = typeof has === "object" && has !== null ? has.value : has;
|
||||
if (exists) {
|
||||
// A pre-existing collection may predate BM25 support. Inspect its schema
|
||||
// so insert/search/keywordSearch know whether the text + sparse fields
|
||||
// exist instead of assuming and getting rejected by the server.
|
||||
const desc = await this.client.describeCollection({
|
||||
collection_name: collectionName,
|
||||
});
|
||||
const names = new Set(
|
||||
(desc?.schema?.fields || []).map((f: any) => f.name),
|
||||
);
|
||||
this.hasBm25Schema = names.has("text") && names.has("sparse");
|
||||
if (!this.hasBm25Schema) {
|
||||
console.warn(
|
||||
`Milvus collection '${collectionName}' predates BM25 hybrid search ` +
|
||||
"(no 'text'/'sparse' fields). Keyword scoring is disabled for it; " +
|
||||
"semantic search still works. Use a fresh collection to enable it.",
|
||||
);
|
||||
}
|
||||
return;
|
||||
}
|
||||
|
||||
const DataType = this.DataType || {};
|
||||
const fields = [
|
||||
{
|
||||
name: "id",
|
||||
data_type: DataType.VarChar,
|
||||
is_primary_key: true,
|
||||
max_length: 512,
|
||||
},
|
||||
{
|
||||
name: "vectors",
|
||||
data_type: DataType.FloatVector,
|
||||
dim: vectorSize,
|
||||
},
|
||||
{
|
||||
name: "metadata",
|
||||
data_type: DataType.JSON,
|
||||
},
|
||||
// Analyzer-enabled text field that feeds the BM25 function below.
|
||||
{
|
||||
name: "text",
|
||||
data_type: DataType.VarChar,
|
||||
max_length: 65535,
|
||||
enable_analyzer: true,
|
||||
},
|
||||
// Sparse vectors are generated automatically by the BM25 function.
|
||||
{
|
||||
name: "sparse",
|
||||
data_type: DataType.SparseFloatVector,
|
||||
},
|
||||
];
|
||||
|
||||
await this.client.createCollection({
|
||||
collection_name: collectionName,
|
||||
fields,
|
||||
enable_dynamic_field: true,
|
||||
// BM25 turns the `text` field into `sparse` vectors for full-text search.
|
||||
functions: [
|
||||
{
|
||||
name: "bm25",
|
||||
type: this.FunctionType?.BM25,
|
||||
input_field_names: ["text"],
|
||||
output_field_names: ["sparse"],
|
||||
params: {},
|
||||
},
|
||||
],
|
||||
index_params: [
|
||||
{
|
||||
field_name: "vectors",
|
||||
index_type: "AUTOINDEX",
|
||||
metric_type: this.metricType,
|
||||
index_name: "vector_index",
|
||||
},
|
||||
{
|
||||
field_name: "sparse",
|
||||
index_type: "SPARSE_INVERTED_INDEX",
|
||||
metric_type: "BM25",
|
||||
index_name: "sparse_index",
|
||||
},
|
||||
],
|
||||
});
|
||||
|
||||
await this.client.loadCollection({ collection_name: collectionName });
|
||||
this.hasBm25Schema = true;
|
||||
}
|
||||
|
||||
/**
|
||||
* Filter keys are interpolated straight into the expression, so restrict them
|
||||
* to safe identifiers (same rule as the Python provider) to block injection.
|
||||
*/
|
||||
private static readonly SAFE_FILTER_KEY = /^[a-zA-Z_][a-zA-Z0-9_]*$/;
|
||||
|
||||
/**
|
||||
* Build a Milvus boolean filter expression from a flat filters object.
|
||||
* Mirrors the Python `_create_filter` (equality only, AND-combined): validate
|
||||
* each key, escape string values (backslash first, then double-quote), and
|
||||
* reject value types Milvus can't compare against a scalar field.
|
||||
*/
|
||||
private createFilter(filters?: SearchFilters): string | undefined {
|
||||
if (!filters || Object.keys(filters).length === 0) return undefined;
|
||||
const operands: string[] = [];
|
||||
for (const [key, value] of Object.entries(filters)) {
|
||||
if (value === undefined || value === null) continue;
|
||||
if (!Milvus.SAFE_FILTER_KEY.test(key)) {
|
||||
throw new Error(`Invalid filter key: ${JSON.stringify(key)}`);
|
||||
}
|
||||
if (typeof value === "string") {
|
||||
// Escape backslashes before quotes so a value can't break out of the
|
||||
// string literal (order matters, exactly as in the Python provider).
|
||||
const escaped = value.replace(/\\/g, "\\\\").replace(/"/g, '\\"');
|
||||
operands.push(`(metadata["${key}"] == "${escaped}")`);
|
||||
} else if (typeof value === "number" || typeof value === "boolean") {
|
||||
operands.push(`(metadata["${key}"] == ${value})`);
|
||||
} else {
|
||||
throw new Error(
|
||||
`Filter value for ${JSON.stringify(key)} must be a string, number, or boolean, got ${typeof value}`,
|
||||
);
|
||||
}
|
||||
}
|
||||
return operands.length > 0 ? operands.join(" and ") : undefined;
|
||||
}
|
||||
|
||||
/**
|
||||
* Text fed to the BM25 sparse index for a payload. Prefers the lemmatized
|
||||
* text, falls back to the raw memory `data`, and truncates to the VarChar
|
||||
* limit (mirrors the Python provider).
|
||||
*/
|
||||
private bm25Text(payload?: Record<string, any>): string {
|
||||
if (!payload) return "";
|
||||
const raw = payload.text_lemmatized || payload.data || "";
|
||||
return String(raw).slice(0, 65535);
|
||||
}
|
||||
|
||||
async insert(
|
||||
vectors: number[][],
|
||||
ids: string[],
|
||||
payloads: Record<string, any>[],
|
||||
): Promise<void> {
|
||||
const data = vectors.map((vector, idx) => {
|
||||
const metadata = payloads[idx] || {};
|
||||
const row: Record<string, any> = {
|
||||
id: ids[idx],
|
||||
vectors: vector,
|
||||
metadata,
|
||||
};
|
||||
// Only write `text` when the collection has the BM25 schema; legacy
|
||||
// collections reject an unknown top-level field.
|
||||
if (this.hasBm25Schema) row.text = this.bm25Text(metadata);
|
||||
return row;
|
||||
});
|
||||
await this.client.insert({
|
||||
collection_name: this.collectionName,
|
||||
data,
|
||||
});
|
||||
}
|
||||
|
||||
/**
|
||||
* Map raw Milvus search hits to VectorStoreResult. For the L2 metric,
|
||||
* distances are unbounded and smaller-is-better, so normalise them to a 0..1
|
||||
* similarity; every other metric passes the raw score through. Shared by
|
||||
* search and keywordSearch so both match the Python provider's `_parse_output`
|
||||
* (which likewise normalises by metric type regardless of dense vs BM25 score).
|
||||
*/
|
||||
private parseHits(hits: any[]): VectorStoreResult[] {
|
||||
return hits.map((hit: any) => {
|
||||
const rawDistance = hit.score ?? hit.distance;
|
||||
let score = rawDistance;
|
||||
if (rawDistance != null && this.metricType === "L2") {
|
||||
score = 1.0 / (1.0 + rawDistance);
|
||||
}
|
||||
return {
|
||||
id: String(hit.id),
|
||||
payload: hit.metadata || {},
|
||||
score,
|
||||
} as VectorStoreResult;
|
||||
});
|
||||
}
|
||||
|
||||
async search(
|
||||
query: number[],
|
||||
topK: number = 5,
|
||||
filters?: SearchFilters,
|
||||
): Promise<VectorStoreResult[]> {
|
||||
const filter = this.createFilter(filters);
|
||||
const req: Record<string, any> = {
|
||||
collection_name: this.collectionName,
|
||||
data: [query],
|
||||
limit: topK,
|
||||
filter,
|
||||
output_fields: ["*"],
|
||||
};
|
||||
// A BM25 collection has both a dense `vectors` and a sparse `sparse` field,
|
||||
// so anns_field is ambiguous unless we name the dense one explicitly.
|
||||
if (this.hasBm25Schema) req.anns_field = "vectors";
|
||||
const res = await this.client.search(req);
|
||||
return this.parseHits(res?.results || []);
|
||||
}
|
||||
|
||||
/**
|
||||
* BM25 full-text keyword search over the sparse field. Milvus tokenizes the
|
||||
* raw query string via the collection's BM25 function. Returns null when the
|
||||
* collection has no BM25 schema so callers fall back to dense search only
|
||||
* (mirrors the Python provider).
|
||||
*/
|
||||
async keywordSearch(
|
||||
query: string,
|
||||
topK: number = 5,
|
||||
filters?: SearchFilters,
|
||||
): Promise<VectorStoreResult[] | null> {
|
||||
if (!this.hasBm25Schema) return null;
|
||||
try {
|
||||
const filter = this.createFilter(filters);
|
||||
const res = await this.client.search({
|
||||
collection_name: this.collectionName,
|
||||
data: [query],
|
||||
anns_field: "sparse",
|
||||
limit: topK,
|
||||
filter,
|
||||
output_fields: ["*"],
|
||||
});
|
||||
return this.parseHits(res?.results || []);
|
||||
} catch (_) {
|
||||
// Keyword search is best-effort; degrade to null rather than failing the
|
||||
// whole retrieval path.
|
||||
return null;
|
||||
}
|
||||
}
|
||||
|
||||
async get(vectorId: string): Promise<VectorStoreResult | null> {
|
||||
const res = await this.client.get({
|
||||
collection_name: this.collectionName,
|
||||
ids: [vectorId],
|
||||
output_fields: ["id", "metadata"],
|
||||
});
|
||||
const rows = res?.data || [];
|
||||
if (!rows.length) return null;
|
||||
return {
|
||||
id: String(rows[0].id),
|
||||
payload: rows[0].metadata || {},
|
||||
};
|
||||
}
|
||||
|
||||
async update(
|
||||
vectorId: string,
|
||||
vector: number[],
|
||||
payload: Record<string, any>,
|
||||
): Promise<void> {
|
||||
const row: Record<string, any> = {
|
||||
id: vectorId,
|
||||
vectors: vector,
|
||||
metadata: payload,
|
||||
};
|
||||
if (this.hasBm25Schema) row.text = this.bm25Text(payload);
|
||||
await this.client.upsert({
|
||||
collection_name: this.collectionName,
|
||||
data: [row],
|
||||
});
|
||||
}
|
||||
|
||||
async delete(vectorId: string): Promise<void> {
|
||||
await this.client.delete({
|
||||
collection_name: this.collectionName,
|
||||
ids: [vectorId],
|
||||
});
|
||||
}
|
||||
|
||||
async deleteCol(): Promise<void> {
|
||||
await this.client.dropCollection({
|
||||
collection_name: this.collectionName,
|
||||
});
|
||||
}
|
||||
|
||||
async list(
|
||||
filters?: SearchFilters,
|
||||
topK: number = 100,
|
||||
): Promise<[VectorStoreResult[], number]> {
|
||||
const filter = this.createFilter(filters);
|
||||
const res = await this.client.query({
|
||||
collection_name: this.collectionName,
|
||||
filter: filter ?? "",
|
||||
limit: topK,
|
||||
output_fields: ["id", "metadata"],
|
||||
});
|
||||
const rows = res?.data || [];
|
||||
const results: VectorStoreResult[] = rows.map((row: any) => ({
|
||||
id: String(row.id),
|
||||
payload: row.metadata || {},
|
||||
}));
|
||||
return [results, results.length];
|
||||
}
|
||||
|
||||
private generateUUID(): string {
|
||||
return "xxxxxxxx-xxxx-4xxx-yxxx-xxxxxxxxxxxx".replace(/[xy]/g, (c) => {
|
||||
const r = (Math.random() * 16) | 0;
|
||||
const v = c === "x" ? r : (r & 0x3) | 0x8;
|
||||
return v.toString(16);
|
||||
});
|
||||
}
|
||||
|
||||
private async ensureMigrationsCol(): Promise<void> {
|
||||
const name = "memory_migrations";
|
||||
const has = await this.client.hasCollection({ collection_name: name });
|
||||
const exists = typeof has === "object" && has !== null ? has.value : has;
|
||||
if (exists) return;
|
||||
|
||||
const DataType = this.DataType || {};
|
||||
await this.client.createCollection({
|
||||
collection_name: name,
|
||||
fields: [
|
||||
{
|
||||
name: "id",
|
||||
data_type: DataType.VarChar,
|
||||
is_primary_key: true,
|
||||
max_length: 512,
|
||||
},
|
||||
// Milvus rejects a FloatVector with dim < 2, so use the minimum (2). This
|
||||
// helper collection is never vector-searched; the vector is a fixed
|
||||
// placeholder that exists only to satisfy the required-vector-field schema.
|
||||
{ name: "vectors", data_type: DataType.FloatVector, dim: 2 },
|
||||
{ name: "user_id", data_type: DataType.VarChar, max_length: 512 },
|
||||
],
|
||||
index_params: [
|
||||
{
|
||||
field_name: "vectors",
|
||||
index_type: "AUTOINDEX",
|
||||
// Fixed metric: this helper collection is never vector-searched, and a
|
||||
// zero vector has no direction, so COSINE would be degenerate here.
|
||||
metric_type: "L2",
|
||||
index_name: "vector_index",
|
||||
},
|
||||
],
|
||||
});
|
||||
await this.client.loadCollection({ collection_name: name });
|
||||
}
|
||||
|
||||
async getUserId(): Promise<string> {
|
||||
await this.ensureMigrationsCol();
|
||||
const res = await this.client.query({
|
||||
collection_name: "memory_migrations",
|
||||
filter: "",
|
||||
limit: 1,
|
||||
output_fields: ["*"],
|
||||
});
|
||||
const rows = res?.data || [];
|
||||
if (rows.length > 0 && rows[0].user_id) {
|
||||
return String(rows[0].user_id);
|
||||
}
|
||||
const randomUserId =
|
||||
Math.random().toString(36).substring(2, 15) +
|
||||
Math.random().toString(36).substring(2, 15);
|
||||
await this.client.insert({
|
||||
collection_name: "memory_migrations",
|
||||
data: [
|
||||
{ id: this.generateUUID(), vectors: [0, 0], user_id: randomUserId },
|
||||
],
|
||||
});
|
||||
return randomUserId;
|
||||
}
|
||||
|
||||
async setUserId(userId: string): Promise<void> {
|
||||
await this.ensureMigrationsCol();
|
||||
// Keep a single row: reuse the existing row's id when present so the upsert
|
||||
// overwrites it in place instead of appending a new row on every call
|
||||
// (mirrors the qdrant provider's single-row telemetry id).
|
||||
const existing = await this.client.query({
|
||||
collection_name: "memory_migrations",
|
||||
filter: "",
|
||||
limit: 1,
|
||||
output_fields: ["id"],
|
||||
});
|
||||
const rows = existing?.data || [];
|
||||
const id =
|
||||
rows.length > 0 && rows[0].id ? String(rows[0].id) : this.generateUUID();
|
||||
await this.client.upsert({
|
||||
collection_name: "memory_migrations",
|
||||
data: [{ id, vectors: [0, 0], user_id: userId }],
|
||||
});
|
||||
}
|
||||
}
|
||||
@@ -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 })
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -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",
|
||||
|
||||
@@ -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 });
|
||||
}
|
||||
|
||||
@@ -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";
|
||||
|
||||
|
||||
@@ -0,0 +1,240 @@
|
||||
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 {
|
||||
/** 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>;
|
||||
collectionName: string;
|
||||
embeddingModelDims: number;
|
||||
}
|
||||
|
||||
const RETURN_PROPERTIES = [
|
||||
"ids",
|
||||
"hash",
|
||||
"metadata",
|
||||
"data",
|
||||
"created_at",
|
||||
"category",
|
||||
"updated_at",
|
||||
"user_id",
|
||||
"agent_id",
|
||||
"run_id",
|
||||
];
|
||||
|
||||
export class WeaviateDB implements VectorStore {
|
||||
private _config: WeaviateConfig;
|
||||
private _client!: WeaviateClient;
|
||||
private _sdk: any;
|
||||
private _col!: any;
|
||||
private _userId: string;
|
||||
private _initPromise?: Promise<void>;
|
||||
|
||||
constructor(config: WeaviateConfig) {
|
||||
this._config = config;
|
||||
this._userId = "";
|
||||
this.initialize().catch(console.error);
|
||||
}
|
||||
|
||||
initialize(): Promise<void> {
|
||||
return (this._initPromise ??= this._doInitialize());
|
||||
}
|
||||
|
||||
// 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;
|
||||
} else if (clusterUrl?.includes("localhost")) {
|
||||
this._client = await weaviate.connectToLocal({
|
||||
headers: additionalHeaders,
|
||||
});
|
||||
} else if (apiKey) {
|
||||
this._client = await weaviate.connectToWeaviateCloud(clusterUrl!, {
|
||||
authCredentials: new weaviate.ApiKey(apiKey),
|
||||
headers: additionalHeaders,
|
||||
});
|
||||
} else {
|
||||
if (!clusterUrl) {
|
||||
throw new Error(
|
||||
"WeaviateDB: clusterUrl is required when client and apiKey are not provided",
|
||||
);
|
||||
}
|
||||
const parsed = new URL(clusterUrl);
|
||||
const httpSecure = parsed.protocol === "https:";
|
||||
this._client = await weaviate.connectToCustom({
|
||||
httpHost: parsed.hostname,
|
||||
httpPort: parsed.port
|
||||
? parseInt(parsed.port, 10)
|
||||
: httpSecure
|
||||
? 443
|
||||
: 8080,
|
||||
httpSecure,
|
||||
grpcHost: parsed.hostname,
|
||||
grpcPort: 50051,
|
||||
grpcSecure: false,
|
||||
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) {
|
||||
await this._client.collections.create({
|
||||
name: collectionName,
|
||||
properties: RETURN_PROPERTIES.map((name) => ({
|
||||
name,
|
||||
dataType: "text" as const,
|
||||
})),
|
||||
vectorizers: weaviate.configure.vectorizer.none(),
|
||||
vectorIndex: weaviate.configure.vectorIndex.hnsw(),
|
||||
} as any);
|
||||
}
|
||||
|
||||
this._col = this._client.collections.get(collectionName);
|
||||
}
|
||||
|
||||
private _buildFilters(filters?: SearchFilters) {
|
||||
if (!filters) return undefined;
|
||||
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 ? this._sdk.Filters.and(...conditions) : undefined;
|
||||
}
|
||||
|
||||
async insert(
|
||||
vectors: number[][],
|
||||
ids: string[],
|
||||
payloads: Record<string, any>[],
|
||||
): Promise<void> {
|
||||
await this.initialize();
|
||||
const objects = vectors.map((vector, i) => ({
|
||||
id: ids[i],
|
||||
properties: payloads[i],
|
||||
vectors: vector,
|
||||
}));
|
||||
await this._col.data.insertMany(objects);
|
||||
}
|
||||
|
||||
async search(
|
||||
query: number[],
|
||||
topK?: number,
|
||||
filters?: SearchFilters,
|
||||
): Promise<VectorStoreResult[]> {
|
||||
await this.initialize();
|
||||
const result = await this._col.query.nearVector(query, {
|
||||
limit: topK ?? 10,
|
||||
filters: this._buildFilters(filters),
|
||||
returnMetadata: ["distance"],
|
||||
});
|
||||
return result.objects.map((obj: any) => ({
|
||||
id: obj.uuid,
|
||||
payload: obj.properties,
|
||||
score: 1 - obj.metadata.distance,
|
||||
}));
|
||||
}
|
||||
|
||||
async keywordSearch(
|
||||
query: string,
|
||||
topK?: number,
|
||||
filters?: SearchFilters,
|
||||
): Promise<VectorStoreResult[] | null> {
|
||||
await this.initialize();
|
||||
const result = await this._col.query.bm25(query, {
|
||||
queryProperties: ["data"],
|
||||
limit: topK ?? 10,
|
||||
filters: this._buildFilters(filters),
|
||||
returnMetadata: ["score"],
|
||||
});
|
||||
return result.objects.map((obj: any) => ({
|
||||
id: obj.uuid,
|
||||
payload: obj.properties,
|
||||
score: obj.metadata.score,
|
||||
}));
|
||||
}
|
||||
|
||||
async get(vectorId: string): Promise<VectorStoreResult | null> {
|
||||
await this.initialize();
|
||||
const obj = await this._col.query.fetchObjectById(vectorId, {
|
||||
returnProperties: RETURN_PROPERTIES,
|
||||
});
|
||||
if (!obj) return null;
|
||||
return { id: obj.uuid, payload: obj.properties };
|
||||
}
|
||||
|
||||
async update(
|
||||
vectorId: string,
|
||||
vector: number[],
|
||||
payload: Record<string, any>,
|
||||
): Promise<void> {
|
||||
await this.initialize();
|
||||
await this._col.data.update({
|
||||
id: vectorId,
|
||||
properties: payload,
|
||||
vectors: vector,
|
||||
});
|
||||
}
|
||||
|
||||
async delete(vectorId: string): Promise<void> {
|
||||
await this.initialize();
|
||||
await this._col.data.deleteById(vectorId);
|
||||
}
|
||||
|
||||
async deleteCol(): Promise<void> {
|
||||
await this.initialize();
|
||||
await this._client.collections.delete(this._config.collectionName);
|
||||
}
|
||||
|
||||
async list(
|
||||
filters?: SearchFilters,
|
||||
topK?: number,
|
||||
): Promise<[VectorStoreResult[], number]> {
|
||||
await this.initialize();
|
||||
const result = await this._col.query.fetchObjects({
|
||||
limit: topK ?? 100,
|
||||
filters: this._buildFilters(filters),
|
||||
returnProperties: RETURN_PROPERTIES,
|
||||
});
|
||||
const results = result.objects.map((obj: any) => ({
|
||||
id: obj.uuid,
|
||||
payload: obj.properties,
|
||||
}));
|
||||
return [results, results.length];
|
||||
}
|
||||
|
||||
async getUserId(): Promise<string> {
|
||||
if (!this._userId) {
|
||||
this._userId = uuidv4();
|
||||
}
|
||||
return this._userId;
|
||||
}
|
||||
|
||||
async setUserId(userId: string): Promise<void> {
|
||||
this._userId = userId;
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,617 @@
|
||||
import {
|
||||
AutoBuildPolicyType,
|
||||
FieldType,
|
||||
IndexType,
|
||||
InvertedIndexFieldAttribute,
|
||||
MetricType,
|
||||
PartitionType,
|
||||
ServerErrCode,
|
||||
TableState,
|
||||
} from "@mochow/mochow-sdk-node";
|
||||
import { BaiduDB } from "../src/vector_stores/baidu";
|
||||
|
||||
// No jest.mock() here on purpose: the real SDK supplies the enums and the search request
|
||||
// classes (which carry an internal `set` map the client reads, so they cannot be hand-rolled
|
||||
// as plain literals). Only the network-facing MochowClient is faked, via the `client` config.
|
||||
|
||||
const OK = { code: 0, msg: "" };
|
||||
const DIMS = 1536;
|
||||
|
||||
const normalTable = (schema: unknown = { fields: [], indexes: [] }) => ({
|
||||
...OK,
|
||||
table: { state: TableState.Normal, schema },
|
||||
});
|
||||
|
||||
const CORE_FIELDS = [
|
||||
{ fieldName: "id", fieldType: FieldType.String },
|
||||
{ fieldName: "data", fieldType: FieldType.Text },
|
||||
{ fieldName: "vector", fieldType: FieldType.FloatVector, dimension: DIMS },
|
||||
{ fieldName: "metadata", fieldType: "JSON" },
|
||||
];
|
||||
|
||||
const BM25_FIELDS = [
|
||||
...CORE_FIELDS,
|
||||
{ fieldName: "textLemmatized", fieldType: FieldType.Text },
|
||||
];
|
||||
|
||||
/** Records call order, so ordering regressions (deleteCol vs. in-flight init) are visible. */
|
||||
function fakeClient(overrides: Record<string, (...args: any[]) => any> = {}) {
|
||||
const calls: string[] = [];
|
||||
const track =
|
||||
(name: string, impl: (...args: any[]) => any) =>
|
||||
(...args: any[]) => {
|
||||
calls.push(name);
|
||||
return impl(...args);
|
||||
};
|
||||
|
||||
const client: any = {
|
||||
calls,
|
||||
createDatabase: jest.fn(track("createDatabase", async () => OK)),
|
||||
createTable: jest.fn(track("createTable", async () => OK)),
|
||||
dropTable: jest.fn(track("dropTable", async () => OK)),
|
||||
descTable: jest.fn(track("descTable", async () => normalTable())),
|
||||
upsert: jest.fn(async () => OK),
|
||||
delete: jest.fn(async () => OK),
|
||||
query: jest.fn(),
|
||||
select: jest.fn(),
|
||||
vectorSearch: jest.fn(),
|
||||
bm25Search: jest.fn(),
|
||||
};
|
||||
|
||||
for (const [name, impl] of Object.entries(overrides)) {
|
||||
client[name] = jest.fn(track(name, impl));
|
||||
}
|
||||
return client;
|
||||
}
|
||||
|
||||
const makeStore = (client: any, extra: Record<string, unknown> = {}) =>
|
||||
new BaiduDB({
|
||||
endpoint: "http://127.0.0.1:5287",
|
||||
account: "root",
|
||||
apiKey: "test-key",
|
||||
databaseName: "mem0_db",
|
||||
tableName: "mem0",
|
||||
embeddingModelDims: DIMS,
|
||||
client,
|
||||
...extra,
|
||||
} as any);
|
||||
|
||||
/** Run the poll loop's setTimeout inline so tests never wait the real 2s interval. */
|
||||
const runTimersInline = () =>
|
||||
jest.spyOn(global, "setTimeout").mockImplementation(((fn: () => void) => {
|
||||
fn();
|
||||
return 0;
|
||||
}) as any);
|
||||
|
||||
beforeEach(() => {
|
||||
jest.spyOn(console, "warn").mockImplementation(() => {});
|
||||
jest.spyOn(console, "error").mockImplementation(() => {});
|
||||
});
|
||||
|
||||
afterEach(() => jest.restoreAllMocks());
|
||||
|
||||
describe("BaiduDB config", () => {
|
||||
it("rejects a missing required field", () => {
|
||||
expect(() => makeStore(fakeClient(), { tableName: "" })).toThrow(
|
||||
/non-empty 'tableName'/,
|
||||
);
|
||||
});
|
||||
|
||||
it("does not require endpoint credentials when a client is injected", () => {
|
||||
expect(() =>
|
||||
makeStore(fakeClient(), { endpoint: "", account: "", apiKey: "" }),
|
||||
).not.toThrow();
|
||||
});
|
||||
});
|
||||
|
||||
describe("BaiduDB table provisioning", () => {
|
||||
it("creates the table with the schema mem0 needs", async () => {
|
||||
const client = fakeClient();
|
||||
await makeStore(client).initialize();
|
||||
|
||||
const spec = client.createTable.mock.calls[0][0];
|
||||
expect(client.createDatabase).toHaveBeenCalledWith("mem0_db");
|
||||
expect(spec.database).toBe("mem0_db");
|
||||
expect(spec.table).toBe("mem0");
|
||||
expect(spec.enableDynamicField).toBe(false);
|
||||
// Mochow rejects a partition without partitionType.
|
||||
expect(spec.partition).toEqual({
|
||||
partitionType: PartitionType.HASH,
|
||||
partitionNum: 1,
|
||||
});
|
||||
|
||||
const fields = spec.schema.fields.map((f: any) => [
|
||||
f.fieldName,
|
||||
f.fieldType,
|
||||
]);
|
||||
expect(fields).toEqual([
|
||||
["id", FieldType.String],
|
||||
["data", FieldType.Text],
|
||||
["vector", FieldType.FloatVector],
|
||||
["textLemmatized", FieldType.Text],
|
||||
["metadata", "JSON"],
|
||||
]);
|
||||
expect(spec.schema.fields[0]).toMatchObject({
|
||||
primaryKey: true,
|
||||
partitionKey: true,
|
||||
notNull: true,
|
||||
});
|
||||
expect(spec.schema.fields[2].dimension).toBe(DIMS);
|
||||
});
|
||||
|
||||
it("builds a vector index with a genuine row-count-increment auto-build policy", async () => {
|
||||
const client = fakeClient();
|
||||
await makeStore(client).initialize();
|
||||
|
||||
const [vectorIndex] = client.createTable.mock.calls[0][0].schema.indexes;
|
||||
expect(vectorIndex).toMatchObject({
|
||||
indexName: "vector_idx",
|
||||
indexType: IndexType.HNSW,
|
||||
field: "vector",
|
||||
metricType: MetricType.L2,
|
||||
params: { M: 16, efConstruction: 200 },
|
||||
autoBuild: true,
|
||||
});
|
||||
// Regression guard: sdk.AutoBuildIncrement() stamps policyType "TIMING" in 2.1.5.
|
||||
expect(vectorIndex.autoBuildPolicy).toEqual({
|
||||
policyType: AutoBuildPolicyType.Increment,
|
||||
rowCountIncrement: 10000,
|
||||
});
|
||||
expect(AutoBuildPolicyType.Increment).not.toBe(AutoBuildPolicyType.Timing);
|
||||
});
|
||||
|
||||
it("declares the filtering and BM25 indexes with an indexType", async () => {
|
||||
const client = fakeClient();
|
||||
await makeStore(client).initialize();
|
||||
|
||||
const [, filtering, bm25] =
|
||||
client.createTable.mock.calls[0][0].schema.indexes;
|
||||
expect(filtering).toEqual({
|
||||
indexName: "metadata_filtering_idx",
|
||||
indexType: IndexType.FilteringIndex,
|
||||
fields: ["metadata"],
|
||||
});
|
||||
// Memory.search() hands keywordSearch() an already-lemmatized query, so raw `data` is
|
||||
// not worth indexing — only the lemmatized column is. The index is therefore named for
|
||||
// that column and must never be called "data_bm25_idx": that is the name Python's
|
||||
// keyword_search() queries with a raw, unstemmed query, and it must keep missing (and so
|
||||
// falling back to vector search) rather than half-matching this stemmed index.
|
||||
expect(bm25).toMatchObject({
|
||||
indexName: "text_lemmatized_bm25_idx",
|
||||
indexType: IndexType.InvertedIndex,
|
||||
fields: ["textLemmatized"],
|
||||
fieldAttributes: [InvertedIndexFieldAttribute.Analyzed],
|
||||
});
|
||||
});
|
||||
|
||||
it("honours a configured metric type", async () => {
|
||||
const client = fakeClient();
|
||||
await makeStore(client, { metricType: "COSINE" }).initialize();
|
||||
const [vectorIndex] = client.createTable.mock.calls[0][0].schema.indexes;
|
||||
expect(vectorIndex.metricType).toBe(MetricType.COSINE);
|
||||
});
|
||||
|
||||
it("tolerates an existing database and table", async () => {
|
||||
const client = fakeClient({
|
||||
createDatabase: async () => ({
|
||||
code: ServerErrCode.DBAlreadyExist,
|
||||
msg: "db exists",
|
||||
}),
|
||||
createTable: async () => ({
|
||||
code: ServerErrCode.TableAlreadyExist,
|
||||
msg: "table exists",
|
||||
}),
|
||||
descTable: async () => normalTable({ fields: BM25_FIELDS, indexes: [] }),
|
||||
});
|
||||
await expect(makeStore(client).initialize()).resolves.toBeUndefined();
|
||||
});
|
||||
|
||||
it("waits for a CREATING table to become NORMAL", async () => {
|
||||
runTimersInline();
|
||||
const states = [
|
||||
TableState.Creating,
|
||||
TableState.Creating,
|
||||
TableState.Normal,
|
||||
];
|
||||
const client = fakeClient({
|
||||
descTable: async () => ({
|
||||
...OK,
|
||||
table: { state: states.shift(), schema: { fields: [], indexes: [] } },
|
||||
}),
|
||||
});
|
||||
|
||||
await makeStore(client).initialize();
|
||||
expect(client.descTable).toHaveBeenCalledTimes(3);
|
||||
});
|
||||
|
||||
it("surfaces a non-zero envelope as an error rather than succeeding", async () => {
|
||||
const client = fakeClient({
|
||||
createTable: async () => ({
|
||||
code: ServerErrCode.InvalidTableSchema,
|
||||
msg: "bad schema",
|
||||
}),
|
||||
});
|
||||
await expect(makeStore(client).initialize()).rejects.toThrow(
|
||||
/createTable 'mem0' failed \(code 60\): bad schema/,
|
||||
);
|
||||
});
|
||||
|
||||
it("rejects an existing table whose vector dimension disagrees", async () => {
|
||||
const client = fakeClient({
|
||||
createTable: async () => ({
|
||||
code: ServerErrCode.TableAlreadyExist,
|
||||
msg: "",
|
||||
}),
|
||||
descTable: async () =>
|
||||
normalTable({
|
||||
fields: [
|
||||
CORE_FIELDS[0],
|
||||
CORE_FIELDS[1],
|
||||
{
|
||||
fieldName: "vector",
|
||||
fieldType: FieldType.FloatVector,
|
||||
dimension: 768,
|
||||
},
|
||||
CORE_FIELDS[3],
|
||||
],
|
||||
indexes: [],
|
||||
}),
|
||||
});
|
||||
await expect(makeStore(client).initialize()).rejects.toThrow(
|
||||
/stores 768-dimensional vectors, but 'embeddingModelDims' is 1536/,
|
||||
);
|
||||
});
|
||||
|
||||
it("rejects an existing table missing the core schema", async () => {
|
||||
const client = fakeClient({
|
||||
createTable: async () => ({
|
||||
code: ServerErrCode.TableAlreadyExist,
|
||||
msg: "",
|
||||
}),
|
||||
descTable: async () =>
|
||||
normalTable({ fields: [CORE_FIELDS[0]], indexes: [] }),
|
||||
});
|
||||
await expect(makeStore(client).initialize()).rejects.toThrow(
|
||||
/missing the id\/data\/vector\/metadata schema/,
|
||||
);
|
||||
});
|
||||
});
|
||||
|
||||
describe("BaiduDB keyword search support detection", () => {
|
||||
it("fails closed when an existing table has no inverted index", async () => {
|
||||
const client = fakeClient({
|
||||
createTable: async () => ({
|
||||
code: ServerErrCode.TableAlreadyExist,
|
||||
msg: "",
|
||||
}),
|
||||
descTable: async () => normalTable({ fields: CORE_FIELDS, indexes: [] }),
|
||||
});
|
||||
const store = makeStore(client);
|
||||
|
||||
await expect(store.keywordSearch("hello")).resolves.toBeNull();
|
||||
expect(client.bm25Search).not.toHaveBeenCalled();
|
||||
expect(console.warn).toHaveBeenCalledWith(
|
||||
expect.stringContaining("text_lemmatized_bm25_idx"),
|
||||
);
|
||||
});
|
||||
|
||||
it("enables keyword search when the existing table carries the BM25 index", async () => {
|
||||
const client = fakeClient({
|
||||
createTable: async () => ({
|
||||
code: ServerErrCode.TableAlreadyExist,
|
||||
msg: "",
|
||||
}),
|
||||
descTable: async () =>
|
||||
normalTable({
|
||||
fields: BM25_FIELDS,
|
||||
indexes: [{ indexName: "text_lemmatized_bm25_idx" }],
|
||||
}),
|
||||
});
|
||||
client.bm25Search.mockResolvedValue({ ...OK, rows: [] });
|
||||
|
||||
await expect(makeStore(client).keywordSearch("hello")).resolves.toEqual([]);
|
||||
expect(client.bm25Search).toHaveBeenCalled();
|
||||
});
|
||||
|
||||
it("queries the inverted index with the caller's already-lemmatized text", async () => {
|
||||
const client = fakeClient();
|
||||
client.bm25Search.mockResolvedValue({
|
||||
...OK,
|
||||
rows: [
|
||||
{ row: { id: "m1", data: "loves pizza", metadata: {} }, score: 3.5 },
|
||||
],
|
||||
});
|
||||
|
||||
const results = await makeStore(client).keywordSearch("love pizza", 7, {
|
||||
userId: "alice",
|
||||
});
|
||||
expect(results).toEqual([
|
||||
{ id: "m1", payload: { data: "loves pizza" }, score: 3.5 },
|
||||
]);
|
||||
|
||||
const { request, ...ns } = client.bm25Search.mock.calls[0][0];
|
||||
expect(ns).toEqual({ database: "mem0_db", table: "mem0" });
|
||||
expect(request.indexName).toBe("text_lemmatized_bm25_idx");
|
||||
expect(request.searchText).toBe("love pizza");
|
||||
expect(request.limit).toBe(7);
|
||||
expect(request.filter).toBe('metadata["userId"] = "alice"');
|
||||
});
|
||||
});
|
||||
|
||||
describe("BaiduDB writes", () => {
|
||||
it("upserts the whole batch in one call and mirrors textLemmatized out of the payload", async () => {
|
||||
const client = fakeClient();
|
||||
|
||||
await makeStore(client).insert(
|
||||
[
|
||||
[1, 2],
|
||||
[3, 4],
|
||||
],
|
||||
["a", "b"],
|
||||
[
|
||||
{ data: "loves pizza", textLemmatized: "love pizza" },
|
||||
{ data: "runs daily" },
|
||||
],
|
||||
);
|
||||
|
||||
expect(client.upsert).toHaveBeenCalledTimes(1);
|
||||
expect(client.upsert.mock.calls[0][0]).toEqual({
|
||||
database: "mem0_db",
|
||||
table: "mem0",
|
||||
rows: [
|
||||
{
|
||||
id: "a",
|
||||
data: "loves pizza",
|
||||
vector: [1, 2],
|
||||
textLemmatized: "love pizza",
|
||||
metadata: {},
|
||||
},
|
||||
// Falls back to `data` when the caller did not lemmatize.
|
||||
{
|
||||
id: "b",
|
||||
data: "runs daily",
|
||||
vector: [3, 4],
|
||||
textLemmatized: "runs daily",
|
||||
metadata: {},
|
||||
},
|
||||
],
|
||||
});
|
||||
});
|
||||
|
||||
it("refuses a ragged batch instead of silently truncating it", async () => {
|
||||
await expect(
|
||||
makeStore(fakeClient()).insert([[1]], ["a", "b"], [{}]),
|
||||
).rejects.toThrow(/equal length \(got 1\/2\/1\)/);
|
||||
});
|
||||
|
||||
it("updates and deletes by primary key", async () => {
|
||||
const client = fakeClient();
|
||||
const store = makeStore(client);
|
||||
|
||||
await store.update("m1", [9], { data: "new" });
|
||||
expect(client.upsert.mock.calls[0][0].rows).toEqual([
|
||||
{
|
||||
id: "m1",
|
||||
data: "new",
|
||||
vector: [9],
|
||||
textLemmatized: "new",
|
||||
metadata: {},
|
||||
},
|
||||
]);
|
||||
|
||||
await store.delete("m1");
|
||||
expect(client.delete).toHaveBeenCalledWith({
|
||||
database: "mem0_db",
|
||||
table: "mem0",
|
||||
primaryKey: { id: "m1" },
|
||||
});
|
||||
});
|
||||
|
||||
it("throws when the server rejects an upsert", async () => {
|
||||
const client = fakeClient();
|
||||
client.upsert.mockResolvedValue({ code: 100, msg: "duplicate key" });
|
||||
|
||||
await expect(makeStore(client).insert([[1]], ["a"], [{}])).rejects.toThrow(
|
||||
/upsert failed \(code 100\): duplicate key/,
|
||||
);
|
||||
});
|
||||
});
|
||||
|
||||
describe("BaiduDB reads", () => {
|
||||
it("maps vector search hits out of the nested row envelope", async () => {
|
||||
const client = fakeClient();
|
||||
client.vectorSearch.mockResolvedValue({
|
||||
...OK,
|
||||
rows: [
|
||||
{
|
||||
row: { id: "m1", data: "x", metadata: {} },
|
||||
distance: 0.2,
|
||||
score: 0.8,
|
||||
},
|
||||
],
|
||||
});
|
||||
|
||||
const results = await makeStore(client).search([1, 2, 3], 5, {
|
||||
userId: "alice",
|
||||
});
|
||||
expect(results).toEqual([{ id: "m1", payload: { data: "x" }, score: 0.8 }]);
|
||||
|
||||
const { request } = client.vectorSearch.mock.calls[0][0];
|
||||
expect(request.vectorField).toBe("vector");
|
||||
expect(request.vector).toEqual({ vector: [1, 2, 3] });
|
||||
expect(request.limit).toBe(5);
|
||||
expect(request.filter).toBe('metadata["userId"] = "alice"');
|
||||
expect(request.projections).toEqual(["id", "data", "metadata"]);
|
||||
expect(request.config.params).toEqual({ ef: 200 });
|
||||
});
|
||||
|
||||
it("omits the filter when no filters are supplied", async () => {
|
||||
const client = fakeClient();
|
||||
client.vectorSearch.mockResolvedValue({ ...OK, rows: [] });
|
||||
await makeStore(client).search([1], 5);
|
||||
expect(client.vectorSearch.mock.calls[0][0].request.filter).toBeUndefined();
|
||||
});
|
||||
|
||||
it("escapes quotes and rejects unsafe filter keys and values", async () => {
|
||||
const client = fakeClient();
|
||||
client.vectorSearch.mockResolvedValue({ ...OK, rows: [] });
|
||||
const store = makeStore(client);
|
||||
|
||||
await store.search([1], 5, { userId: 'a"b', runId: 3, agentId: true });
|
||||
expect(client.vectorSearch.mock.calls[0][0].request.filter).toBe(
|
||||
'metadata["userId"] = "a\\"b" AND metadata["runId"] = 3 AND metadata["agentId"] = true',
|
||||
);
|
||||
|
||||
await expect(store.search([1], 5, { "bad key": "x" })).rejects.toThrow(
|
||||
/Invalid filter key/,
|
||||
);
|
||||
await expect(
|
||||
store.search([1], 5, { userId: ["a"] as any }),
|
||||
).rejects.toThrow(/must be str, int, float, or bool, got array/);
|
||||
});
|
||||
|
||||
it("returns null for a missing id and throws on a real query failure", async () => {
|
||||
const client = fakeClient();
|
||||
const store = makeStore(client);
|
||||
|
||||
client.query.mockResolvedValue({ ...OK, row: {} });
|
||||
await expect(store.get("nope")).resolves.toBeNull();
|
||||
|
||||
client.query.mockResolvedValue({
|
||||
...OK,
|
||||
row: { id: "m1", data: "stored text", metadata: { a: 1 } },
|
||||
});
|
||||
await expect(store.get("m1")).resolves.toEqual({
|
||||
id: "m1",
|
||||
payload: { a: 1, data: "stored text" },
|
||||
});
|
||||
|
||||
client.query.mockResolvedValue({ code: 2, msg: "invalid parameter" });
|
||||
await expect(store.get("m1")).rejects.toThrow(
|
||||
/query 'm1' failed \(code 2\): invalid parameter/,
|
||||
);
|
||||
});
|
||||
|
||||
// The server signals a missing primary key with code 101; pymochow and mochow-sdk-go both
|
||||
// name it (ROW_KEY_NOT_FOUND / RowKeyNotFound). The Node SDK's ServerErrCode stops at 100,
|
||||
// so it has to be spelled out. Memory.get()/update()/delete() all branch on a null here.
|
||||
it("returns null when the server reports the row key is missing", async () => {
|
||||
const client = fakeClient();
|
||||
const store = makeStore(client);
|
||||
|
||||
client.query.mockResolvedValue({ code: 101, msg: "row key not found" });
|
||||
await expect(store.get("nope")).resolves.toBeNull();
|
||||
});
|
||||
|
||||
it("lists flat select rows and reports how many came back", async () => {
|
||||
const client = fakeClient();
|
||||
client.select.mockResolvedValue({
|
||||
...OK,
|
||||
isTruncated: false,
|
||||
nextMarker: "",
|
||||
rows: [{ id: "m1", data: "x", metadata: {} }, { id: "m2" }],
|
||||
});
|
||||
|
||||
await expect(
|
||||
makeStore(client).list({ userId: "alice" }, 50),
|
||||
).resolves.toEqual([
|
||||
[
|
||||
{ id: "m1", payload: { data: "x" } },
|
||||
{ id: "m2", payload: {} },
|
||||
],
|
||||
2,
|
||||
]);
|
||||
expect(client.select).toHaveBeenCalledWith({
|
||||
database: "mem0_db",
|
||||
table: "mem0",
|
||||
filter: 'metadata["userId"] = "alice"',
|
||||
projections: ["id", "data", "metadata"],
|
||||
limit: 50,
|
||||
});
|
||||
});
|
||||
});
|
||||
|
||||
describe("BaiduDB deleteCol", () => {
|
||||
it("waits for the drop to land before returning", async () => {
|
||||
runTimersInline();
|
||||
const client = fakeClient();
|
||||
const store = makeStore(client);
|
||||
await store.initialize();
|
||||
|
||||
client.descTable
|
||||
.mockResolvedValueOnce({ ...OK, table: { state: TableState.Deleting } })
|
||||
.mockResolvedValueOnce({
|
||||
code: ServerErrCode.TableNotExist,
|
||||
msg: "gone",
|
||||
});
|
||||
|
||||
await store.deleteCol();
|
||||
expect(client.dropTable).toHaveBeenCalledWith("mem0_db", "mem0");
|
||||
expect(client.descTable).toHaveBeenCalledTimes(3); // 1 from initialize + 2 polls
|
||||
});
|
||||
|
||||
it("is a no-op when the table is already gone", async () => {
|
||||
const client = fakeClient();
|
||||
const store = makeStore(client);
|
||||
await store.initialize();
|
||||
|
||||
client.dropTable.mockResolvedValue({
|
||||
code: ServerErrCode.TableNotExist,
|
||||
msg: "gone",
|
||||
});
|
||||
await expect(store.deleteCol()).resolves.toBeUndefined();
|
||||
});
|
||||
|
||||
// Regression: deleteCol() used to run alongside the fire-and-forget initialize() the
|
||||
// constructor starts, so the in-flight createTable landed *after* dropTable and the table
|
||||
// survived reset().
|
||||
it("does not race the initialize() the constructor kicks off", async () => {
|
||||
runTimersInline();
|
||||
let release: () => void = () => {};
|
||||
const gate = new Promise<void>((resolve) => {
|
||||
release = resolve;
|
||||
});
|
||||
|
||||
// The table exists after the first init, is gone once dropTable lands, then exists again.
|
||||
const descQueue: unknown[] = [
|
||||
normalTable(),
|
||||
{ code: ServerErrCode.TableNotExist, msg: "gone" },
|
||||
];
|
||||
const client = fakeClient({
|
||||
createDatabase: async () => {
|
||||
await gate;
|
||||
return OK;
|
||||
},
|
||||
descTable: async () => descQueue.shift() ?? normalTable(),
|
||||
});
|
||||
|
||||
const store = makeStore(client); // initialize() is now in flight, parked on `gate`
|
||||
const resetting = store.reset();
|
||||
release();
|
||||
await resetting;
|
||||
|
||||
expect(client.calls).toEqual([
|
||||
"createDatabase",
|
||||
"createTable",
|
||||
"descTable",
|
||||
"dropTable",
|
||||
"descTable",
|
||||
"createDatabase",
|
||||
"createTable",
|
||||
"descTable",
|
||||
]);
|
||||
expect(client.calls.indexOf("dropTable")).toBeGreaterThan(
|
||||
client.calls.indexOf("createTable"),
|
||||
);
|
||||
expect(client.createTable).toHaveBeenCalledTimes(2);
|
||||
});
|
||||
});
|
||||
|
||||
describe("BaiduDB user id", () => {
|
||||
it("round-trips the store user id", async () => {
|
||||
const store = makeStore(fakeClient());
|
||||
await expect(store.getUserId()).resolves.toBe("anonymous-baidu-user");
|
||||
await store.setUserId("alice");
|
||||
await expect(store.getUserId()).resolves.toBe("alice");
|
||||
});
|
||||
});
|
||||
@@ -378,6 +378,55 @@ describe("ConfigManager", () => {
|
||||
});
|
||||
});
|
||||
|
||||
describe("mergeConfig - AWS Bedrock credential passthrough", () => {
|
||||
const baseEmbedder = { provider: "openai", config: { apiKey: "k" } };
|
||||
const baseVectorStore = { provider: "memory", config: {} };
|
||||
|
||||
it("preserves explicit AWS Bedrock credentials through MemoryConfigSchema.parse", () => {
|
||||
const cfg = ConfigManager.mergeConfig({
|
||||
embedder: baseEmbedder,
|
||||
vectorStore: baseVectorStore,
|
||||
llm: {
|
||||
provider: "aws_bedrock",
|
||||
config: {
|
||||
model: "anthropic.claude-3-sonnet-20240229-v1:0",
|
||||
awsRegion: "us-east-1",
|
||||
awsAccessKeyId: "AKIAEXAMPLE",
|
||||
awsSecretAccessKey: "secret-example",
|
||||
awsSessionToken: "session-example",
|
||||
} as any,
|
||||
},
|
||||
});
|
||||
|
||||
// Regression: zod default parse strips undeclared keys, which silently
|
||||
// dropped credentials before .passthrough() + typed fields were added.
|
||||
expect(cfg.llm.config.awsRegion).toBe("us-east-1");
|
||||
expect(cfg.llm.config.awsAccessKeyId).toBe("AKIAEXAMPLE");
|
||||
expect(cfg.llm.config.awsSecretAccessKey).toBe("secret-example");
|
||||
expect(cfg.llm.config.awsSessionToken).toBe("session-example");
|
||||
});
|
||||
|
||||
it("normalizes snake_case AWS credentials from the raw config", () => {
|
||||
const cfg = ConfigManager.mergeConfig({
|
||||
embedder: baseEmbedder,
|
||||
vectorStore: baseVectorStore,
|
||||
llm: {
|
||||
provider: "aws_bedrock",
|
||||
config: {
|
||||
model: "anthropic.claude-3-sonnet-20240229-v1:0",
|
||||
aws_region: "eu-west-1",
|
||||
aws_access_key_id: "AKIASNAKE",
|
||||
aws_secret_access_key: "snake-secret",
|
||||
} as any,
|
||||
},
|
||||
});
|
||||
|
||||
expect(cfg.llm.config.awsRegion).toBe("eu-west-1");
|
||||
expect(cfg.llm.config.awsAccessKeyId).toBe("AKIASNAKE");
|
||||
expect(cfg.llm.config.awsSecretAccessKey).toBe("snake-secret");
|
||||
});
|
||||
});
|
||||
|
||||
describe("mergeConfig - full OpenClaw-style LM Studio config", () => {
|
||||
it("handles the exact config from issue #4235", () => {
|
||||
const cfg = ConfigManager.mergeConfig({
|
||||
@@ -422,6 +471,85 @@ describe("ConfigManager", () => {
|
||||
expect(cfg.vectorStore.config.port).toBe(6333);
|
||||
});
|
||||
});
|
||||
|
||||
describe("mergeConfig - provider-specific embedder fields", () => {
|
||||
// The embedder config used to be rebuilt from a fixed key list, which
|
||||
// dropped every provider-specific field before the embedder was
|
||||
// constructed. Vertex AI then authenticated against whatever ambient
|
||||
// project ADC resolved to and ignored the configured task types.
|
||||
it("preserves Vertex AI fields through the merge", () => {
|
||||
const cfg = ConfigManager.mergeConfig({
|
||||
embedder: {
|
||||
provider: "vertexai",
|
||||
config: {
|
||||
model: "gemini-embedding-001",
|
||||
googleProjectId: "my-proj",
|
||||
location: "europe-west4",
|
||||
vertexCredentialsJson: "/creds.json",
|
||||
memoryAddEmbeddingType: "SEMANTIC_SIMILARITY",
|
||||
},
|
||||
},
|
||||
vectorStore: { provider: "memory", config: { collectionName: "test" } },
|
||||
llm: { provider: "openai", config: { apiKey: "test-key" } },
|
||||
});
|
||||
|
||||
expect(cfg.embedder.config).toMatchObject({
|
||||
model: "gemini-embedding-001",
|
||||
googleProjectId: "my-proj",
|
||||
location: "europe-west4",
|
||||
vertexCredentialsJson: "/creds.json",
|
||||
memoryAddEmbeddingType: "SEMANTIC_SIMILARITY",
|
||||
});
|
||||
});
|
||||
|
||||
it("still lets normalized values win over the raw user config", () => {
|
||||
const cfg = ConfigManager.mergeConfig({
|
||||
embedder: {
|
||||
provider: "lmstudio",
|
||||
config: {
|
||||
lmstudio_base_url: "http://localhost:1234/v1",
|
||||
embedding_dims: 768,
|
||||
},
|
||||
} as never,
|
||||
vectorStore: { provider: "memory", config: { collectionName: "test" } },
|
||||
llm: { provider: "openai", config: { apiKey: "test-key" } },
|
||||
});
|
||||
|
||||
expect(cfg.embedder.config.baseURL).toBe("http://localhost:1234/v1");
|
||||
expect(cfg.embedder.config.embeddingDims).toBe(768);
|
||||
// Snake_case aliases are normalized, not passed through to the provider.
|
||||
expect(cfg.embedder.config).not.toHaveProperty("lmstudio_base_url");
|
||||
expect(cfg.embedder.config).not.toHaveProperty("embedding_dims");
|
||||
});
|
||||
});
|
||||
|
||||
describe("mergeConfig - vector store provider normalization", () => {
|
||||
const baseLlm = { provider: "openai", config: { apiKey: "test-key" } };
|
||||
const baseEmbedder = { provider: "openai", config: { apiKey: "test-key" } };
|
||||
|
||||
it.each(["Memory", "DATABRICKS", "QdRaNt"])(
|
||||
"lowercases the vector store provider %p",
|
||||
(provider) => {
|
||||
const config = ConfigManager.mergeConfig({
|
||||
embedder: baseEmbedder,
|
||||
vectorStore: { provider, config: { collectionName: "test" } },
|
||||
llm: baseLlm,
|
||||
});
|
||||
|
||||
expect(config.vectorStore.provider).toBe(provider.toLowerCase());
|
||||
},
|
||||
);
|
||||
|
||||
it("still falls back to the default provider when none is given", () => {
|
||||
const config = ConfigManager.mergeConfig({
|
||||
embedder: baseEmbedder,
|
||||
vectorStore: { config: { collectionName: "test" } } as any,
|
||||
llm: baseLlm,
|
||||
});
|
||||
|
||||
expect(config.vectorStore.provider).toBe("memory");
|
||||
});
|
||||
});
|
||||
});
|
||||
|
||||
// ─────────────────────────────────────────────────────────────────────────
|
||||
@@ -629,7 +757,10 @@ describe("Memory – LM Studio end-to-end flow", () => {
|
||||
filters: { user_id: "u1" },
|
||||
});
|
||||
|
||||
expect(mockEmbedder.embed).toHaveBeenCalledWith("What does the user like?");
|
||||
expect(mockEmbedder.embed).toHaveBeenCalledWith(
|
||||
"What does the user like?",
|
||||
"search",
|
||||
);
|
||||
expect(mockVStore.search).toHaveBeenCalled();
|
||||
expect(result.results).toHaveLength(1);
|
||||
expect(result.results[0].memory).toBe("User likes hiking");
|
||||
@@ -659,6 +790,47 @@ describe("Memory – LM Studio end-to-end flow", () => {
|
||||
);
|
||||
});
|
||||
|
||||
it("passes AWS Bedrock credentials to LLMFactory through the Memory stack", async () => {
|
||||
const mem = new MemoryClass({
|
||||
embedder: {
|
||||
provider: "lmstudio",
|
||||
config: {
|
||||
model: "nomic-embed-text-v1.5",
|
||||
baseURL: "http://localhost:1234/v1",
|
||||
embeddingDims: 768,
|
||||
},
|
||||
},
|
||||
vectorStore: {
|
||||
provider: "memory",
|
||||
config: { collectionName: "test", dimension: 768 },
|
||||
},
|
||||
llm: {
|
||||
provider: "aws_bedrock",
|
||||
config: {
|
||||
model: "anthropic.claude-3-sonnet-20240229-v1:0",
|
||||
awsRegion: "us-east-1",
|
||||
awsAccessKeyId: "AKIAEXAMPLE",
|
||||
awsSecretAccessKey: "secret-example",
|
||||
},
|
||||
},
|
||||
disableHistory: true,
|
||||
});
|
||||
|
||||
await mem.getAll({ filters: { user_id: "u1" } });
|
||||
|
||||
// Regression: credentials must survive ConfigManager -> Memory -> LLMFactory,
|
||||
// otherwise Bedrock silently falls back to the ambient AWS credential chain.
|
||||
expect(mockLlmFactory.create).toHaveBeenCalledWith(
|
||||
"aws_bedrock",
|
||||
expect.objectContaining({
|
||||
model: "anthropic.claude-3-sonnet-20240229-v1:0",
|
||||
awsRegion: "us-east-1",
|
||||
awsAccessKeyId: "AKIAEXAMPLE",
|
||||
awsSecretAccessKey: "secret-example",
|
||||
}),
|
||||
);
|
||||
});
|
||||
|
||||
it("add flow works with lmstudio LLM for fact extraction", async () => {
|
||||
mockLlm.generateResponse.mockResolvedValueOnce(
|
||||
'{"facts":["User loves sushi"]}',
|
||||
|
||||
@@ -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({
|
||||
|
||||
@@ -40,6 +40,11 @@ jest.mock("../src/embeddings/lmstudio", () => ({
|
||||
.fn()
|
||||
.mockImplementation((config) => ({ type: "lmstudio-embedder", config })),
|
||||
}));
|
||||
jest.mock("../src/embeddings/vertexai", () => ({
|
||||
VertexAIEmbedder: jest
|
||||
.fn()
|
||||
.mockImplementation((config) => ({ type: "vertexai-embedder", config })),
|
||||
}));
|
||||
jest.mock("../src/embeddings/together", () => ({
|
||||
TogetherEmbedder: jest
|
||||
.fn()
|
||||
@@ -107,6 +112,11 @@ jest.mock("../src/llms/xai", () => ({
|
||||
.fn()
|
||||
.mockImplementation((config) => ({ type: "xai-llm", config })),
|
||||
}));
|
||||
jest.mock("../src/llms/sarvam", () => ({
|
||||
SarvamLLM: jest
|
||||
.fn()
|
||||
.mockImplementation((config) => ({ type: "sarvam-llm", config })),
|
||||
}));
|
||||
jest.mock("../src/llms/litellm", () => ({
|
||||
LiteLLM: jest
|
||||
.fn()
|
||||
@@ -133,6 +143,11 @@ jest.mock("../src/vector_stores/qdrant", () => ({
|
||||
.fn()
|
||||
.mockImplementation((config) => ({ type: "qdrant", config })),
|
||||
}));
|
||||
jest.mock("../src/vector_stores/baidu", () => ({
|
||||
BaiduDB: jest
|
||||
.fn()
|
||||
.mockImplementation((config) => ({ type: "baidu", config })),
|
||||
}));
|
||||
jest.mock("../src/vector_stores/redis", () => ({
|
||||
RedisDB: jest
|
||||
.fn()
|
||||
@@ -168,6 +183,17 @@ jest.mock("../src/vector_stores/pgvector", () => ({
|
||||
.fn()
|
||||
.mockImplementation((config) => ({ type: "pgvector", config })),
|
||||
}));
|
||||
jest.mock("../src/vector_stores/databricks", () => ({
|
||||
DatabricksVectorStore: jest
|
||||
.fn()
|
||||
.mockImplementation((config) => ({ type: "databricks", config })),
|
||||
}));
|
||||
jest.mock("../src/vector_stores/neptune_analytics", () => ({
|
||||
NeptuneAnalyticsVectorStore: jest.fn().mockImplementation((config) => ({
|
||||
type: "neptune-analytics",
|
||||
config,
|
||||
})),
|
||||
}));
|
||||
jest.mock("../src/vector_stores/upstash_vector", () => ({
|
||||
UpstashVector: jest
|
||||
.fn()
|
||||
@@ -188,6 +214,11 @@ jest.mock("../src/vector_stores/s3_vectors", () => ({
|
||||
.fn()
|
||||
.mockImplementation((config) => ({ type: "s3-vectors", config })),
|
||||
}));
|
||||
jest.mock("../src/vector_stores/weaviate", () => ({
|
||||
WeaviateDB: jest
|
||||
.fn()
|
||||
.mockImplementation((config) => ({ type: "weaviate", config })),
|
||||
}));
|
||||
jest.mock("../src/storage/SupabaseHistoryManager", () => ({
|
||||
SupabaseHistoryManager: jest
|
||||
.fn()
|
||||
@@ -226,6 +257,7 @@ describe("EmbedderFactory", () => {
|
||||
["fastembed"],
|
||||
["langchain"],
|
||||
["lmstudio"],
|
||||
["vertexai"],
|
||||
["together"],
|
||||
])("creates embedder for provider '%s'", (provider) => {
|
||||
expect(() =>
|
||||
@@ -269,6 +301,7 @@ describe("LLMFactory", () => {
|
||||
["lmstudio"],
|
||||
["deepseek"],
|
||||
["xai"],
|
||||
["sarvam"],
|
||||
["litellm"],
|
||||
["minimax"],
|
||||
["together"],
|
||||
@@ -309,6 +342,7 @@ describe("VectorStoreFactory", () => {
|
||||
});
|
||||
|
||||
test.each([
|
||||
["baidu"],
|
||||
["qdrant"],
|
||||
["redis"],
|
||||
["valkey"],
|
||||
@@ -317,15 +351,42 @@ describe("VectorStoreFactory", () => {
|
||||
["vectorize"],
|
||||
["azure-ai-search"],
|
||||
["pgvector"],
|
||||
["databricks"],
|
||||
["neptune"],
|
||||
["neptune-analytics"],
|
||||
["upstash_vector"],
|
||||
["azure_mysql"],
|
||||
["cassandra"],
|
||||
["s3-vectors"],
|
||||
["s3_vectors"],
|
||||
["weaviate"],
|
||||
])("creates vector store for provider '%s'", (provider) => {
|
||||
expect(() =>
|
||||
VectorStoreFactory.create(provider, dummyVSConfig),
|
||||
).not.toThrow();
|
||||
const result = VectorStoreFactory.create(provider, dummyVSConfig) as any;
|
||||
expect(result.config).toBe(dummyVSConfig);
|
||||
});
|
||||
|
||||
test("passes Neptune endpoint URI config through the factory", () => {
|
||||
const config = {
|
||||
collectionName: "test",
|
||||
dimension: 4,
|
||||
endpoint: "neptune-graph://g-1234567890",
|
||||
region: "us-east-1",
|
||||
};
|
||||
const store = VectorStoreFactory.create("neptune", config) as any;
|
||||
|
||||
expect(store.config).toEqual(config);
|
||||
});
|
||||
|
||||
test("keeps neptune-analytics as a compatibility alias", () => {
|
||||
const config = {
|
||||
collectionName: "test",
|
||||
dimension: 4,
|
||||
endpoint: "neptune-graph://g-1234567890",
|
||||
region: "us-east-1",
|
||||
};
|
||||
const store = VectorStoreFactory.create("neptune-analytics", config) as any;
|
||||
|
||||
expect(store.config).toEqual(config);
|
||||
});
|
||||
|
||||
test("throws for unsupported provider", () => {
|
||||
|
||||
@@ -0,0 +1,149 @@
|
||||
/// <reference types="jest" />
|
||||
/**
|
||||
* HuggingFace Embedder unit tests (mocked OpenAI client).
|
||||
* The TS provider targets a HuggingFace TEI / OpenAI-compatible inference
|
||||
* endpoint, so it reuses the `openai` client with a HuggingFace baseURL.
|
||||
* These tests verify the required base URL, request shape, and batch ordering.
|
||||
*/
|
||||
|
||||
const mockEmbeddingsCreate = jest.fn();
|
||||
const mockOpenAICtor = jest.fn();
|
||||
|
||||
jest.mock("openai", () => {
|
||||
return {
|
||||
__esModule: true,
|
||||
default: jest.fn().mockImplementation((opts: any) => {
|
||||
mockOpenAICtor(opts);
|
||||
return { embeddings: { create: mockEmbeddingsCreate } };
|
||||
}),
|
||||
};
|
||||
});
|
||||
|
||||
import { HuggingFaceEmbedder } from "../src/embeddings/huggingface";
|
||||
|
||||
const mockEmbedding = [0.1, 0.2, 0.3, 0.4, 0.5];
|
||||
|
||||
describe("HuggingFaceEmbedder (unit)", () => {
|
||||
const OLD_ENV = process.env;
|
||||
|
||||
beforeEach(() => {
|
||||
mockEmbeddingsCreate.mockReset();
|
||||
mockOpenAICtor.mockReset();
|
||||
mockEmbeddingsCreate.mockResolvedValue({
|
||||
data: [{ index: 0, embedding: mockEmbedding }],
|
||||
});
|
||||
process.env = { ...OLD_ENV };
|
||||
delete process.env.HUGGINGFACE_BASE_URL;
|
||||
});
|
||||
|
||||
afterAll(() => {
|
||||
process.env = OLD_ENV;
|
||||
});
|
||||
|
||||
describe("configuration", () => {
|
||||
it("throws when no inference endpoint is configured", () => {
|
||||
expect(() => new HuggingFaceEmbedder({ apiKey: "test-key" })).toThrow(
|
||||
/requires an inference endpoint/,
|
||||
);
|
||||
});
|
||||
|
||||
it("uses huggingfaceBaseUrl and the default model", async () => {
|
||||
const embedder = new HuggingFaceEmbedder({
|
||||
apiKey: "test-key",
|
||||
huggingfaceBaseUrl: "http://localhost:8080/v1",
|
||||
});
|
||||
await embedder.embed("hello");
|
||||
|
||||
expect(mockOpenAICtor.mock.calls[0][0]).toMatchObject({
|
||||
apiKey: "test-key",
|
||||
baseURL: "http://localhost:8080/v1",
|
||||
});
|
||||
const callArgs = mockEmbeddingsCreate.mock.calls[0][0];
|
||||
expect(callArgs).toEqual({ model: "tei", input: "hello" });
|
||||
});
|
||||
|
||||
it("falls back to baseURL and honors a custom model", async () => {
|
||||
const embedder = new HuggingFaceEmbedder({
|
||||
baseURL: "https://tei.example.com/v1",
|
||||
model: "BAAI/bge-small-en-v1.5",
|
||||
});
|
||||
await embedder.embed("hello");
|
||||
|
||||
expect(mockOpenAICtor.mock.calls[0][0]).toMatchObject({
|
||||
baseURL: "https://tei.example.com/v1",
|
||||
});
|
||||
expect(mockEmbeddingsCreate.mock.calls[0][0].model).toBe(
|
||||
"BAAI/bge-small-en-v1.5",
|
||||
);
|
||||
});
|
||||
|
||||
it("reads HUGGINGFACE_BASE_URL from the environment", async () => {
|
||||
process.env.HUGGINGFACE_BASE_URL = "http://env-host:8080/v1";
|
||||
const embedder = new HuggingFaceEmbedder({ apiKey: "test-key" });
|
||||
await embedder.embed("hello");
|
||||
|
||||
expect(mockOpenAICtor.mock.calls[0][0]).toMatchObject({
|
||||
baseURL: "http://env-host:8080/v1",
|
||||
});
|
||||
});
|
||||
|
||||
it("never forwards a dimensions parameter", async () => {
|
||||
const embedder = new HuggingFaceEmbedder({
|
||||
huggingfaceBaseUrl: "http://localhost:8080/v1",
|
||||
embeddingDims: 384,
|
||||
});
|
||||
await embedder.embed("hello");
|
||||
|
||||
expect(mockEmbeddingsCreate.mock.calls[0][0]).not.toHaveProperty(
|
||||
"dimensions",
|
||||
);
|
||||
});
|
||||
});
|
||||
|
||||
describe("basic functionality", () => {
|
||||
const cfg = { huggingfaceBaseUrl: "http://localhost:8080/v1" };
|
||||
|
||||
it("embed() returns the embedding vector", async () => {
|
||||
const embedder = new HuggingFaceEmbedder(cfg);
|
||||
expect(await embedder.embed("hello")).toEqual(mockEmbedding);
|
||||
});
|
||||
|
||||
it("embedBatch() returns [] for empty input without calling the API", async () => {
|
||||
const embedder = new HuggingFaceEmbedder(cfg);
|
||||
expect(await embedder.embedBatch([])).toEqual([]);
|
||||
expect(mockEmbeddingsCreate).not.toHaveBeenCalled();
|
||||
});
|
||||
|
||||
it("embedBatch() sorts results by index", async () => {
|
||||
mockEmbeddingsCreate.mockResolvedValue({
|
||||
data: [
|
||||
{ index: 1, embedding: [0.3, 0.4] },
|
||||
{ index: 0, embedding: [0.1, 0.2] },
|
||||
],
|
||||
});
|
||||
const embedder = new HuggingFaceEmbedder(cfg);
|
||||
expect(await embedder.embedBatch(["a", "b"])).toEqual([
|
||||
[0.1, 0.2],
|
||||
[0.3, 0.4],
|
||||
]);
|
||||
});
|
||||
|
||||
it("embedBatch() throws when the count mismatches the input", async () => {
|
||||
mockEmbeddingsCreate.mockResolvedValue({
|
||||
data: [{ index: 0, embedding: [0.1, 0.2] }],
|
||||
});
|
||||
const embedder = new HuggingFaceEmbedder(cfg);
|
||||
await expect(embedder.embedBatch(["a", "b"])).rejects.toThrow(
|
||||
/returned 1 embeddings for 2 texts/,
|
||||
);
|
||||
});
|
||||
|
||||
it("embed() throws when the endpoint returns no embeddings", async () => {
|
||||
mockEmbeddingsCreate.mockResolvedValue({ data: [] });
|
||||
const embedder = new HuggingFaceEmbedder(cfg);
|
||||
await expect(embedder.embed("hello")).rejects.toThrow(
|
||||
/returned no embeddings/,
|
||||
);
|
||||
});
|
||||
});
|
||||
});
|
||||
@@ -5,6 +5,7 @@
|
||||
/// <reference types="jest" />
|
||||
import { Memory } from "../src/memory";
|
||||
import type { MemoryItem, SearchResult } from "../src/types";
|
||||
import { logger } from "../src/utils/logger";
|
||||
|
||||
jest.setTimeout(30000);
|
||||
|
||||
@@ -150,7 +151,7 @@ describe("Memory - update()", () => {
|
||||
infer: false,
|
||||
});
|
||||
const id = addResult.results[0].id;
|
||||
const result = await memory.update(id, "Updated");
|
||||
const result = await memory.update(id, { text: "Updated" });
|
||||
expect(result.message).toBe("Memory updated successfully!");
|
||||
});
|
||||
|
||||
@@ -160,7 +161,7 @@ describe("Memory - update()", () => {
|
||||
infer: false,
|
||||
});
|
||||
const id = addResult.results[0].id;
|
||||
await memory.update(id, "After update");
|
||||
await memory.update(id, { text: "After update" });
|
||||
const item: MemoryItem | null = await memory.get(id);
|
||||
expect(item!.memory).toBe("After update");
|
||||
});
|
||||
@@ -174,7 +175,7 @@ describe("Memory - update()", () => {
|
||||
const before: MemoryItem | null = await memory.get(id);
|
||||
const originalCreatedAt = before!.createdAt;
|
||||
|
||||
await memory.update(id, "New text");
|
||||
await memory.update(id, { text: "New text" });
|
||||
const after: MemoryItem | null = await memory.get(id);
|
||||
expect(after!.createdAt).toBe(originalCreatedAt);
|
||||
expect(after!.updatedAt).toBeDefined();
|
||||
@@ -187,7 +188,7 @@ describe("Memory - update()", () => {
|
||||
});
|
||||
const id = addResult.results[0].id;
|
||||
const before: MemoryItem | null = await memory.get(id);
|
||||
await memory.update(id, "Completely different text");
|
||||
await memory.update(id, { text: "Completely different text" });
|
||||
const after: MemoryItem | null = await memory.get(id);
|
||||
expect(after!.hash).not.toBe(before!.hash);
|
||||
});
|
||||
@@ -199,7 +200,7 @@ describe("Memory - update()", () => {
|
||||
infer: false,
|
||||
});
|
||||
const id = addResult.results[0].id;
|
||||
await memory.update(id, "Updated text");
|
||||
await memory.update(id, { text: "Updated text" });
|
||||
const after: MemoryItem | null = await memory.get(id);
|
||||
expect(after!.memory).toBe("Updated text");
|
||||
expect(after!.metadata).toEqual(
|
||||
@@ -208,6 +209,309 @@ describe("Memory - update()", () => {
|
||||
});
|
||||
});
|
||||
|
||||
// ─── update() options: text / data / metadata / expirationDate ───
|
||||
|
||||
describe("Memory - update() options", () => {
|
||||
let memory: Memory;
|
||||
let warnSpy: jest.SpyInstance;
|
||||
const userId = `update_options_${Date.now()}`;
|
||||
|
||||
beforeAll(async () => {
|
||||
memory = createMemory();
|
||||
});
|
||||
|
||||
beforeEach(() => {
|
||||
warnSpy = jest.spyOn(logger, "warn").mockImplementation(() => {});
|
||||
});
|
||||
|
||||
afterEach(() => {
|
||||
warnSpy.mockRestore();
|
||||
});
|
||||
|
||||
afterAll(async () => {
|
||||
await memory.reset();
|
||||
});
|
||||
|
||||
async function seed(text: string): Promise<string> {
|
||||
const addResult: SearchResult = await memory.add(text, {
|
||||
userId,
|
||||
infer: false,
|
||||
});
|
||||
return addResult.results[0].id;
|
||||
}
|
||||
|
||||
test("accepts an options object with text, without warning", async () => {
|
||||
const id = await seed("Options before");
|
||||
await memory.update(id, { text: "Options after" });
|
||||
const after: MemoryItem | null = await memory.get(id);
|
||||
expect(after!.memory).toBe("Options after");
|
||||
expect(warnSpy).not.toHaveBeenCalled();
|
||||
});
|
||||
|
||||
test("accepts a bare text string, as on main", async () => {
|
||||
const id = await seed("Bare before");
|
||||
await memory.update(id, "Bare after");
|
||||
const after: MemoryItem | null = await memory.get(id);
|
||||
expect(after!.memory).toBe("Bare after");
|
||||
});
|
||||
|
||||
test("accepts the deprecated data alias and warns", async () => {
|
||||
const id = await seed("Alias before");
|
||||
await memory.update(id, { data: "Alias after" });
|
||||
const after: MemoryItem | null = await memory.get(id);
|
||||
expect(after!.memory).toBe("Alias after");
|
||||
expect(warnSpy).toHaveBeenCalledWith(expect.stringContaining("deprecated"));
|
||||
});
|
||||
|
||||
// An empty `text` is content, so it must beat `data`. Guards against `||=`,
|
||||
// which would treat "" as absent and store the `data` value instead.
|
||||
test("text wins over data, even when text is empty", async () => {
|
||||
const id = await seed("Both before");
|
||||
await memory.update(id, { text: "", data: "From data" });
|
||||
const after: MemoryItem | null = await memory.get(id);
|
||||
expect(after!.memory).toBe("");
|
||||
});
|
||||
|
||||
// Same trap on the guard: "" must not read as "no field provided".
|
||||
test("an empty bare string is content, not a missing argument", async () => {
|
||||
const id = await seed("Empty bare");
|
||||
await memory.update(id, "");
|
||||
const after: MemoryItem | null = await memory.get(id);
|
||||
expect(after!.memory).toBe("");
|
||||
});
|
||||
|
||||
test("updates metadata without touching the stored text", async () => {
|
||||
const id = await seed("Metadata only");
|
||||
await memory.update(id, { metadata: { category: "solo" } });
|
||||
const after: MemoryItem | null = await memory.get(id);
|
||||
expect(after!.memory).toBe("Metadata only");
|
||||
expect(after!.metadata).toEqual(
|
||||
expect.objectContaining({ category: "solo" }),
|
||||
);
|
||||
});
|
||||
|
||||
test("sets an expiration date without touching the stored text", async () => {
|
||||
const id = await seed("Expiry only");
|
||||
await memory.update(id, { expirationDate: "2099-12-31" });
|
||||
const after: MemoryItem | null = await memory.get(id);
|
||||
expect(after!.memory).toBe("Expiry only");
|
||||
expect(after!.metadata).toEqual(
|
||||
expect.objectContaining({ expiration_date: "2099-12-31" }),
|
||||
);
|
||||
});
|
||||
|
||||
// Python raises on `data=None` / `metadata=None`, so loose `== null` is the
|
||||
// right check for both. `{}` is a real value and must not raise.
|
||||
test.each([{}, { data: null }, { metadata: null }])(
|
||||
"throws when %p provides nothing updatable",
|
||||
async (options) => {
|
||||
const id = await seed("Nothing to update");
|
||||
await expect(memory.update(id, options as any)).rejects.toThrow(
|
||||
"At least one of text, metadata, or expirationDate must be provided.",
|
||||
);
|
||||
expect(warnSpy).not.toHaveBeenCalled();
|
||||
},
|
||||
);
|
||||
|
||||
test("accepts empty metadata as an updatable field", async () => {
|
||||
const id = await seed("Empty metadata");
|
||||
await expect(memory.update(id, { metadata: {} })).resolves.toEqual({
|
||||
message: "Memory updated successfully!",
|
||||
});
|
||||
});
|
||||
});
|
||||
|
||||
// ─── expiration date parsing ─────────────────────────────
|
||||
|
||||
describe("Memory - expiration date parsing", () => {
|
||||
let memory: Memory;
|
||||
const userId = `expiry_parse_test_${Date.now()}`;
|
||||
|
||||
beforeAll(async () => {
|
||||
memory = createMemory();
|
||||
});
|
||||
|
||||
afterAll(async () => {
|
||||
await memory.reset();
|
||||
});
|
||||
|
||||
// `new Date(...)` accepts all of these; Python's date.fromisoformat rejects
|
||||
// them. The first two die to the format regex, the last two only to the
|
||||
// UTC component round-trip: they match YYYY-MM-DD but are not real days.
|
||||
const rejected = [
|
||||
"12/31/2099", // also shifts a day west of UTC
|
||||
"2099-12-31T23:00:00",
|
||||
"2099-02-30", // rolls over to 2099-03-02
|
||||
"2100-02-29", // 2100 is not a leap year
|
||||
];
|
||||
|
||||
test.each(rejected)("add() rejects %p", async (value) => {
|
||||
await expect(
|
||||
memory.add("Bad expiry", { userId, infer: false, expirationDate: value }),
|
||||
).rejects.toThrow("YYYY-MM-DD");
|
||||
});
|
||||
|
||||
// add() and update() share normalizeExpirationDate(); this checks the wiring.
|
||||
test("update() rejects a malformed expiration date", async () => {
|
||||
const addResult: SearchResult = await memory.add("Good", {
|
||||
userId,
|
||||
infer: false,
|
||||
});
|
||||
await expect(
|
||||
memory.update(addResult.results[0].id, { expirationDate: "12/31/2099" }),
|
||||
).rejects.toThrow("YYYY-MM-DD");
|
||||
});
|
||||
|
||||
// "2096-02-29" is a real leap day: it must survive the component check.
|
||||
test.each(["2099-12-31", "2096-02-29"])(
|
||||
"stores %p verbatim",
|
||||
async (value) => {
|
||||
const addResult: SearchResult = await memory.add("Good expiry", {
|
||||
userId,
|
||||
infer: false,
|
||||
expirationDate: value,
|
||||
});
|
||||
const item: MemoryItem | null = await memory.get(addResult.results[0].id);
|
||||
expect(item!.metadata!.expiration_date).toBe(value);
|
||||
},
|
||||
);
|
||||
});
|
||||
|
||||
// ─── expired memories are hidden on read ─────────────────
|
||||
|
||||
describe("Memory - expired memories", () => {
|
||||
let memory: Memory;
|
||||
const userId = `expired_test_${Date.now()}`;
|
||||
const today = new Date().toISOString().slice(0, 10);
|
||||
|
||||
beforeAll(async () => {
|
||||
memory = createMemory();
|
||||
});
|
||||
|
||||
afterAll(async () => {
|
||||
await memory.reset();
|
||||
});
|
||||
|
||||
async function seed(text: string, expirationDate?: string): Promise<string> {
|
||||
const addResult: SearchResult = await memory.add(text, {
|
||||
userId,
|
||||
infer: false,
|
||||
...(expirationDate ? { expirationDate } : {}),
|
||||
});
|
||||
return addResult.results[0].id;
|
||||
}
|
||||
|
||||
async function seedLiveAndDead(scopedUser: string): Promise<void> {
|
||||
await memory.add("Live memory", { userId: scopedUser, infer: false });
|
||||
await memory.add("Dead memory", {
|
||||
userId: scopedUser,
|
||||
infer: false,
|
||||
expirationDate: "2020-01-01",
|
||||
});
|
||||
}
|
||||
|
||||
test("getAll() hides expired memories unless showExpired", async () => {
|
||||
const scopedUser = `${userId}_getall`;
|
||||
await seedLiveAndDead(scopedUser);
|
||||
const filters = { user_id: scopedUser };
|
||||
|
||||
const hidden: SearchResult = await memory.getAll({ filters });
|
||||
expect(hidden.results.map((r) => r.memory)).toEqual(["Live memory"]);
|
||||
|
||||
const shown: SearchResult = await memory.getAll({
|
||||
filters,
|
||||
showExpired: true,
|
||||
});
|
||||
expect(shown.results.map((r) => r.memory).sort()).toEqual([
|
||||
"Dead memory",
|
||||
"Live memory",
|
||||
]);
|
||||
});
|
||||
|
||||
test("search() hides expired memories unless showExpired", async () => {
|
||||
const scopedUser = `${userId}_search`;
|
||||
await seedLiveAndDead(scopedUser);
|
||||
const filters = { user_id: scopedUser };
|
||||
|
||||
const hidden: SearchResult = await memory.search("memory", { filters });
|
||||
expect(hidden.results.map((r) => r.memory)).toEqual(["Live memory"]);
|
||||
|
||||
const shown: SearchResult = await memory.search("memory", {
|
||||
filters,
|
||||
showExpired: true,
|
||||
});
|
||||
expect(shown.results.map((r) => r.memory).sort()).toEqual([
|
||||
"Dead memory",
|
||||
"Live memory",
|
||||
]);
|
||||
});
|
||||
|
||||
test("a memory expiring today is not yet expired", async () => {
|
||||
const scopedUser = `${userId}_today`;
|
||||
await memory.add("Expires today", {
|
||||
userId: scopedUser,
|
||||
infer: false,
|
||||
expirationDate: today,
|
||||
});
|
||||
const result: SearchResult = await memory.getAll({
|
||||
filters: { user_id: scopedUser },
|
||||
});
|
||||
expect(result.results).toHaveLength(1);
|
||||
});
|
||||
|
||||
test("get() still returns an expired memory by ID", async () => {
|
||||
const id = await seed("Fetch by id", "2020-01-01");
|
||||
const item: MemoryItem | null = await memory.get(id);
|
||||
expect(item).not.toBeNull();
|
||||
expect(item!.memory).toBe("Fetch by id");
|
||||
});
|
||||
|
||||
test("clearing the expiration date makes a memory visible again", async () => {
|
||||
const scopedUser = `${userId}_revive`;
|
||||
const addResult: SearchResult = await memory.add("Revived", {
|
||||
userId: scopedUser,
|
||||
infer: false,
|
||||
expirationDate: "2020-01-01",
|
||||
});
|
||||
const id = addResult.results[0].id;
|
||||
|
||||
const before: SearchResult = await memory.getAll({
|
||||
filters: { user_id: scopedUser },
|
||||
});
|
||||
expect(before.results).toHaveLength(0);
|
||||
|
||||
await memory.update(id, { expirationDate: null });
|
||||
|
||||
const after: SearchResult = await memory.getAll({
|
||||
filters: { user_id: scopedUser },
|
||||
});
|
||||
expect(after.results.map((r) => r.memory)).toEqual(["Revived"]);
|
||||
});
|
||||
|
||||
test("getAll() still fills topK when expired memories are present", async () => {
|
||||
const scopedUser = `${userId}_topk`;
|
||||
for (let i = 0; i < 3; i++) {
|
||||
await memory.add(`Dead ${i}`, {
|
||||
userId: scopedUser,
|
||||
infer: false,
|
||||
expirationDate: "2020-01-01",
|
||||
});
|
||||
}
|
||||
for (let i = 0; i < 3; i++) {
|
||||
await memory.add(`Live ${i}`, { userId: scopedUser, infer: false });
|
||||
}
|
||||
|
||||
const result: SearchResult = await memory.getAll({
|
||||
filters: { user_id: scopedUser },
|
||||
topK: 3,
|
||||
});
|
||||
expect(result.results).toHaveLength(3);
|
||||
expect(result.results.every((r) => r.memory!.startsWith("Live"))).toBe(
|
||||
true,
|
||||
);
|
||||
});
|
||||
});
|
||||
|
||||
// ─── delete() ────────────────────────────────────────────
|
||||
|
||||
describe("Memory - delete()", () => {
|
||||
@@ -428,7 +732,7 @@ describe("Memory - history()", () => {
|
||||
userId,
|
||||
});
|
||||
const id = addResult.results[0].id;
|
||||
await memory.update(id, "After");
|
||||
await memory.update(id, { text: "After" });
|
||||
const history = await memory.history(id);
|
||||
expect(history.length).toBeGreaterThanOrEqual(2);
|
||||
});
|
||||
|
||||
@@ -0,0 +1,115 @@
|
||||
/// <reference types="jest" />
|
||||
/**
|
||||
* Verifies the memory pipeline threads the correct memory action
|
||||
* ("add" | "update" | "search") into the embedder. Task-type-aware providers
|
||||
* (e.g. Vertex AI) embed queries and documents differently based on this, and
|
||||
* the argument is silently ignored by every other embedder, so only a
|
||||
* pipeline-level test catches a dropped action.
|
||||
*/
|
||||
import { Memory } from "../src/memory";
|
||||
|
||||
const mockEmbedding = new Array(1536).fill(0.1);
|
||||
// Prefixed `mock*` so jest's hoisted module factory may reference them.
|
||||
const mockEmbed = jest.fn().mockResolvedValue(mockEmbedding);
|
||||
const mockEmbedBatch = jest
|
||||
.fn()
|
||||
.mockImplementation((texts: string[]) =>
|
||||
Promise.resolve(texts.map(() => mockEmbedding)),
|
||||
);
|
||||
|
||||
const mockGenerateResponse = jest
|
||||
.fn()
|
||||
.mockResolvedValue(JSON.stringify({ memory: [] }));
|
||||
|
||||
jest.mock("../src/embeddings/google", () => ({ GoogleEmbedder: jest.fn() }));
|
||||
jest.mock("../src/llms/google", () => ({ GoogleLLM: jest.fn() }));
|
||||
jest.mock("../src/llms/openai", () => ({
|
||||
OpenAILLM: jest.fn().mockImplementation(() => ({
|
||||
generateResponse: mockGenerateResponse,
|
||||
})),
|
||||
}));
|
||||
jest.mock("../src/embeddings/openai", () => ({
|
||||
OpenAIEmbedder: jest.fn().mockImplementation(() => ({
|
||||
embed: mockEmbed,
|
||||
embedBatch: mockEmbedBatch,
|
||||
embeddingDims: 1536,
|
||||
})),
|
||||
}));
|
||||
|
||||
function createMemory(): Memory {
|
||||
return new Memory({
|
||||
version: "v1.1",
|
||||
embedder: {
|
||||
provider: "openai",
|
||||
config: { apiKey: "test-key", model: "text-embedding-3-small" },
|
||||
},
|
||||
vectorStore: {
|
||||
provider: "memory",
|
||||
config: {
|
||||
collectionName: `test-action-${Date.now()}-${Math.random()}`,
|
||||
dimension: 1536,
|
||||
dbPath: ":memory:",
|
||||
},
|
||||
},
|
||||
llm: {
|
||||
provider: "openai",
|
||||
config: { apiKey: "test-key", model: "gpt-5-mini" },
|
||||
},
|
||||
historyDbPath: ":memory:",
|
||||
});
|
||||
}
|
||||
|
||||
describe("embedder memory-action threading", () => {
|
||||
let memory: Memory;
|
||||
|
||||
beforeEach(() => {
|
||||
memory = createMemory();
|
||||
mockEmbed.mockClear();
|
||||
mockEmbedBatch.mockClear();
|
||||
mockGenerateResponse.mockResolvedValue(JSON.stringify({ memory: [] }));
|
||||
});
|
||||
|
||||
afterEach(async () => {
|
||||
await memory.reset();
|
||||
});
|
||||
|
||||
test("search() embeds the query with the 'search' action", async () => {
|
||||
await memory.search("what do I like", { filters: { user_id: "u1" } });
|
||||
expect(mockEmbed).toHaveBeenCalledWith("what do I like", "search");
|
||||
});
|
||||
|
||||
test("update() embeds the new value with the 'update' action", async () => {
|
||||
// Missing id: update embeds the value before it throws on the absent row.
|
||||
await memory.update("missing-id", "new value").catch(() => {});
|
||||
expect(mockEmbed).toHaveBeenCalledWith("new value", "update");
|
||||
});
|
||||
|
||||
test("add() batch-embeds extracted memories and entities with the 'add' action", async () => {
|
||||
mockGenerateResponse.mockResolvedValue(
|
||||
JSON.stringify({
|
||||
memory: [
|
||||
{ id: "1", text: "John loves sci-fi movies", attributed_to: "user" },
|
||||
],
|
||||
}),
|
||||
);
|
||||
|
||||
await memory.add("I love sci-fi movies", { userId: "u1" });
|
||||
|
||||
// Phase 1 retrieval embeds the incoming turn as a query.
|
||||
expect(mockEmbed).toHaveBeenCalledWith(
|
||||
expect.stringContaining("I love sci-fi movies"),
|
||||
"search",
|
||||
);
|
||||
// Phase 3 (extracted memories) and phase 7 (linked entities) both batch
|
||||
// embed as documents. Without an explicit action, a task-type-aware
|
||||
// embedder falls back to its own default and silently mis-embeds.
|
||||
expect(mockEmbedBatch).toHaveBeenCalledWith(
|
||||
["John loves sci-fi movies"],
|
||||
"add",
|
||||
);
|
||||
expect(mockEmbedBatch.mock.calls.length).toBeGreaterThan(0);
|
||||
for (const call of mockEmbedBatch.mock.calls) {
|
||||
expect(call[1]).toBe("add");
|
||||
}
|
||||
});
|
||||
});
|
||||
@@ -0,0 +1,104 @@
|
||||
/**
|
||||
* Regression test for #6101 / #5903: LLM extraction transport failures
|
||||
* (rate limits, timeouts, connection errors) must propagate as a typed
|
||||
* `LLMError`, not be silently swallowed into an empty result. Mirrors the
|
||||
* Python SDK regression test added in #5878
|
||||
* (`test_llm_extraction_exception_is_reraised`).
|
||||
*/
|
||||
/// <reference types="jest" />
|
||||
import { Memory, LLMError } from "../src/memory";
|
||||
import type { SearchResult } from "../src/types";
|
||||
|
||||
jest.setTimeout(15000);
|
||||
|
||||
// Mock Google modules to prevent @google/genai crash in CI
|
||||
jest.mock("../src/embeddings/google", () => ({
|
||||
GoogleEmbedder: jest.fn(),
|
||||
}));
|
||||
jest.mock("../src/llms/google", () => ({
|
||||
GoogleLLM: jest.fn(),
|
||||
}));
|
||||
|
||||
class _ProviderError extends Error {}
|
||||
|
||||
jest.mock("../src/llms/openai", () => ({
|
||||
OpenAILLM: jest.fn().mockImplementation(() => ({
|
||||
generateResponse: jest
|
||||
.fn()
|
||||
.mockRejectedValue(new _ProviderError("429 rate limit")),
|
||||
})),
|
||||
}));
|
||||
|
||||
const mockEmbedding = new Array(1536).fill(0.1);
|
||||
jest.mock("../src/embeddings/openai", () => ({
|
||||
OpenAIEmbedder: jest.fn().mockImplementation(() => ({
|
||||
embed: jest.fn().mockResolvedValue(mockEmbedding),
|
||||
embedBatch: jest
|
||||
.fn()
|
||||
.mockImplementation((texts: string[]) =>
|
||||
Promise.resolve(texts.map(() => mockEmbedding)),
|
||||
),
|
||||
embeddingDims: 1536,
|
||||
})),
|
||||
}));
|
||||
|
||||
function createMemory(): Memory {
|
||||
return new Memory({
|
||||
version: "v1.1",
|
||||
embedder: {
|
||||
provider: "openai",
|
||||
config: { apiKey: "test-key", model: "text-embedding-3-small" },
|
||||
},
|
||||
vectorStore: {
|
||||
provider: "memory",
|
||||
config: {
|
||||
collectionName: `test-llm-error-${Date.now()}`,
|
||||
dimension: 1536,
|
||||
dbPath: ":memory:",
|
||||
},
|
||||
},
|
||||
llm: {
|
||||
provider: "openai",
|
||||
config: { apiKey: "test-key", model: "gpt-5-mini" },
|
||||
},
|
||||
historyDbPath: ":memory:",
|
||||
});
|
||||
}
|
||||
|
||||
describe("Memory - LLM extraction transport failures", () => {
|
||||
let memory: Memory;
|
||||
const userId = `llm_error_test_${Date.now()}`;
|
||||
|
||||
beforeAll(async () => {
|
||||
memory = createMemory();
|
||||
});
|
||||
|
||||
afterAll(async () => {
|
||||
await memory.reset();
|
||||
});
|
||||
|
||||
test("add() rejects with LLMError instead of returning an empty result", async () => {
|
||||
await expect(
|
||||
memory.add("this should trigger a provider failure", { userId }),
|
||||
).rejects.toBeInstanceOf(LLMError);
|
||||
});
|
||||
|
||||
test("thrown LLMError preserves the original error as its cause", async () => {
|
||||
let caught: unknown;
|
||||
try {
|
||||
const result: SearchResult = await memory.add("trigger failure again", {
|
||||
userId,
|
||||
});
|
||||
// Should never reach here — fail loudly if the call resolves.
|
||||
throw new Error(
|
||||
`Expected add() to reject, but it resolved with: ${JSON.stringify(result)}`,
|
||||
);
|
||||
} catch (e) {
|
||||
caught = e;
|
||||
}
|
||||
|
||||
expect(caught).toBeInstanceOf(LLMError);
|
||||
expect((caught as LLMError).cause).toBeInstanceOf(_ProviderError);
|
||||
expect((caught as LLMError).message).toContain("429 rate limit");
|
||||
});
|
||||
});
|
||||
@@ -0,0 +1,184 @@
|
||||
/**
|
||||
* Reranker integration tests for Memory.search().
|
||||
*
|
||||
* Verifies the per-search `rerank` flag: when a reranker is configured and
|
||||
* `rerank: true` is passed, search reorders results by the reranker's output;
|
||||
* otherwise results pass through unchanged. Failures degrade gracefully.
|
||||
*/
|
||||
/// <reference types="jest" />
|
||||
import { Memory } from "../src/memory";
|
||||
import { CohereReranker } from "../src/rerankers/cohere";
|
||||
import type { RerankResult } from "../src/rerankers/base";
|
||||
|
||||
jest.setTimeout(15000);
|
||||
|
||||
jest.mock("../src/embeddings/google", () => ({
|
||||
GoogleEmbedder: jest.fn(),
|
||||
}));
|
||||
jest.mock("../src/llms/google", () => ({
|
||||
GoogleLLM: jest.fn(),
|
||||
}));
|
||||
jest.mock("../src/llms/openai", () => ({
|
||||
OpenAILLM: jest.fn().mockImplementation(() => ({
|
||||
generateResponse: jest
|
||||
.fn()
|
||||
.mockResolvedValue(JSON.stringify({ memory: [] })),
|
||||
})),
|
||||
}));
|
||||
|
||||
const mockEmbedding = new Array(1536).fill(0.1);
|
||||
jest.mock("../src/embeddings/openai", () => ({
|
||||
OpenAIEmbedder: jest.fn().mockImplementation(() => ({
|
||||
embed: jest.fn().mockResolvedValue(mockEmbedding),
|
||||
embedBatch: jest
|
||||
.fn()
|
||||
.mockImplementation((texts: string[]) =>
|
||||
Promise.resolve(texts.map(() => mockEmbedding)),
|
||||
),
|
||||
embeddingDims: 1536,
|
||||
})),
|
||||
}));
|
||||
|
||||
function createMemory(config: Record<string, any> = {}): Memory {
|
||||
return new Memory({
|
||||
version: "v1.1",
|
||||
embedder: {
|
||||
provider: "openai",
|
||||
config: { apiKey: "test-key", model: "text-embedding-3-small" },
|
||||
},
|
||||
vectorStore: {
|
||||
provider: "memory",
|
||||
config: {
|
||||
collectionName: `test-rerank-${Date.now()}-${Math.random()}`,
|
||||
dimension: 1536,
|
||||
dbPath: ":memory:",
|
||||
},
|
||||
},
|
||||
llm: {
|
||||
provider: "openai",
|
||||
config: { apiKey: "test-key", model: "gpt-5-mini" },
|
||||
},
|
||||
historyDbPath: ":memory:",
|
||||
...config,
|
||||
});
|
||||
}
|
||||
|
||||
// Two semantic results whose natural (score-sorted) order is [alpha, bravo].
|
||||
async function primeSearch(m: any) {
|
||||
await m._ensureInitialized();
|
||||
m.embedder = { embed: jest.fn().mockResolvedValue(mockEmbedding) };
|
||||
m.vectorStore.search = jest.fn().mockResolvedValue([
|
||||
{ id: "a", score: 0.9, payload: { data: "alpha" } },
|
||||
{ id: "b", score: 0.8, payload: { data: "bravo" } },
|
||||
]);
|
||||
m.vectorStore.keywordSearch = jest.fn().mockResolvedValue(null);
|
||||
}
|
||||
|
||||
describe("Memory.search reranking", () => {
|
||||
it("reorders results by the reranker when rerank:true, adding rerankScore while preserving the original vector score", async () => {
|
||||
const memory = createMemory();
|
||||
const m = memory as any;
|
||||
await primeSearch(m);
|
||||
|
||||
const rerank = jest
|
||||
.fn<Promise<RerankResult[]>, [string, string[], number?]>()
|
||||
.mockResolvedValue([
|
||||
{ index: 1, rerankScore: 0.99 }, // bravo
|
||||
{ index: 0, rerankScore: 0.4 }, // alpha
|
||||
]);
|
||||
m.reranker = { rerank };
|
||||
|
||||
const result = await m.search("what did i eat", {
|
||||
filters: { user_id: "u1" },
|
||||
rerank: true,
|
||||
});
|
||||
|
||||
expect(rerank).toHaveBeenCalledWith(
|
||||
"what did i eat",
|
||||
["alpha", "bravo"],
|
||||
expect.any(Number),
|
||||
);
|
||||
expect(result.results.map((r: any) => r.memory)).toEqual([
|
||||
"bravo",
|
||||
"alpha",
|
||||
]);
|
||||
expect(result.results[0].rerankScore).toBe(0.99);
|
||||
expect(result.results[1].rerankScore).toBe(0.4);
|
||||
// The original vector similarity `score` must survive reranking.
|
||||
expect(result.results[0].score).toBe(0.8); // bravo's original vector score
|
||||
expect(result.results[1].score).toBe(0.9); // alpha's original vector score
|
||||
|
||||
await memory.reset();
|
||||
});
|
||||
|
||||
it("leaves results untouched and does not call the reranker when rerank is omitted", async () => {
|
||||
const memory = createMemory();
|
||||
const m = memory as any;
|
||||
await primeSearch(m);
|
||||
const rerank = jest.fn();
|
||||
m.reranker = { rerank };
|
||||
|
||||
const result = await m.search("what did i eat", {
|
||||
filters: { user_id: "u1" },
|
||||
});
|
||||
|
||||
expect(rerank).not.toHaveBeenCalled();
|
||||
expect(result.results.map((r: any) => r.memory)).toEqual([
|
||||
"alpha",
|
||||
"bravo",
|
||||
]);
|
||||
expect(result.results[0].rerankScore).toBeUndefined();
|
||||
|
||||
await memory.reset();
|
||||
});
|
||||
|
||||
it("is a no-op (no throw) when rerank:true but no reranker is configured", async () => {
|
||||
const memory = createMemory();
|
||||
const m = memory as any;
|
||||
await primeSearch(m);
|
||||
|
||||
const result = await m.search("what did i eat", {
|
||||
filters: { user_id: "u1" },
|
||||
rerank: true,
|
||||
});
|
||||
|
||||
expect(result.results.map((r: any) => r.memory)).toEqual([
|
||||
"alpha",
|
||||
"bravo",
|
||||
]);
|
||||
|
||||
await memory.reset();
|
||||
});
|
||||
|
||||
it("falls back to the original results when the reranker throws", async () => {
|
||||
const memory = createMemory();
|
||||
const m = memory as any;
|
||||
await primeSearch(m);
|
||||
m.reranker = {
|
||||
rerank: jest.fn().mockRejectedValue(new Error("provider down")),
|
||||
};
|
||||
const warnSpy = jest.spyOn(console, "warn").mockImplementation(() => {});
|
||||
|
||||
const result = await m.search("what did i eat", {
|
||||
filters: { user_id: "u1" },
|
||||
rerank: true,
|
||||
});
|
||||
|
||||
expect(result.results.map((r: any) => r.memory)).toEqual([
|
||||
"alpha",
|
||||
"bravo",
|
||||
]);
|
||||
expect(warnSpy).toHaveBeenCalled();
|
||||
|
||||
warnSpy.mockRestore();
|
||||
await memory.reset();
|
||||
});
|
||||
|
||||
it("wires a reranker from config in the constructor", () => {
|
||||
const memory = createMemory({
|
||||
reranker: { provider: "cohere", config: { apiKey: "test-key" } },
|
||||
});
|
||||
|
||||
expect((memory as any).reranker).toBeInstanceOf(CohereReranker);
|
||||
});
|
||||
});
|
||||
@@ -0,0 +1,70 @@
|
||||
import { readdirSync, readFileSync, statSync } from "fs";
|
||||
import { join, relative, resolve } from "path";
|
||||
|
||||
// Optional peers are not installed by npm/pnpm. A static value import of one therefore throws
|
||||
// MODULE_NOT_FOUND the moment anything pulls in `mem0ai/oss`, because src/index.ts re-exports every
|
||||
// vector store. Load them with `await import(...)` inside the code path that needs them instead.
|
||||
|
||||
const packageRoot = resolve(__dirname, "../../..");
|
||||
|
||||
function optionalPeers(): string[] {
|
||||
const pkg = JSON.parse(
|
||||
readFileSync(join(packageRoot, "package.json"), "utf8"),
|
||||
);
|
||||
return Object.entries(
|
||||
(pkg.peerDependenciesMeta ?? {}) as Record<string, { optional?: boolean }>,
|
||||
)
|
||||
.filter(([, meta]) => meta.optional)
|
||||
.map(([name]) => name);
|
||||
}
|
||||
|
||||
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);
|
||||
} else if (entry.endsWith(".ts") && !entry.endsWith(".test.ts")) {
|
||||
acc.push(full);
|
||||
}
|
||||
}
|
||||
return acc;
|
||||
}
|
||||
|
||||
// Matches `import ... from "pkg"` and `import "pkg"`, but not `import type ... from "pkg"`,
|
||||
// `typeof import("pkg")`, or `await import("pkg")` — those are erased or already lazy.
|
||||
function hasStaticValueImport(pkg: string, source: string): boolean {
|
||||
const escaped = pkg.replace(/[.*+?^${}()|[\]\\]/g, "\\$&");
|
||||
const specifier = `["']${escaped}(?:\\/[^"']*)?["']`;
|
||||
return (
|
||||
new RegExp(
|
||||
`(?:^|\\n)\\s*import\\s+(?!type\\b)[^;]*?from\\s+${specifier}`,
|
||||
).test(source) ||
|
||||
new RegExp(`(?:^|\\n)\\s*import\\s+${specifier}`).test(source)
|
||||
);
|
||||
}
|
||||
|
||||
describe("optional peer dependencies", () => {
|
||||
const peers = optionalPeers();
|
||||
|
||||
it("are discoverable from package.json", () => {
|
||||
expect(peers.length).toBeGreaterThan(0);
|
||||
});
|
||||
|
||||
it("are never statically imported by src", () => {
|
||||
const files = sourceFiles(join(packageRoot, "src"));
|
||||
const sources = new Map(
|
||||
files.map((file) => [file, readFileSync(file, "utf8")]),
|
||||
);
|
||||
|
||||
const offenders: string[] = [];
|
||||
for (const peer of peers) {
|
||||
for (const [file, source] of sources) {
|
||||
if (hasStaticValueImport(peer, source)) {
|
||||
offenders.push(`${relative(packageRoot, file)} imports ${peer}`);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
expect(offenders).toEqual([]);
|
||||
});
|
||||
});
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user