Compare commits

..

24 Commits

Author SHA1 Message Date
Saket Aryan 50db9e428d chore(release): bump SDK versions to next beta (#4859) 2026-04-16 16:23:50 +05:30
soumil-rathi fb87349664 fix(oss): v3 entity cleanup, filter fixes, and QA hardening (TS + Python) (#4858)
Co-authored-by: Soumil Rathi <soumilrathi@gmail.com>
Co-authored-by: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
2026-04-16 12:17:13 +05:30
Saket Aryan 8827553576 fix: adopt new v3 memory endpoints in Python + TS clients (#4856) 2026-04-16 05:28:49 +05:30
Kabir Kohli c8e20a9bb5 fix(docs): resolve duplicate operationIds and expiration_date type in openapi spec (#4854) 2026-04-16 04:09:51 +05:30
Kartik 93a51f4763 test: update integration tests for v1.1 output_format (#4847) 2026-04-16 01:37:51 +05:30
Saket Aryan 86fe275f53 fix(ts): entity store isolation, backward compat, pgvector + redis init fixes (#4841) 2026-04-15 21:01:01 +05:30
Kartik e6d6276bb9 refactor: add entity ID and search param validation, rename textLemmatized field, update tests (#4843) 2026-04-15 20:57:09 +05:30
Chaithanya Kumar 9692726db4 fix(ts-oss): isolate entity store from memory store by default (#4829)
Co-authored-by: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
Co-authored-by: kartik-mem0 <kartik.labhshetwar@mem0.ai>
2026-04-15 14:47:43 +05:30
soumil-rathi d8d776636f fix(v3): migration crashes + entity linking on OSS (#4836)
Co-authored-by: Soumil Rathi <soumilrathi@gmail.com>
Co-authored-by: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
2026-04-15 14:46:28 +05:30
Kartik a5a688295e fix: prevent arbitrary code execution via pickle in FAISS vector store (#4833) 2026-04-15 00:02:07 +05:30
Saket Aryan 5d40592e42 chore: version bump to beta1 (#4827) 2026-04-14 18:05:22 +05:30
soumil-rathi a488e19044 feat(oss): port v3 pipeline with hybrid search, entity extraction, and additive scoring (#4805)
Co-authored-by: Soumil Rathi <soumilrathi@gmail.com>
Co-authored-by: Saket Aryan <saketaryan2002@gmail.com>
Co-authored-by: chaithanyak42 <chaithanya.kumar42a@gmail.com>
Co-authored-by: kartik-mem0 <kartik.labhshetwar@mem0.ai>
2026-04-14 18:00:58 +05:30
Parteeksachdeva 57f944e18a fix: allow anonymousTelemetryId in openclaw.json config (#4826)
Co-authored-by: parteeksachdeva-123 <parteek.sachdeva@aerchain.io>
2026-04-14 17:31:28 +05:30
Gabriel Stein fe3f7ae618 fix(plugin): remove invalid keys from Claude plugin config (#4821) 2026-04-14 01:13:03 +05:30
shafdev 4a7e166f9a fix(tests): use top_k instead of limit in test_server_params (#4820) 2026-04-14 01:11:31 +05:30
Kartik 85768e78e7 fix(docs): remove chrome extension cookbooks (#4813) 2026-04-13 22:14:00 +05:30
Yunsu 7b395f3bf7 fix(openai): make store opt-in so it stops leaking to non-OpenAI backends (#4757) 2026-04-13 21:53:22 +05:30
HUANG XIAO 4180409b09 fix(s3vectors): handle vector=None in update() to prevent boto3 validation error (#4594) 2026-04-13 21:12:49 +05:30
Joe Wu 649e719ce6 fix: LLM config manager falls back to userConf.url for baseURL (#4715) (#4761) 2026-04-13 20:46:18 +05:30
Kartik 1a53852d93 test: update valkey cluster search test to use top_k parameter (#4815) 2026-04-13 20:44:18 +05:30
Chinnu Abey ac9cdd4840 Fix incorrect use of SentenceTransformer for cross-encoder reranker models (#4806) 2026-04-13 20:14:22 +05:30
Swarnaprakash Udayakumar cf530c4bec feat(valkey): add cluster mode enabled (CME) support (#4759) 2026-04-13 20:07:11 +05:30
Saket Aryan 92b958c1cc chore: bump Python SDK to v2.0.0b0 and Node SDK to v3.0.0-beta.0 (#4810) 2026-04-13 15:51:33 +05:30
Asish Kumar c239d8a483 fix(client): prevent feedback telemetry TypeError (#4795) 2026-04-12 03:06:05 +05:30
74 changed files with 2392 additions and 419 deletions
+21
View File
@@ -50,4 +50,25 @@ Here are the parameters available for configuring Valkey:
| `hnsw_m` | Number of bi-directional links for HNSW | `16` |
| `hnsw_ef_construction` | Size of dynamic candidate list for HNSW | `200` |
| `hnsw_ef_runtime` | Size of dynamic candidate list for search | `10` |
| `cluster_mode` | Enable cluster mode for Valkey cluster (CME) deployments | `false` |
| `distance_metric` | Distance metric for vector similarity | `cosine` |
## Cluster Mode
To use Valkey with cluster mode enabled (CME), set `cluster_mode` to `true`:
```python
config = {
"vector_store": {
"provider": "valkey",
"config": {
"collection_name": "memories",
"valkey_url": "valkey://cluster-endpoint:6379",
"embedding_model_dims": 1536,
"cluster_mode": True
}
}
}
```
When cluster mode is enabled, the connector uses `ValkeyCluster` instead of the standalone client, which handles `MOVED`/`ASK` redirections automatically. Search queries are coordinated across all shards by the valkey-search module's built-in coordinator. See the [valkey-search documentation](https://github.com/valkey-io/valkey-search) for details on cluster mode behavior.
@@ -1,70 +0,0 @@
---
title: Browser Extension Memory
description: "Add Mem0's universal memory layer to Chrome chat surfaces."
---
Enhance your AI interactions with Mem0, a Chrome extension that introduces a universal memory layer across platforms like ChatGPT, Claude, and Perplexity. Mem0 ensures seamless context sharing, making your AI experiences more personalized and efficient.
<Note>
We now support Grok! The Mem0 Chrome Extension has been updated to work with Grok, bringing the same powerful memory capabilities to your Grok conversations.
</Note>
## Features
- **Universal Memory Layer**: Share context seamlessly across ChatGPT, Claude, Perplexity, and Grok.
- **Smart Context Detection**: Automatically captures relevant information from your conversations.
- **Intelligent Memory Retrieval**: Surfaces pertinent memories at the right time.
- **One-Click Sync**: Easily synchronize with existing ChatGPT memories.
- **Memory Dashboard**: Manage all your memories in one centralized location.
## Installation
You can install the Mem0 Chrome Extension using one of the following methods:
### Method 1: Chrome Web Store Installation
1. **Download the Extension**: Open Google Chrome and navigate to the [Mem0 Chrome Extension page](https://chromewebstore.google.com/detail/mem0/onihkkbipkfeijkadecaafbgagkhglop?hl=en).
2. **Add to Chrome**: Click on the "Add to Chrome" button.
3. **Confirm Installation**: In the pop-up dialog, click "Add extension" to confirm. The Mem0 icon should now appear in your Chrome toolbar.
### Method 2: Manual Installation
1. **Download the Extension**: Clone or download the extension files from the [Mem0 Chrome Extension GitHub repository](https://github.com/mem0ai/mem0-chrome-extension).
2. **Access Chrome Extensions**: Open Google Chrome and navigate to `chrome://extensions`.
3. **Enable Developer Mode**: Toggle the "Developer mode" switch in the top right corner.
4. **Load Unpacked Extension**: Click "Load unpacked" and select the directory containing the extension files.
5. **Confirm Installation**: The Mem0 Chrome Extension should now appear in your Chrome toolbar.
## Usage
1. **Locate the Mem0 Icon**: After installation, find the Mem0 icon in your Chrome toolbar.
2. **Sign In**: Click the icon and sign in with your Google account.
3. **Interact with AI Assistants**:
- **ChatGPT and Perplexity**: Continue your conversations as usual; Mem0 operates seamlessly in the background.
- **Claude**: Click the Mem0 button or use the shortcut `Ctrl + M` to activate memory functions.
## Configuration
- **API Key**: Obtain your API key from the Mem0 Dashboard to connect the extension to the Mem0 API.
- **User ID**: This is your unique identifier in the Mem0 system. If not provided, it defaults to `chrome-extension-user`.
## Demo Video
<iframe width="700" height="400" src="https://www.youtube.com/embed/dqenCMMlfwQ?si=zhGVrkq6IS_0Jwyj" title="YouTube video player" frameborder="0" allow="accelerometer; autoplay; clipboard-write; encrypted-media; gyroscope; picture-in-picture; web-share" referrerpolicy="strict-origin-when-cross-origin" allowfullscreen></iframe>
## Privacy and Data Security
Your messages are sent to the Mem0 API for extracting and retrieving memories. Mem0 is committed to ensuring your data's privacy and security.
---
<CardGroup cols={2}>
<Card title="Build a Mem0 Companion" icon="users" href="/cookbooks/essentials/building-ai-companion">
Learn the foundations of memory-powered assistants that work across platforms.
</Card>
<Card title="Multimodal Support" icon="image" href="/platform/features/multimodal-support">
Extend your browser interactions with vision and audio memory.
</Card>
</CardGroup>
+1 -4
View File
@@ -201,10 +201,7 @@ Here are some examples of how Mem0 can be integrated into various applications:
>
Persistent personality for Eliza agents.
</Card>
<Card title="Browser Extension Memory" icon="globe" href="/cookbooks/frameworks/chrome-extension">
Universal memory layer for Chrome.
</Card>
</CardGroup>
</CardGroup>
---
+6 -3
View File
@@ -373,7 +373,6 @@
"cookbooks/frameworks/llamaindex-multiagent",
"cookbooks/frameworks/multimodal-retrieval",
"cookbooks/frameworks/eliza-os-character",
"cookbooks/frameworks/chrome-extension",
"cookbooks/frameworks/gemini-3-with-mem0-mcp",
"cookbooks/frameworks/mirofish-swarm-memory"
]
@@ -798,7 +797,11 @@
},
{
"source": "/examples/chrome-extension",
"destination": "/cookbooks/frameworks/chrome-extension"
"destination": "/cookbooks/overview"
},
{
"source": "/cookbooks/frameworks/chrome-extension",
"destination": "/cookbooks/overview"
},
{
"source": "/examples",
@@ -838,7 +841,7 @@
},
{
"source": "/v0x/examples/chrome-extension",
"destination": "/cookbooks/frameworks/chrome-extension"
"destination": "/cookbooks/overview"
},
{
"source": "/v0x/examples/youtube-assistant",
-1
View File
@@ -245,7 +245,6 @@ Key differentiators:
- [LlamaIndex Multiagent](https://docs.mem0.ai/cookbooks/frameworks/llamaindex-multiagent): Multi-agent systems with shared memory
- [Multimodal Retrieval](https://docs.mem0.ai/cookbooks/frameworks/multimodal-retrieval): Memory systems handling text, images, and documents
- [Eliza OS Character](https://docs.mem0.ai/cookbooks/frameworks/eliza-os-character): Character-based AI with persistent personality
- [Chrome Extension](https://docs.mem0.ai/cookbooks/frameworks/chrome-extension): Browser extensions that remember user interactions
- [Gemini with Mem0 MCP](https://docs.mem0.ai/cookbooks/frameworks/gemini-3-with-mem0-mcp): Google Gemini integration using MCP server
- [Mirofish Swarm Memory](https://docs.mem0.ai/cookbooks/frameworks/mirofish-swarm-memory): Swarm-based multi-agent memory patterns
+12 -11
View File
@@ -1048,8 +1048,8 @@
},
"expiration_date": {
"type": "string",
"format": "date-time",
"description": "The date and time when the memory will expire. Format: YYYY-MM-DD.",
"format": "date",
"description": "The date when the memory will expire. Format: YYYY-MM-DD.",
"title": "Expiration date",
"nullable": true,
"default": null
@@ -1233,7 +1233,7 @@
"memories"
],
"description": "Delete memories by filter. At least one filter is required \u2014 previously omitting all filters silently deleted everything; now it returns a validation error.",
"operationId": "memories_delete",
"operationId": "memories_delete_all",
"parameters": [
{
"name": "user_id",
@@ -1393,8 +1393,8 @@
},
"expiration_date": {
"type": "string",
"format": "date-time",
"description": "The date and time when the memory will expire. Format: YYYY-MM-DD.",
"format": "date",
"description": "The date when the memory will expire. Format: YYYY-MM-DD.",
"title": "Expiration date",
"nullable": true,
"default": null
@@ -1540,8 +1540,8 @@
},
"expiration_date": {
"type": "string",
"format": "date-time",
"description": "The date and time when the memory will expire. Format: YYYY-MM-DD.",
"format": "date",
"description": "The date when the memory will expire. Format: YYYY-MM-DD.",
"title": "Expiration date",
"nullable": true,
"default": null
@@ -1675,8 +1675,8 @@
},
"expiration_date": {
"type": "string",
"format": "date-time",
"description": "The date and time when the memory will expire. Format: YYYY-MM-DD.",
"format": "date",
"description": "The date when the memory will expire. Format: YYYY-MM-DD.",
"title": "Expiration date",
"nullable": true,
"default": null
@@ -1739,7 +1739,7 @@
"tags": [
"memories"
],
"operationId": "memories_read",
"operationId": "memories_entity_read",
"responses": {
"200": {
"description": "Successfully retrieved memories.",
@@ -5176,9 +5176,10 @@
"nullable": true
},
"expiration_date": {
"description": "The date and time when the memory will expire. Format: YYYY-MM-DD",
"description": "The date when the memory will expire. Format: YYYY-MM-DD",
"title": "Expiration date",
"type": "string",
"format": "date",
"nullable": true
},
"org_id": {
+2 -2
View File
@@ -8,6 +8,6 @@
},
"homepage": "https://mem0.ai",
"repository": "https://github.com/mem0ai/mem0",
"logo": "logo.svg",
"license": "Apache-2.0"
"license": "Apache-2.0",
"keywords": ["memory", "personalization", "mcp", "semantic-search"]
}
-1
View File
@@ -1,5 +1,4 @@
{
"description": "Mem0 memory capture hooks — automatic memory extraction at key lifecycle points",
"hooks": {
"SessionStart": [
{
+1 -1
View File
@@ -1,6 +1,6 @@
{
"name": "mem0ai",
"version": "2.4.6",
"version": "3.0.0-beta.2",
"description": "The Memory Layer For Your AI Apps",
"main": "./dist/index.js",
"module": "./dist/index.mjs",
+5 -5
View File
@@ -1,6 +1,7 @@
import axios from "axios";
import {
AllUsers,
PaginatedMemories,
ProjectOptions,
Memory,
MemoryHistory,
@@ -221,7 +222,7 @@ export default class MemoryClient {
this._captureEvent("add", [payloadKeys]);
const response = await this._fetchWithErrorHandling(
`${this.host}/v3/memories/`,
`${this.host}/v3/memories/add/`,
{
method: "POST",
headers: this.headers,
@@ -284,7 +285,7 @@ export default class MemoryClient {
);
}
async getAll(options?: GetAllMemoryOptions): Promise<Array<Memory>> {
async getAll(options?: GetAllMemoryOptions): Promise<PaginatedMemories> {
// Reject top-level entity params - must use filters instead
rejectTopLevelEntityParams(options as Record<string, any>, "getAll");
@@ -293,12 +294,11 @@ export default class MemoryClient {
this._captureEvent("get_all", [payloadKeys]);
const { page, pageSize, filters, ...rest } = options ?? {};
const body: Record<string, any> = {
output_format: "v1.1",
...camelToSnakeKeys(rest),
...(filters && { filters }),
};
let url = `${this.host}/v2/memories/`;
let url = `${this.host}/v3/memories/`;
if (page && pageSize) {
url += `?page=${page}&page_size=${pageSize}`;
}
@@ -308,7 +308,7 @@ export default class MemoryClient {
headers: this.headers,
body: JSON.stringify(body),
});
return Array.isArray(response) ? response : (response?.results ?? response);
return response;
}
async search(
+7
View File
@@ -142,6 +142,13 @@ export interface AllUsers {
previous: any;
}
export interface PaginatedMemories {
count: number;
next: string | null;
previous: string | null;
results: Array<Memory>;
}
export interface ProjectResponse {
customInstructions?: string;
customCategories?: string[];
@@ -26,7 +26,7 @@ describeIntegration("MemoryClient Integration — Batch Operations", () => {
beforeAll(async () => {
cleanup = suppressTelemetryNoise();
client = createTestClient();
client = await createTestClient();
memoryIds = await seedTestMemories(client, TEST_USER_ID);
});
@@ -26,9 +26,9 @@ describeIntegration("MemoryClient Integration — CRUD", () => {
let cleanup: () => void;
let memoryIds: string[] = [];
beforeAll(() => {
beforeAll(async () => {
cleanup = suppressTelemetryNoise();
client = createTestClient();
client = await createTestClient();
});
afterAll(async () => {
@@ -120,14 +120,20 @@ describeIntegration("MemoryClient Integration — CRUD", () => {
// ─── Get all ──────────────────────────────────────────────
describe("get all memories", () => {
test("returns all memories for test user", async () => {
const memories = await client.getAll({
const response = await client.getAll({
filters: { user_id: TEST_USER_ID },
});
expect(Array.isArray(memories)).toBe(true);
expect(memories.length).toBeGreaterThanOrEqual(memoryIds.length);
// Paginated shape: { count, next, previous, results: [...] }
expect(response).toHaveProperty("count");
expect(response).toHaveProperty("next");
expect(response).toHaveProperty("previous");
expect(response).toHaveProperty("results");
expect(typeof response.count).toBe("number");
expect(Array.isArray(response.results)).toBe(true);
expect(response.results.length).toBeGreaterThanOrEqual(memoryIds.length);
for (const mem of memories) {
for (const mem of response.results) {
expect(typeof mem.id).toBe("string");
expect(typeof mem.memory).toBe("string");
}
@@ -140,8 +146,12 @@ describeIntegration("MemoryClient Integration — CRUD", () => {
pageSize: 1,
});
// Paginated response is an object with results array
expect(page1).toBeDefined();
expect(page1).toHaveProperty("count");
expect(page1).toHaveProperty("next");
expect(page1).toHaveProperty("previous");
expect(Array.isArray(page1.results)).toBe(true);
expect(page1.results.length).toBeLessThanOrEqual(1);
});
});
@@ -197,13 +207,15 @@ describeIntegration("MemoryClient Integration — CRUD", () => {
expect(result).toHaveProperty("eventId");
});
test("getAll for non-existent user returns empty array", async () => {
const memories = await client.getAll({
test("getAll for non-existent user returns empty results", async () => {
const response = await client.getAll({
filters: { user_id: `nonexistent-user-${randomUUID()}` },
});
expect(Array.isArray(memories)).toBe(true);
expect(memories.length).toBe(0);
expect(response).toHaveProperty("results");
expect(Array.isArray(response.results)).toBe(true);
expect(response.results.length).toBe(0);
expect(response.count).toBe(0);
});
test("deleteAll for non-existent user does not throw", async () => {
@@ -18,9 +18,13 @@ export const describeIntegration = API_KEY ? describe : describe.skip;
* Create a MemoryClient with the real API key.
* Call this inside beforeAll — not at module scope — so it only
* runs when the suite is not skipped.
*
* Returns an initialized client ready for immediate use.
*/
export function createTestClient(): MemoryClient {
return new MemoryClient({ apiKey: API_KEY! });
export async function createTestClient(): Promise<MemoryClient> {
const client = new MemoryClient({ apiKey: API_KEY! });
await client.ping();
return client;
}
/**
@@ -63,10 +67,11 @@ export async function waitForMemories(
maxRetries = 4,
): Promise<Memory[]> {
for (let attempt = 1; attempt <= maxRetries; attempt++) {
const memories = await withRetry(() =>
const response = await withRetry(() =>
client.getAll({ filters: { user_id: userId } }),
);
if (Array.isArray(memories) && memories.length >= minCount) {
const memories = response.results ?? [];
if (memories.length >= minCount) {
return memories;
}
if (attempt < maxRetries) {
@@ -24,9 +24,9 @@ describeIntegration("MemoryClient Integration — Initialization", () => {
let client: MemoryClient;
let cleanup: () => void;
beforeAll(() => {
beforeAll(async () => {
cleanup = suppressTelemetryNoise();
client = createTestClient();
client = await createTestClient();
});
afterAll(() => cleanup());
@@ -27,7 +27,7 @@ describeIntegration("MemoryClient Integration — Users & Project", () => {
beforeAll(async () => {
cleanup = suppressTelemetryNoise();
client = createTestClient();
client = await createTestClient();
await seedTestMemories(client, TEST_USER_ID);
});
@@ -27,7 +27,7 @@ describeIntegration("MemoryClient Integration — Search & History", () => {
beforeAll(async () => {
cleanup = suppressTelemetryNoise();
client = createTestClient();
client = await createTestClient();
memoryIds = await seedTestMemories(client, TEST_USER_ID);
});
@@ -21,33 +21,33 @@ installConsoleSuppression();
// ─── add() ───────────────────────────────────────────────
describe("MemoryClient - add()", () => {
test("sends POST to /v3/memories/", async () => {
test("sends POST to /v3/memories/add/", async () => {
const extra = new Map<string, { status: number; body: unknown }>();
extra.set("/v3/memories/", { status: 200, body: [createMockMemory()] });
extra.set("/v3/memories/add/", { status: 200, body: [createMockMemory()] });
const mock = setupMockFetch(extra);
const client = new MemoryClient({ apiKey: TEST_API_KEY });
await client.add([{ role: "user", content: "Hello" }], { userId: "u1" });
expect(findFetchCall(mock, "/v3/memories/", "POST")).toBeDefined();
expect(findFetchCall(mock, "/v3/memories/add/", "POST")).toBeDefined();
});
test("includes messages in request body", async () => {
const messages = [{ role: "user" as const, content: "Hello, I am Alex" }];
const extra = new Map<string, { status: number; body: unknown }>();
extra.set("/v3/memories/", { status: 200, body: [createMockMemory()] });
extra.set("/v3/memories/add/", { status: 200, body: [createMockMemory()] });
const mock = setupMockFetch(extra);
const client = new MemoryClient({ apiKey: TEST_API_KEY });
await client.add(messages, { userId: "u1" });
const call = findFetchCall(mock, "/v3/memories/", "POST");
const call = findFetchCall(mock, "/v3/memories/add/", "POST");
expect(getFetchBody(call!).messages).toEqual(messages);
});
test("includes user_id in request body", async () => {
const extra = new Map<string, { status: number; body: unknown }>();
extra.set("/v3/memories/", { status: 200, body: [createMockMemory()] });
extra.set("/v3/memories/add/", { status: 200, body: [createMockMemory()] });
const mock = setupMockFetch(extra);
const client = new MemoryClient({ apiKey: TEST_API_KEY });
@@ -55,19 +55,19 @@ describe("MemoryClient - add()", () => {
user_id: "user_1",
});
const call = findFetchCall(mock, "/v3/memories/", "POST");
const call = findFetchCall(mock, "/v3/memories/add/", "POST");
expect(getFetchBody(call!).user_id).toBe("user_1");
});
test("sends empty messages array without crashing", async () => {
const extra = new Map<string, { status: number; body: unknown }>();
extra.set("/v3/memories/", { status: 200, body: [] });
extra.set("/v3/memories/add/", { status: 200, body: [] });
const mock = setupMockFetch(extra);
const client = new MemoryClient({ apiKey: TEST_API_KEY });
await client.add([], { userId: "u1" });
const call = findFetchCall(mock, "/v3/memories/", "POST");
const call = findFetchCall(mock, "/v3/memories/add/", "POST");
expect(getFetchBody(call!).messages).toEqual([]);
});
});
@@ -254,7 +254,7 @@ describe("MemoryClient - getAll() entity param rejection", () => {
test("accepts filters with user_id", async () => {
const extra = new Map<string, { status: number; body: unknown }>();
extra.set("/v2/memories/", {
extra.set("/v3/memories/", {
status: 200,
body: { results: [] },
});
@@ -262,6 +262,6 @@ describe("MemoryClient - getAll() entity param rejection", () => {
const client = new MemoryClient({ apiKey: TEST_API_KEY });
await client.getAll({ filters: { user_id: "u1" } });
expect(findFetchCall(mock, "/v2/memories/", "POST")).toBeDefined();
expect(findFetchCall(mock, "/v3/memories/", "POST")).toBeDefined();
});
});
+1 -1
View File
@@ -22,7 +22,7 @@ export const DEFAULT_MEMORY_CONFIG: MemoryConfig = {
config: {
baseURL: "https://api.openai.com/v1",
apiKey: process.env.OPENAI_API_KEY || "",
model: "gpt-4.1-nano-2025-04-14",
model: "gpt-5-mini",
modelProperties: undefined,
},
},
+1
View File
@@ -110,6 +110,7 @@ export class ConfigManager {
((userConf as Record<string, unknown>)?.lmstudio_base_url as
| string
| undefined) ??
userConf?.url ??
defaultConf.baseURL;
return {
+1 -1
View File
@@ -18,7 +18,7 @@ export class AzureOpenAILLM implements LLM {
endpoint: endpoint as string,
...rest,
});
this.model = config.model || "gpt-4";
this.model = config.model || "gpt-5-mini";
}
async generateResponse(
+1 -1
View File
@@ -11,7 +11,7 @@ export class OpenAILLM implements LLM {
apiKey: config.apiKey,
baseURL: config.baseURL,
});
this.model = config.model || "gpt-4.1-nano-2025-04-14";
this.model = config.model || "gpt-5-mini";
}
async generateResponse(
@@ -8,7 +8,7 @@ export class OpenAIStructuredLLM implements LLM {
constructor(config: LLMConfig) {
this.openai = new OpenAI({ apiKey: config.apiKey });
this.model = config.model || "gpt-4-turbo-preview";
this.model = config.model || "gpt-5-mini";
}
async generateResponse(
+315 -28
View File
@@ -52,6 +52,7 @@ import {
ENTITY_BOOST_WEIGHT,
ScoredResult,
} from "../utils/scoring";
import { getDefaultVectorStoreDbPath } from "../utils/sqlite";
// Entity params that must be passed via filters - check both snake_case and camelCase
const ENTITY_PARAMS = [
@@ -82,6 +83,58 @@ function rejectTopLevelEntityParams(
}
}
/**
* Validates and normalizes an entity ID.
* - Trims leading/trailing whitespace
* - Rejects empty or whitespace-only strings
* - Rejects strings containing internal whitespace
* @returns The trimmed entity ID, or undefined if input is undefined
* @throws Error if entity ID is invalid
*/
function validateAndTrimEntityId(
value: string | undefined,
name: string,
): string | undefined {
if (value === undefined) return undefined;
const trimmed = value.trim();
if (trimmed === "") {
throw new Error(
`Invalid ${name}: cannot be empty or whitespace-only. Provide a valid identifier.`,
);
}
if (/\s/.test(trimmed)) {
throw new Error(
`Invalid ${name}: cannot contain whitespace. Provide a valid identifier without spaces.`,
);
}
return trimmed;
}
/**
* Validates search parameters.
* @throws Error if threshold or topK are invalid
*/
function validateSearchParams(threshold?: number, topK?: number): void {
if (threshold !== undefined) {
if (typeof threshold !== "number" || isNaN(threshold)) {
throw new Error("threshold must be a valid number");
}
if (threshold < 0 || threshold > 1) {
throw new Error(
`Invalid threshold: ${threshold}. Must be between 0 and 1 (inclusive).`,
);
}
}
if (topK !== undefined) {
if (typeof topK !== "number" || isNaN(topK) || !Number.isInteger(topK)) {
throw new Error("topK must be a valid integer");
}
if (topK < 0) {
throw new Error(`Invalid topK: ${topK}. Must be a non-negative integer.`);
}
}
}
export class Memory {
private config: MemoryConfig;
private customInstructions: string | undefined;
@@ -195,15 +248,10 @@ export class Memory {
...this.config.vectorStore.config,
collectionName: entityCollectionName,
};
// For file-based stores (memory/SQLite), use a separate DB path for entities.
// If dbPath was set explicitly, derive the entity path from it. If it's unset,
// leave it unset — the vector store's default-path logic now scopes by
// collectionName, so the entity store will land in its own file automatically.
if (entityConfig.dbPath) {
entityConfig.dbPath = entityConfig.dbPath.replace(
/\.db$/,
"_entities.db",
);
// For file-based stores (memory/SQLite), always use a separate DB for entities
if (this.config.vectorStore.provider === "memory") {
const basePath = entityConfig.dbPath || getDefaultVectorStoreDbPath();
entityConfig.dbPath = basePath.replace(/\.db$/, "_entities.db");
}
this._entityStore = VectorStoreFactory.create(
this.config.vectorStore.provider,
@@ -214,6 +262,181 @@ export class Memory {
return this._entityStore;
}
/**
* Normalize a filters object for entity-store scoping: keeps only
* user_id/agent_id/run_id keys whose values are defined.
*/
private _sessionFiltersFromPayload(
payload: Record<string, any>,
): Record<string, any> {
const filters: Record<string, any> = {};
if (payload.user_id) filters.user_id = payload.user_id;
if (payload.agent_id) filters.agent_id = payload.agent_id;
if (payload.run_id) filters.run_id = payload.run_id;
return filters;
}
/**
* Remove `memoryId` from every entity record scoped to `filters`.
* If an entity's `linkedMemoryIds` becomes empty after removal, the
* entity record itself is deleted. Errors on individual entities are
* swallowed so one bad record does not break the whole operation.
*
* No-op if the entity store has not been initialized yet.
*/
private async _removeMemoryFromEntityStore(
memoryId: string,
filters: Record<string, any>,
): Promise<void> {
let entityStore: VectorStore;
try {
entityStore = await this.getEntityStore();
} catch (e) {
console.debug(`Entity store unavailable during cleanup: ${e}`);
return;
}
let rows: Array<{ id: string; payload: Record<string, any> }> = [];
try {
const listed = await entityStore.list(filters, 10000);
rows = (
Array.isArray(listed) && Array.isArray(listed[0])
? listed[0]
: (listed as any)
) as Array<{ id: string; payload: Record<string, any> }>;
} catch (e) {
console.debug(`Entity store list failed during cleanup: ${e}`);
return;
}
for (const row of rows) {
try {
const payload = row.payload || {};
const linked: string[] = Array.isArray(payload.linkedMemoryIds)
? payload.linkedMemoryIds
: [];
if (!linked.includes(memoryId)) continue;
const remaining = linked.filter((id) => id !== memoryId);
if (remaining.length === 0) {
try {
await entityStore.delete(row.id);
} catch (e) {
console.debug(`Entity delete failed for id=${row.id}: ${e}`);
}
} else {
const newPayload = { ...payload, linkedMemoryIds: remaining };
// entityStore.update requires a vector — re-embed entity text.
const entityText =
typeof payload.data === "string" ? payload.data : "";
if (!entityText) {
// Can't re-embed without text; skip gracefully.
console.debug(
`Entity id=${row.id} missing 'data'; skipping update during cleanup`,
);
continue;
}
let vec: number[];
try {
vec = await this.embedder.embed(entityText);
} catch (e) {
console.debug(`Entity re-embed failed for '${entityText}': ${e}`);
continue;
}
try {
await entityStore.update(row.id, vec, newPayload);
} catch (e) {
console.debug(`Entity update failed for id=${row.id}: ${e}`);
}
}
} catch (e) {
console.debug(`Entity cleanup error for id=${row?.id}: ${e}`);
}
}
}
/**
* Extract entities from `text` and link them to `memoryId` in the
* entity store, scoped to `filters` (user_id / agent_id / run_id).
*
* Simpler single-memory variant of Phase 7 in add(): no cross-memory
* dedup, but still does per-entity "search for existing, update if
* match >= 0.95 else insert new". Non-fatal errors are swallowed.
*/
private async _linkEntitiesForMemory(
memoryId: string,
text: string,
filters: Record<string, any>,
): Promise<void> {
try {
const entities = extractEntities(text);
if (entities.length === 0) return;
const entityStore = await this.getEntityStore();
for (const entity of entities) {
try {
let entityVec: number[];
try {
entityVec = await this.embedder.embed(entity.text);
} catch (e) {
console.debug(`Entity embed failed for '${entity.text}': ${e}`);
continue;
}
let matches: Array<{
id: string;
score?: number;
payload: Record<string, any>;
}> = [];
try {
matches = await entityStore.search(entityVec, 1, filters);
} catch {}
if (matches.length > 0 && (matches[0].score ?? 0) >= 0.95) {
const match = matches[0];
const payload = match.payload || {};
const linked = new Set<string>(
Array.isArray(payload.linkedMemoryIds)
? payload.linkedMemoryIds
: [],
);
linked.add(memoryId);
payload.linkedMemoryIds = Array.from(linked).sort();
try {
await entityStore.update(match.id, entityVec, payload);
} catch (e) {
console.debug(`Entity update failed for '${entity.text}': ${e}`);
}
} else {
const entityPayload: Record<string, any> = {
data: entity.text,
entityType: entity.type,
linkedMemoryIds: [memoryId],
};
if (filters.user_id) entityPayload.user_id = filters.user_id;
if (filters.agent_id) entityPayload.agent_id = filters.agent_id;
if (filters.run_id) entityPayload.run_id = filters.run_id;
try {
await entityStore.insert(
[entityVec],
[uuidv4()],
[entityPayload],
);
} catch (e) {
console.debug(`Entity insert failed for '${entity.text}': ${e}`);
}
}
} catch (e) {
console.debug(`Entity link error for '${entity.text}': ${e}`);
}
}
} catch (e) {
console.warn(`Entity linking failed during update: ${e}`);
}
}
private buildSessionScope(filters: SearchFilters): string {
const parts: string[] = [];
for (const key of ["agent_id", "run_id", "user_id"].sort()) {
@@ -279,6 +502,13 @@ export class Memory {
messages: string | Message[],
config: AddMemoryOptions,
): Promise<SearchResult> {
// Validate messages input
if (messages === undefined || messages === null) {
throw new Error(
"messages is required and cannot be undefined or null. Provide a string or array of messages.",
);
}
await this._ensureInitialized();
await this._captureEvent("add", {
message_count: Array.isArray(messages) ? messages.length : 1,
@@ -286,14 +516,12 @@ export class Memory {
has_filters: !!config.filters,
infer: config.infer,
});
const {
userId,
agentId,
runId,
metadata = {},
filters = {},
infer = true,
} = config;
const { metadata = {}, filters = {}, infer = true } = config;
// Validate and trim entity IDs
const userId = validateAndTrimEntityId(config.userId, "userId");
const agentId = validateAndTrimEntityId(config.agentId, "agentId");
const runId = validateAndTrimEntityId(config.runId, "runId");
// Convert camelCase entity params to snake_case for storage (matches API and search/getAll filters)
if (userId) filters.user_id = metadata.user_id = userId;
@@ -807,14 +1035,38 @@ export class Memory {
// Reject top-level entity params - must use filters instead
rejectTopLevelEntityParams(config as Record<string, any>, "search");
// Validate search parameters (before applying defaults)
validateSearchParams(config.threshold, config.topK);
// Validate and trim entity IDs in filters. Only include keys whose
// validated value is defined — otherwise downstream vector stores
// receive `agent_id: undefined` / `run_id: undefined` and fail
// (Qdrant rejects the malformed match, pgvector binds NULL, Redis
// emits a literal "undefined" string in TAG filters).
const normalizedFilters: Record<string, any> = config.filters
? Object.fromEntries(
Object.entries({
...config.filters,
user_id: validateAndTrimEntityId(config.filters.user_id, "user_id"),
agent_id: validateAndTrimEntityId(
config.filters.agent_id,
"agent_id",
),
run_id: validateAndTrimEntityId(config.filters.run_id, "run_id"),
}).filter(([, v]) => v !== undefined),
)
: {};
await this._ensureInitialized();
const { topK = 20, threshold = 0.1 } = config;
await this._captureEvent("search", {
query_length: query.length,
topK: config.topK,
topK,
has_filters: !!config.filters,
});
const { topK = 100, threshold = 0.1 } = config;
let effectiveFilters: Record<string, any> = { ...(config.filters || {}) };
let effectiveFilters: Record<string, any> = { ...normalizedFilters };
// Apply enhanced metadata filtering if advanced operators are detected
if (this._hasAdvancedOperators(effectiveFilters)) {
@@ -996,12 +1248,9 @@ export class Memory {
createdAt: payload.createdAt,
updatedAt: payload.updatedAt,
score: scored.score,
metadata: {
...Object.entries(payload)
.filter(([key]) => !excludedKeys.has(key))
.reduce((acc, [key, value]) => ({ ...acc, [key]: value }), {}),
scoreBreakdown: scored.scoreBreakdown,
},
metadata: Object.entries(payload)
.filter(([key]) => !excludedKeys.has(key))
.reduce((acc, [key, value]) => ({ ...acc, [key]: value }), {}),
...(payload.user_id && { user_id: payload.user_id }),
...(payload.agent_id && { agent_id: payload.agent_id }),
...(payload.run_id && { run_id: payload.run_id }),
@@ -1119,11 +1368,27 @@ export class Memory {
// Reject top-level entity params - must use filters instead
rejectTopLevelEntityParams(config as Record<string, any>, "getAll");
// Validate topK if provided (before applying defaults)
validateSearchParams(undefined, config.topK);
await this._ensureInitialized();
const { topK = 100, filters = {} } = config;
const { topK = 20 } = config;
// Validate and trim entity IDs in filters. Drop keys that resolve to
// undefined so downstream vector stores don't receive
// `agent_id: undefined` / `run_id: undefined` and fail.
const filters: Record<string, any> = Object.fromEntries(
Object.entries({
...(config.filters || {}),
user_id: validateAndTrimEntityId(config.filters?.user_id, "user_id"),
agent_id: validateAndTrimEntityId(config.filters?.agent_id, "agent_id"),
run_id: validateAndTrimEntityId(config.filters?.run_id, "run_id"),
}).filter(([, v]) => v !== undefined),
);
await this._captureEvent("get_all", {
topK: topK,
topK,
has_user_id: !!filters.user_id,
has_agent_id: !!filters.agent_id,
has_run_id: !!filters.run_id,
@@ -1180,6 +1445,7 @@ export class Memory {
...metadata,
data,
hash: createHash("md5").update(data).digest("hex"),
textLemmatized: lemmatizeForBm25(data),
createdAt: new Date().toISOString(),
};
@@ -1237,6 +1503,16 @@ export class Memory {
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}`);
}
return memoryId;
}
@@ -1247,6 +1523,9 @@ export class Memory {
}
const prevValue = existingMemory.payload.data;
const sessionFilters = this._sessionFiltersFromPayload(
existingMemory.payload || {},
);
await this.vectorStore.delete(memoryId);
await this.db.addHistory(
memoryId,
@@ -1258,6 +1537,14 @@ export class Memory {
1,
);
// Entity-store cleanup: strip this memory's id from any entity records
// that linked to it. Non-fatal — log and continue on error.
try {
await this._removeMemoryFromEntityStore(memoryId, sessionFilters);
} catch (e) {
console.warn(`Entity store cleanup failed during delete: ${e}`);
}
return memoryId;
}
@@ -197,7 +197,7 @@ describe("MemoryVectorStore (better-sqlite3)", () => {
expect(result).not.toBeNull();
expect(result!.id).toBe("id-1");
expect(result!.payload.data).toBe("hello");
expect(result!.payload.userId).toBe("u1");
expect(result!.payload.user_id).toBe("u1");
});
it("get returns null for non-existent id", async () => {
@@ -253,7 +253,7 @@ describe("MemoryVectorStore (better-sqlite3)", () => {
],
);
const results = await store.search(v, 10, { userId: "alice" });
const results = await store.search(v, 10, { user_id: "alice" });
expect(results).toHaveLength(1);
expect(results[0].id).toBe("id-1");
});
@@ -320,7 +320,7 @@ describe("MemoryVectorStore (better-sqlite3)", () => {
expect(all).toHaveLength(3);
expect(totalAll).toBe(3);
const [filtered, totalFiltered] = await store.list({ userId: "alice" });
const [filtered, totalFiltered] = await store.list({ user_id: "alice" });
expect(filtered).toHaveLength(2);
expect(totalFiltered).toBe(2);
});
@@ -308,7 +308,7 @@ describe("backward compat: MemoryVectorStore", () => {
);
const results = await store.search(normalize([1, 0, 0]), 10, {
userId: "user2",
user_id: "user2",
});
expect(results).toHaveLength(1);
expect(results[0].id).toBe("b");
@@ -660,6 +660,9 @@ export function extractEntities(text: string): ExtractedEntity[] {
txt = txt.replace(/\s*:+$/, "");
// Strip leading numbered list markers
txt = txt.replace(/^\d+\s*\.\s*/, "");
// Strip trailing sentence punctuation (".", ",", ";", "!", "?") — otherwise
// "Paris." and "Paris" produce different embeddings and break entity dedup.
txt = txt.replace(/[.,;!?]+$/, "").trim();
if (!txt || txt.length <= 2 || hasArtifacts(txt)) {
continue;
-10
View File
@@ -59,11 +59,6 @@ export function normalizeBm25(
export interface ScoredResult {
id: string;
score: number;
scoreBreakdown: {
semantic: number;
bm25: number;
entityBoost: number;
};
payload: Record<string, any>;
}
@@ -134,11 +129,6 @@ export function scoreAndRank(
scored.push({
id: memIdStr,
score: combined,
scoreBreakdown: {
semantic: semanticScore,
bm25: bm25Score,
entityBoost: entityBoost,
},
payload: result.payload,
});
}
+2 -10
View File
@@ -2,16 +2,8 @@ import fs from "fs";
import os from "os";
import path from "path";
export function getDefaultVectorStoreDbPath(collectionName?: string): string {
// Scope the default DB file by collection name so that parallel stores
// (e.g. "memories" vs "memories_entities") don't collide in the same
// SQLite table. Without this, both collections write to the same
// `vectors` table and search results leak across them.
const filename =
collectionName && collectionName.length > 0
? `vector_store_${collectionName.replace(/[^a-zA-Z0-9_-]/g, "_")}.db`
: "vector_store.db";
return path.join(os.homedir(), ".mem0", filename);
export function getDefaultVectorStoreDbPath(): string {
return path.join(os.homedir(), ".mem0", "vector_store.db");
}
export function ensureSQLiteDirectory(dbPath: string): void {
+24 -9
View File
@@ -19,12 +19,27 @@ export class MemoryVectorStore implements VectorStore {
private dimension: number;
private dbPath: string;
private static readonly CAMEL_TO_SNAKE: Record<string, string> = {
userId: "user_id",
agentId: "agent_id",
runId: "run_id",
};
private normalizePayload(payload: Record<string, any>): Record<string, any> {
for (const [camel, snake] of Object.entries(
MemoryVectorStore.CAMEL_TO_SNAKE,
)) {
if (camel in payload && !(snake in payload)) {
payload[snake] = payload[camel];
delete payload[camel];
}
}
return payload;
}
constructor(config: VectorStoreConfig) {
this.dimension = config.dimension || 1536; // Default OpenAI dimension
// Scope the default path by collectionName so that parallel stores
// (e.g. memory vs entity) don't share the same `vectors` table.
this.dbPath =
config.dbPath || getDefaultVectorStoreDbPath(config.collectionName);
this.dbPath = config.dbPath || getDefaultVectorStoreDbPath();
if (!config.dbPath) {
const oldDefault = path.join(process.cwd(), "vector_store.db");
@@ -249,7 +264,7 @@ export class MemoryVectorStore implements VectorStore {
}[] = [];
for (const row of rows) {
const payload = JSON.parse(row.payload);
const payload = this.normalizePayload(JSON.parse(row.payload));
const memoryVector: MemoryVector = {
id: row.id,
vector: Array.from(
@@ -263,7 +278,7 @@ export class MemoryVectorStore implements VectorStore {
};
if (this.filterVector(memoryVector, filters)) {
const text = payload.text_lemmatized || payload.data || "";
const text = payload.textLemmatized || payload.data || "";
candidates.push({ id: row.id, payload, tokens: this.tokenize(text) });
}
}
@@ -354,7 +369,7 @@ export class MemoryVectorStore implements VectorStore {
row.vector.byteOffset,
row.vector.byteLength / 4,
);
const payload = JSON.parse(row.payload);
const payload = this.normalizePayload(JSON.parse(row.payload));
const memoryVector: MemoryVector = {
id: row.id,
vector: Array.from(vector),
@@ -381,7 +396,7 @@ export class MemoryVectorStore implements VectorStore {
.get(vectorId) as any;
if (!row) return null;
const payload = JSON.parse(row.payload);
const payload = this.normalizePayload(JSON.parse(row.payload));
return {
id: row.id,
payload,
@@ -421,7 +436,7 @@ export class MemoryVectorStore implements VectorStore {
const results: VectorStoreResult[] = [];
for (const row of rows) {
const payload = JSON.parse(row.payload);
const payload = this.normalizePayload(JSON.parse(row.payload));
const memoryVector: MemoryVector = {
id: row.id,
vector: Array.from(
+10 -2
View File
@@ -22,6 +22,7 @@ export class PGVector implements VectorStore {
private useHnsw: boolean;
private readonly dbName: string;
private config: PGVectorConfig;
private _initPromise?: Promise<void>;
constructor(config: PGVectorConfig) {
this.collectionName = config.collectionName || "memories";
@@ -41,6 +42,13 @@ export class PGVector implements VectorStore {
}
async initialize(): Promise<void> {
if (!this._initPromise) {
this._initPromise = this._doInitialize();
}
return this._initPromise;
}
private async _doInitialize(): Promise<void> {
try {
await this.client.connect();
@@ -186,9 +194,9 @@ export class PGVector implements VectorStore {
: "";
const searchQuery = `
SELECT id, ts_rank_cd(to_tsvector('simple', payload->>'text_lemmatized'), plainto_tsquery('simple', $1)) AS score, payload
SELECT id, ts_rank_cd(to_tsvector('simple', payload->>'textLemmatized'), plainto_tsquery('simple', $1)) AS score, payload
FROM ${this.collectionName}
WHERE to_tsvector('simple', payload->>'text_lemmatized') @@ plainto_tsquery('simple', $1)
WHERE to_tsvector('simple', payload->>'textLemmatized') @@ plainto_tsquery('simple', $1)
${filterClause}
ORDER BY score DESC
LIMIT $2
+35 -14
View File
@@ -9,6 +9,19 @@ import type {
import { VectorStore } from "./base";
import { SearchFilters, VectorStoreConfig, VectorStoreResult } from "../types";
/**
* Escape RediSearch TAG filter special characters. Any punctuation in the
* value (including `-`, which appears in every UUID) must be backslash-
* escaped, otherwise RediSearch either parses it as an operator (`-` is
* minus, `|` is OR) or rejects the whole expression as a syntax error.
*/
function escapeRedisTagValue(value: unknown): string {
return String(value).replace(
/([,.<>{}\[\]"':;!@#$%^&*()\-+=~|/\\\s])/g,
"\\$1",
);
}
interface RedisConfig extends VectorStoreConfig {
redisUrl: string;
collectionName: string;
@@ -256,17 +269,25 @@ export class RedisDB implements VectorStore {
const modulesResponse =
(await this.client.moduleList()) as unknown as any[];
// Parse module list to find search module
const hasSearch = modulesResponse.some((module: any[]) => {
const moduleMap = new Map();
for (let i = 0; i < module.length; i += 2) {
moduleMap.set(module[i], module[i + 1]);
const hasSearch = modulesResponse.some((mod: any) => {
// node-redis v4+ returns objects: { name: "search", ver: ..., ... }
if (typeof mod === "object" && !Array.isArray(mod) && mod.name) {
const name = String(mod.name).toLowerCase();
return name === "search" || name === "searchlight";
}
const moduleName = moduleMap.get("name");
return (
moduleName?.toLowerCase() === "search" ||
moduleName?.toLowerCase() === "searchlight"
);
// Fallback: legacy flat array format [key, value, key, value, ...]
if (Array.isArray(mod)) {
const moduleMap = new Map();
for (let i = 0; i < mod.length; i += 2) {
moduleMap.set(mod[i], mod[i + 1]);
}
const name = moduleMap.get("name");
return (
name?.toLowerCase() === "search" ||
name?.toLowerCase() === "searchlight"
);
}
return false;
});
if (!hasSearch) {
@@ -369,8 +390,8 @@ export class RedisDB implements VectorStore {
const snakeFilters = filters ? toSnakeCase(filters) : undefined;
const filterExpr = snakeFilters
? Object.entries(snakeFilters)
.filter(([_, value]) => value !== null)
.map(([key, value]) => `@${key}:{${value}}`)
.filter(([_, value]) => value !== null && value !== undefined)
.map(([key, value]) => `@${key}:{${escapeRedisTagValue(value)}}`)
.join(" ")
: "*";
@@ -607,8 +628,8 @@ export class RedisDB implements VectorStore {
const snakeFilters = filters ? toSnakeCase(filters) : undefined;
const filterExpr = snakeFilters
? Object.entries(snakeFilters)
.filter(([_, value]) => value !== null)
.map(([key, value]) => `@${key}:{${value}}`)
.filter(([_, value]) => value !== null && value !== undefined)
.map(([key, value]) => `@${key}:{${escapeRedisTagValue(value)}}`)
.join(" ")
: "*";
@@ -127,6 +127,20 @@ describe("ConfigManager", () => {
expect(config.llm.config.url).toBe("http://fallback:11434");
});
it("should use url as baseURL fallback when no baseURL provided (issue #4715)", () => {
const config = ConfigManager.mergeConfig({
embedder: baseEmbedder,
vectorStore: baseVectorStore,
llm: {
provider: "ollama",
config: { model: "llama3.1:8b", url: "http://my-ollama-host:11434" },
},
});
expect(config.llm.config.baseURL).toBe("http://my-ollama-host:11434");
expect(config.llm.config.url).toBe("http://my-ollama-host:11434");
});
it("should use default baseURL when no url or baseURL provided", () => {
const config = ConfigManager.mergeConfig({
embedder: baseEmbedder,
@@ -469,7 +469,7 @@ describe("Memory – auto-initialization", () => {
},
llm: {
provider: "openai",
config: { apiKey: "sk-fake", model: "gpt-4-turbo-preview" },
config: { apiKey: "sk-fake", model: "gpt-5-mini" },
},
historyDbPath: ":memory:",
disableHistory: true,
+1 -1
View File
@@ -75,7 +75,7 @@ function createMemory(overrides: Partial<MemoryConfig> = {}): Memory {
},
llm: {
provider: "openai",
config: { apiKey: "test-key", model: "gpt-4-turbo-preview" },
config: { apiKey: "test-key", model: "gpt-5-mini" },
},
historyDbPath: ":memory:",
...overrides,
+1 -1
View File
@@ -75,7 +75,7 @@ function createMemory(): Memory {
},
llm: {
provider: "openai",
config: { apiKey: "test-key", model: "gpt-4-turbo-preview" },
config: { apiKey: "test-key", model: "gpt-5-mini" },
},
historyDbPath: ":memory:",
});
+2 -2
View File
@@ -69,7 +69,7 @@ function createMemory(overrides: Partial<MemoryConfig> = {}): Memory {
},
llm: {
provider: "openai",
config: { apiKey: "test-key", model: "gpt-4-turbo-preview" },
config: { apiKey: "test-key", model: "gpt-5-mini" },
},
historyDbPath: ":memory:",
...overrides,
@@ -94,7 +94,7 @@ describe("Memory - Initialization", () => {
},
llm: {
provider: "openai",
config: { apiKey: "test-key", model: "gpt-4" },
config: { apiKey: "test-key", model: "gpt-5-mini" },
},
};
const mem = Memory.fromConfig(config);
@@ -0,0 +1,299 @@
/**
* Unit tests for OSS SDK input validation.
*
* Validates fixes for:
* - Undefined/null message handling in add()
* - Threshold bounds validation (must be 0-1) in search()
* - TopK validation (must be non-negative) in search() and getAll()
* - Whitespace-only entity ID rejection in add(), search(), getAll()
*/
/// <reference types="jest" />
import { Memory } from "../src/memory";
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(),
}));
jest.mock("../src/llms/openai", () => ({
OpenAILLM: jest.fn().mockImplementation(() => ({
generateResponse: jest.fn().mockResolvedValue(
JSON.stringify({
memory: [{ id: "0", text: "test memory", attributed_to: "user" }],
}),
),
})),
}));
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().mockResolvedValue([mockEmbedding]),
})),
}));
describe("Memory Input Validation", () => {
let memory: Memory;
const testUserId = "test-user-validation";
beforeAll(async () => {
memory = new Memory({
version: "v1.1",
embedder: {
provider: "openai",
config: { apiKey: "test-key", model: "text-embedding-3-small" },
},
llm: {
provider: "openai",
config: { apiKey: "test-key", model: "gpt-5-mini" },
},
vectorStore: {
provider: "memory",
config: { collectionName: "validation-test" },
},
});
// Wait for initialization
await new Promise((resolve) => setTimeout(resolve, 2000));
});
afterAll(async () => {
try {
await memory.reset();
} catch (e) {
// ignore cleanup errors
}
});
describe("add() message validation", () => {
it("should throw error when messages is undefined", async () => {
await expect(
// @ts-ignore - intentionally passing undefined
memory.add(undefined, { userId: testUserId }),
).rejects.toThrow("messages is required");
});
it("should throw error when messages is null", async () => {
await expect(
// @ts-ignore - intentionally passing null
memory.add(null, { userId: testUserId }),
).rejects.toThrow("messages is required");
});
});
describe("search() threshold validation", () => {
it("should throw error when threshold > 1.0", async () => {
await expect(
memory.search("test query", {
filters: { user_id: testUserId },
threshold: 1.5,
}),
).rejects.toThrow("Invalid threshold");
});
it("should throw error when threshold = 1.1", async () => {
await expect(
memory.search("test query", {
filters: { user_id: testUserId },
threshold: 1.1,
}),
).rejects.toThrow("Invalid threshold");
});
it("should throw error when threshold is negative", async () => {
await expect(
memory.search("test query", {
filters: { user_id: testUserId },
threshold: -0.5,
}),
).rejects.toThrow("Invalid threshold");
});
it("should throw error when threshold = -0.1", async () => {
await expect(
memory.search("test query", {
filters: { user_id: testUserId },
threshold: -0.1,
}),
).rejects.toThrow("Invalid threshold");
});
it("should accept threshold = 0 (edge case)", async () => {
const result = await memory.search("test query", {
filters: { user_id: testUserId },
threshold: 0,
});
expect(result).toBeDefined();
expect(result.results).toBeDefined();
});
it("should accept threshold = 1.0 (edge case)", async () => {
const result = await memory.search("test query", {
filters: { user_id: testUserId },
threshold: 1.0,
});
expect(result).toBeDefined();
expect(result.results).toBeDefined();
});
it("should accept threshold = 0.5 (normal valid value)", async () => {
const result = await memory.search("test query", {
filters: { user_id: testUserId },
threshold: 0.5,
});
expect(result).toBeDefined();
expect(result.results).toBeDefined();
});
});
describe("search() topK validation", () => {
it("should throw error when topK is negative", async () => {
await expect(
memory.search("test query", {
filters: { user_id: testUserId },
topK: -5,
}),
).rejects.toThrow("Invalid topK");
});
it("should throw error when topK = -1", async () => {
await expect(
memory.search("test query", {
filters: { user_id: testUserId },
topK: -1,
}),
).rejects.toThrow("Invalid topK");
});
it("should accept topK = 0 (returns empty)", async () => {
const result = await memory.search("test query", {
filters: { user_id: testUserId },
topK: 0,
});
expect(result).toBeDefined();
expect(result.results).toBeDefined();
});
it("should accept topK = 20 (normal value)", async () => {
const result = await memory.search("test query", {
filters: { user_id: testUserId },
topK: 20,
});
expect(result).toBeDefined();
expect(result.results).toBeDefined();
});
});
describe("add() entity ID validation", () => {
it("should throw error when userId is whitespace-only", async () => {
await expect(
memory.add("test message", { userId: " " }),
).rejects.toThrow("Invalid userId");
});
it("should throw error when userId is tabs and newlines", async () => {
await expect(
memory.add("test message", { userId: "\t\n\t" }),
).rejects.toThrow("Invalid userId");
});
it("should throw error when agentId is whitespace-only", async () => {
await expect(
memory.add("test message", { agentId: " " }),
).rejects.toThrow("Invalid agentId");
});
it("should throw error when runId is whitespace-only", async () => {
await expect(
memory.add("test message", { runId: " " }),
).rejects.toThrow("Invalid runId");
});
it("should throw error when userId contains internal whitespace", async () => {
await expect(
memory.add("test message", { userId: "user 123" }),
).rejects.toThrow("Invalid userId: cannot contain whitespace");
});
it("should throw error when userId contains tab character", async () => {
await expect(
memory.add("test message", { userId: "user\t123" }),
).rejects.toThrow("Invalid userId: cannot contain whitespace");
});
it("should accept userId with leading/trailing whitespace (trimmed)", async () => {
// Should not throw - leading/trailing whitespace is trimmed
const result = await memory.add("test message", {
userId: " valid-user ",
});
expect(result).toBeDefined();
});
});
describe("search() filter entity ID validation", () => {
it("should throw error when user_id in filters is whitespace-only", async () => {
await expect(
memory.search("test query", {
filters: { user_id: " " },
}),
).rejects.toThrow("Invalid user_id");
});
it("should throw error when agent_id in filters is whitespace-only", async () => {
await expect(
memory.search("test query", {
filters: { agent_id: " " },
}),
).rejects.toThrow("Invalid agent_id");
});
it("should throw error when user_id contains internal whitespace", async () => {
await expect(
memory.search("test query", {
filters: { user_id: "user 123" },
}),
).rejects.toThrow("Invalid user_id: cannot contain whitespace");
});
it("should accept user_id with leading/trailing whitespace (trimmed)", async () => {
const result = await memory.search("test query", {
filters: { user_id: " valid-user " },
});
expect(result).toBeDefined();
expect(result.results).toBeDefined();
});
});
describe("getAll() validation", () => {
it("should throw error when user_id is whitespace-only", async () => {
await expect(
memory.getAll({ filters: { user_id: " " } }),
).rejects.toThrow("Invalid user_id");
});
it("should throw error when topK is negative", async () => {
await expect(
memory.getAll({ filters: { user_id: testUserId }, topK: -1 }),
).rejects.toThrow("Invalid topK");
});
it("should throw error when user_id contains internal whitespace", async () => {
await expect(
memory.getAll({ filters: { user_id: "user 123" } }),
).rejects.toThrow("Invalid user_id: cannot contain whitespace");
});
it("should accept user_id with leading/trailing whitespace (trimmed)", async () => {
const result = await memory.getAll({
filters: { user_id: " valid-user " },
});
expect(result).toBeDefined();
expect(result.results).toBeDefined();
});
});
});
@@ -85,14 +85,16 @@ describe("MemoryVectorStore - search", () => {
expect(results).toHaveLength(1);
});
test("filters by userId", async () => {
const results = await store.search(vec([1, 0, 0, 0]), 10, { userId: "u2" });
expect(results.every((r) => r.payload.userId === "u2")).toBe(true);
test("filters by user_id", async () => {
const results = await store.search(vec([1, 0, 0, 0]), 10, {
user_id: "u2",
});
expect(results.every((r) => r.payload.user_id === "u2")).toBe(true);
});
test("returns empty when filter matches nothing", async () => {
const results = await store.search(vec([1, 0, 0, 0]), 10, {
userId: "nobody",
user_id: "nobody",
});
expect(results).toHaveLength(0);
});
@@ -168,10 +170,10 @@ describe("MemoryVectorStore - list", () => {
expect(results).toHaveLength(3);
});
test("filters by userId", async () => {
const [results, count] = await store.list({ userId: "u1" });
test("filters by user_id", async () => {
const [results, count] = await store.list({ user_id: "u1" });
expect(count).toBe(2);
expect(results.every((r) => r.payload.userId === "u1")).toBe(true);
expect(results.every((r) => r.payload.user_id === "u1")).toBe(true);
});
test("respects limit", async () => {
@@ -95,19 +95,19 @@ describe("MemoryVectorStore – full backward compat", () => {
expect(results.length).toBe(2);
expect(results[0].id).toBe("id-1");
// Search with filters
const filtered = await store.search(vec1, 2, { userId: "u1" });
// Search with filters (camelCase in payload is normalized to snake_case)
const filtered = await store.search(vec1, 2, { user_id: "u1" });
expect(filtered.length).toBe(2);
// Update
const vec3 = new Array(1536).fill(0);
vec3[2] = 1.0;
await store.update("id-1", vec3, { data: "updated", userId: "u1" });
await store.update("id-1", vec3, { data: "updated", user_id: "u1" });
const updated = await store.get("id-1");
expect(updated!.payload.data).toBe("updated");
// List
const [listed, count] = await store.list({ userId: "u1" });
const [listed, count] = await store.list({ user_id: "u1" });
expect(count).toBe(2);
// List with limit
+12 -22
View File
@@ -167,7 +167,7 @@ class MemoryClient:
kwargs = self._prepare_params(kwargs)
payload = self._prepare_payload(messages, kwargs)
response = self.client.post("/v3/memories/", json=payload)
response = self.client.post("/v3/memories/add/", json=payload)
response.raise_for_status()
if "metadata" in kwargs:
del kwargs["metadata"]
@@ -207,7 +207,7 @@ class MemoryClient:
**kwargs: Optional parameters for filtering (filters, page, page_size).
Returns:
A dictionary containing memories in v1.1 format: {"results": [...]}
A paginated dict: {"count": int, "next": str | None, "previous": str | None, "results": [...]}
Raises:
ValidationError: If the input data is invalid.
@@ -233,9 +233,9 @@ class MemoryClient:
"page": params.pop("page"),
"page_size": params.pop("page_size"),
}
response = self.client.post("/v2/memories/", json=params, params=query_params)
response = self.client.post("/v3/memories/", json=params, params=query_params)
else:
response = self.client.post("/v2/memories/", json=params)
response = self.client.post("/v3/memories/", json=params)
response.raise_for_status()
if "metadata" in kwargs:
del kwargs["metadata"]
@@ -247,12 +247,7 @@ class MemoryClient:
"sync_type": "sync",
},
)
result = response.json()
# Ensure v1.1 format (wrap raw list if needed)
if isinstance(result, list):
return {"results": result}
return result
return response.json()
@api_error_handler
def search(self, query: str, options: Optional[SearchMemoryOptions] = None, **kwargs) -> Dict[str, Any]:
@@ -883,7 +878,7 @@ class MemoryClient:
response = self.client.post("/v1/feedback/", json=data)
response.raise_for_status()
capture_client_event("client.feedback", self, data, {"sync_type": "sync"})
capture_client_event("client.feedback", self, {**data, "sync_type": "sync"})
return response.json()
def _prepare_payload(self, messages: List[Dict[str, str]], kwargs: Dict[str, Any]) -> Dict[str, Any]:
@@ -1095,7 +1090,7 @@ class AsyncMemoryClient:
kwargs = self._prepare_params(kwargs)
payload = self._prepare_payload(messages, kwargs)
response = await self.async_client.post("/v3/memories/", json=payload)
response = await self.async_client.post("/v3/memories/add/", json=payload)
response.raise_for_status()
if "metadata" in kwargs:
del kwargs["metadata"]
@@ -1119,7 +1114,7 @@ class AsyncMemoryClient:
**kwargs: Optional parameters for filtering (filters, page, page_size).
Returns:
A dictionary containing memories in v1.1 format: {"results": [...]}
A paginated dict: {"count": int, "next": str | None, "previous": str | None, "results": [...]}
Raises:
ValidationError: If the input data is invalid.
@@ -1145,9 +1140,9 @@ class AsyncMemoryClient:
"page": params.pop("page"),
"page_size": params.pop("page_size"),
}
response = await self.async_client.post("/v2/memories/", json=params, params=query_params)
response = await self.async_client.post("/v3/memories/", json=params, params=query_params)
else:
response = await self.async_client.post("/v2/memories/", json=params)
response = await self.async_client.post("/v3/memories/", json=params)
response.raise_for_status()
if "metadata" in kwargs:
del kwargs["metadata"]
@@ -1159,12 +1154,7 @@ class AsyncMemoryClient:
"sync_type": "async",
},
)
result = response.json()
# Ensure v1.1 format (wrap raw list if needed)
if isinstance(result, list):
return {"results": result}
return result
return response.json()
@api_error_handler
async def search(self, query: str, options: Optional[SearchMemoryOptions] = None, **kwargs) -> Dict[str, Any]:
@@ -1761,5 +1751,5 @@ class AsyncMemoryClient:
response = await self.async_client.post("/v1/feedback/", json=data)
response.raise_for_status()
capture_client_event("client.feedback", self, data, {"sync_type": "async"})
capture_client_event("client.feedback", self, {**data, "sync_type": "async"})
return response.json()
+6 -1
View File
@@ -29,7 +29,7 @@ class OpenAIConfig(BaseLlmConfig):
openrouter_base_url: Optional[str] = None,
site_url: Optional[str] = None,
app_name: Optional[str] = None,
store: bool = False,
store: Optional[bool] = None,
# Response monitoring callback
response_callback: Optional[Callable[[Any, dict, dict], None]] = None,
):
@@ -53,6 +53,11 @@ class OpenAIConfig(BaseLlmConfig):
openrouter_base_url: OpenRouter base URL, defaults to None
site_url: Site URL for OpenRouter, defaults to None
app_name: Application name for OpenRouter, defaults to None
store: Whether to store the conversation on OpenAI's server. Opt-in;
defaults to None (not sent). Set to True or False only if you
want the value forwarded to the OpenAI API. Leaving it None
avoids leaking the field into OpenAI-compatible backends that
reject unknown fields (Gemini, Groq, vLLM, etc.).
response_callback: Optional callback for monitoring LLM responses.
"""
# Initialize base parameters
+1
View File
@@ -14,6 +14,7 @@ class ValkeyConfig(BaseModel):
hnsw_m: int = Field(16, description="HNSW: number of connections per layer")
hnsw_ef_construction: int = Field(200, description="HNSW: search width during index construction")
hnsw_ef_runtime: int = Field(10, description="HNSW: search width during queries")
cluster_mode: bool = Field(False, description="Enable cluster mode for Valkey cluster (CME) deployments")
@model_validator(mode="before")
@classmethod
+1 -1
View File
@@ -39,7 +39,7 @@ class AzureOpenAILLM(LLMBase):
# Model name should match the custom deployment name chosen for it.
if not self.config.model:
self.config.model = "gpt-4.1-nano-2025-04-14"
self.config.model = "gpt-5-mini"
api_key = self.config.azure_kwargs.api_key or os.getenv("LLM_AZURE_OPENAI_API_KEY")
azure_deployment = self.config.azure_kwargs.azure_deployment or os.getenv("LLM_AZURE_DEPLOYMENT")
+1 -1
View File
@@ -18,7 +18,7 @@ class AzureOpenAIStructuredLLM(LLMBase):
# Model name should match the custom deployment name chosen for it.
if not self.config.model:
self.config.model = "gpt-4.1-nano-2025-04-14"
self.config.model = "gpt-5-mini"
api_key = self.config.azure_kwargs.api_key or os.getenv("LLM_AZURE_OPENAI_API_KEY")
azure_deployment = self.config.azure_kwargs.azure_deployment or os.getenv("LLM_AZURE_DEPLOYMENT")
+1 -1
View File
@@ -16,7 +16,7 @@ class LiteLLM(LLMBase):
super().__init__(config)
if not self.config.model:
self.config.model = "gpt-4.1-nano-2025-04-14"
self.config.model = "gpt-5-mini"
def _parse_response(self, response, tools):
"""
+7 -6
View File
@@ -36,7 +36,7 @@ class OpenAILLM(LLMBase):
super().__init__(config)
if not self.config.model:
self.config.model = "gpt-4.1-nano-2025-04-14"
self.config.model = "gpt-5-mini"
if os.environ.get("OPENROUTER_API_KEY"): # Use OpenRouter
self.client = OpenAI(
@@ -126,11 +126,12 @@ class OpenAILLM(LLMBase):
params.update(**openrouter_params)
else:
openai_specific_generation_params = ["store"]
for param in openai_specific_generation_params:
if hasattr(self.config, param):
params[param] = getattr(self.config, param)
# Only send OpenAI-specific parameters when the user has explicitly
# configured them. OpenAI-compatible backends (Gemini, Groq, vLLM, etc.)
# reject unknown fields, so `store` must be opt-in, not opt-out.
if self.config.store is not None:
params["store"] = self.config.store
if response_format:
params["response_format"] = response_format
if tools: # TODO: Remove tools if no issues found with new memory addition logic
+1 -1
View File
@@ -12,7 +12,7 @@ class OpenAIStructuredLLM(LLMBase):
super().__init__(config)
if not self.config.model:
self.config.model = "gpt-4o-2024-08-06"
self.config.model = "gpt-5-mini"
api_key = self.config.api_key or os.getenv("OPENAI_API_KEY")
base_url = self.config.openai_base_url or os.getenv("OPENAI_API_BASE") or "https://api.openai.com/v1"
+379 -24
View File
@@ -110,6 +110,64 @@ def _reject_top_level_entity_params(kwargs: Dict[str, Any], method_name: str) ->
)
def _validate_and_trim_entity_id(value: Optional[str], name: str) -> Optional[str]:
"""
Validates and normalizes an entity ID.
- Trims leading/trailing whitespace
- Rejects empty or whitespace-only strings
- Rejects strings containing internal whitespace
Args:
value: The entity ID value to validate
name: The parameter name (for error messages)
Returns:
The trimmed entity ID, or None if input is None
Raises:
ValueError: If entity ID is invalid
"""
if value is None:
return None
trimmed = value.strip()
if trimmed == "":
raise ValueError(
f"Invalid {name}: cannot be empty or whitespace-only. Provide a valid identifier."
)
if any(c.isspace() for c in trimmed):
raise ValueError(
f"Invalid {name}: cannot contain whitespace. Provide a valid identifier without spaces."
)
return trimmed
def _validate_search_params(threshold: Optional[float] = None, top_k: Optional[int] = None) -> None:
"""
Validates search parameters.
Args:
threshold: Similarity threshold (must be between 0 and 1)
top_k: Number of results to return (must be non-negative integer)
Raises:
ValueError: If threshold or top_k are invalid
"""
if threshold is not None:
if not isinstance(threshold, (int, float)):
raise ValueError("threshold must be a valid number")
if threshold < 0 or threshold > 1:
raise ValueError(
f"Invalid threshold: {threshold}. Must be between 0 and 1 (inclusive)."
)
if top_k is not None:
if not isinstance(top_k, int) or isinstance(top_k, bool):
raise ValueError("top_k must be a valid integer")
if top_k < 0:
raise ValueError(
f"Invalid top_k: {top_k}. Must be a non-negative integer."
)
def _is_sensitive_field(field_name: str) -> bool:
"""Check if a field should be redacted for telemetry safety.
@@ -217,9 +275,14 @@ def _build_filters_and_metadata(
base_metadata_template = deepcopy(input_metadata) if input_metadata else {}
effective_query_filters = deepcopy(input_filters) if input_filters else {}
# ---------- add all provided session ids ----------
# ---------- validate and add all provided session ids ----------
session_ids_provided = []
# Validate and trim entity IDs
user_id = _validate_and_trim_entity_id(user_id, "user_id")
agent_id = _validate_and_trim_entity_id(agent_id, "agent_id")
run_id = _validate_and_trim_entity_id(run_id, "run_id")
if user_id:
base_metadata_template["user_id"] = user_id
effective_query_filters["user_id"] = user_id
@@ -334,6 +397,14 @@ class Memory(MemoryBase):
entity_config.collection_name = entity_collection
elif isinstance(entity_config, dict):
entity_config['collection_name'] = entity_collection
# For Qdrant, share the existing client to avoid RocksDB lock contention
# when using embedded mode (path=...). QdrantConfig.client takes precedence
# over host/port/path.
if self.config.vector_store.provider == "qdrant" and hasattr(self.vector_store, "client"):
if hasattr(entity_config, "client"):
entity_config.client = self.vector_store.client
elif isinstance(entity_config, dict):
entity_config["client"] = self.vector_store.client
self._entity_store = VectorStoreFactory.create(
self.config.vector_store.provider, entity_config
)
@@ -382,6 +453,84 @@ class Memory(MemoryBase):
except Exception as e:
logger.warning(f"Entity upsert failed for '{entity_text}': {e}")
def _remove_memory_from_entity_store(self, memory_id, filters):
"""Strip `memory_id` from every entity record scoped to `filters`.
For each entity whose `linked_memory_ids` contains `memory_id`:
- remove the id; if the list becomes empty, delete the entity record.
- otherwise re-embed the entity text and update the payload
(the vector store's update() requires a vector).
No-op if the entity store has never been initialized in this process.
Errors on individual entities are swallowed at debug level; outer
failures are swallowed at warning level so the primary delete/update
path is never broken by entity cleanup.
"""
if self._entity_store is None:
return
search_filters = {k: v for k, v in filters.items() if k in ("user_id", "agent_id", "run_id") and v}
try:
listed = self.entity_store.list(filters=search_filters, top_k=10000)
rows = listed[0] if isinstance(listed, (list, tuple)) and listed and isinstance(listed[0], list) else listed
for row in rows or []:
try:
payload = getattr(row, "payload", None) or {}
linked = payload.get("linked_memory_ids", [])
if not isinstance(linked, list) or memory_id not in linked:
continue
remaining = [mid for mid in linked if mid != memory_id]
if not remaining:
try:
self.entity_store.delete(vector_id=row.id)
except Exception as e:
logger.debug(f"Entity delete failed for id={row.id}: {e}")
else:
entity_text = payload.get("data")
if not isinstance(entity_text, str) or not entity_text:
logger.debug(f"Entity id={row.id} missing 'data'; skipping update during cleanup")
continue
try:
vec = self.embedding_model.embed(entity_text, "update")
except Exception as e:
logger.debug(f"Entity re-embed failed for '{entity_text}': {e}")
continue
new_payload = {**payload, "linked_memory_ids": remaining}
try:
self.entity_store.update(
vector_id=row.id,
vector=vec,
payload=new_payload,
)
except Exception as e:
logger.debug(f"Entity update failed for id={row.id}: {e}")
except Exception as e:
logger.debug(f"Entity cleanup error: {e}")
except Exception as e:
logger.warning(f"Entity store cleanup failed for memory_id={memory_id}: {e}")
def _link_entities_for_memory(self, memory_id, text, filters):
"""Extract entities from `text` and link them to `memory_id` in the
entity store, scoped to `filters`. Simpler single-memory variant of
Phase 7 in add(): per-entity search-then-update-or-insert via the
existing `_upsert_entity` helper. Non-fatal on any failure.
"""
try:
entities = extract_entities(text)
if not entities:
return
seen = set()
for entity_type, entity_text in entities:
key = entity_text.strip().lower()
if not key or key in seen:
continue
seen.add(key)
try:
self._upsert_entity(entity_text, entity_type, memory_id, filters)
except Exception as e:
logger.debug(f"Entity link failed for '{entity_text}': {e}")
except Exception as e:
logger.warning(f"Entity linking failed for memory_id={memory_id}: {e}")
@classmethod
def from_config(cls, config_dict: Dict[str, Any]):
try:
@@ -868,7 +1017,7 @@ class Memory(MemoryBase):
self,
*,
filters: Optional[Dict[str, Any]] = None,
top_k: int = 100,
top_k: int = 20,
**kwargs,
):
"""
@@ -878,20 +1027,38 @@ class Memory(MemoryBase):
filters (dict): Filter dict containing entity IDs and optional metadata filters.
Must contain at least one of: user_id, agent_id, run_id.
Example: filters={"user_id": "u1", "agent_id": "a1"}
top_k (int, optional): The maximum number of memories to return. Defaults to 100.
top_k (int, optional): The maximum number of memories to return. Defaults to 20.
Returns:
dict: A dictionary containing a list of memories under the "results" key.
Example for v1.1+: `{"results": [{"id": "...", "memory": "...", ...}]}`
Raises:
ValueError: If filters doesn't contain at least one of user_id, agent_id, run_id.
ValueError: If filters doesn't contain at least one of user_id, agent_id, run_id,
or if top_k is invalid.
"""
# Reject top-level entity params - must use filters instead
_reject_top_level_entity_params(kwargs, "get_all")
# Validate top_k
_validate_search_params(top_k=top_k)
# Validate and trim entity IDs in filters
effective_filters = dict(filters) if filters else {}
if "user_id" in effective_filters:
effective_filters["user_id"] = _validate_and_trim_entity_id(
effective_filters["user_id"], "user_id"
)
if "agent_id" in effective_filters:
effective_filters["agent_id"] = _validate_and_trim_entity_id(
effective_filters["agent_id"], "agent_id"
)
if "run_id" in effective_filters:
effective_filters["run_id"] = _validate_and_trim_entity_id(
effective_filters["run_id"], "run_id"
)
# Validate filters contains at least one entity ID
effective_filters = filters or {}
if not any(key in effective_filters for key in ("user_id", "agent_id", "run_id")):
raise ValueError(
"filters must contain at least one of: user_id, agent_id, run_id. "
@@ -960,7 +1127,7 @@ class Memory(MemoryBase):
self,
query: str,
*,
top_k: int = 100,
top_k: int = 20,
filters: Optional[Dict[str, Any]] = None,
threshold: float = 0.1,
rerank: bool = False,
@@ -971,7 +1138,7 @@ class Memory(MemoryBase):
Args:
query (str): Query to search for.
top_k (int, optional): Maximum number of results to return. Defaults to 100.
top_k (int, optional): Maximum number of results to return. Defaults to 20.
filters (dict): Filter dict containing entity IDs and optional metadata filters.
Must contain at least one of: user_id, agent_id, run_id.
Example: filters={"user_id": "u1", "agent_id": "a1"}
@@ -1000,13 +1167,29 @@ class Memory(MemoryBase):
Example for v1.1+: `{"results": [{"id": "...", "memory": "...", "score": 0.8, ...}]}`
Raises:
ValueError: If filters doesn't contain at least one of user_id, agent_id, run_id.
ValueError: If filters doesn't contain at least one of user_id, agent_id, run_id,
or if threshold/top_k values are invalid.
"""
# Reject top-level entity params - must use filters instead
_reject_top_level_entity_params(kwargs, "search")
# Validate filters contains at least one entity ID
# Validate search parameters (before applying defaults)
_validate_search_params(threshold=threshold, top_k=top_k)
# Validate and trim entity IDs in filters
effective_filters = filters.copy() if filters else {}
if "user_id" in effective_filters:
effective_filters["user_id"] = _validate_and_trim_entity_id(
effective_filters["user_id"], "user_id"
)
if "agent_id" in effective_filters:
effective_filters["agent_id"] = _validate_and_trim_entity_id(
effective_filters["agent_id"], "agent_id"
)
if "run_id" in effective_filters:
effective_filters["run_id"] = _validate_and_trim_entity_id(
effective_filters["run_id"], "run_id"
)
if not any(key in effective_filters for key in ("user_id", "agent_id", "run_id")):
raise ValueError(
"filters must contain at least one of: user_id, agent_id, run_id. "
@@ -1188,8 +1371,6 @@ class Memory(MemoryBase):
entity_boosts = self._compute_entity_boosts(query_entities, filters)
# Step 7: Build candidate set from semantic results
# BM25 acts as a boost signal only (not recall-expanding) -- candidates must
# pass the semantic threshold gate, so only semantic results are candidates.
candidates = []
for mem in semantic_results:
mem_id = str(mem.id)
@@ -1234,9 +1415,6 @@ class Memory(MemoryBase):
score=scored["score"],
).model_dump()
# Add score breakdown to metadata
memory_item_dict["score_breakdown"] = scored.get("score_breakdown", {})
for key in promoted_payload_keys:
if key in payload:
memory_item_dict[key] = payload[key]
@@ -1524,6 +1702,13 @@ class Memory(MemoryBase):
actor_id=new_metadata.get("actor_id"),
role=new_metadata.get("role"),
)
# Entity-store cleanup: strip this memory's id from old-text entities,
# then re-extract entities from the new text and link them back.
session_filters = {k: new_metadata[k] for k in ("user_id", "agent_id", "run_id") if new_metadata.get(k)}
self._remove_memory_from_entity_store(memory_id, session_filters)
self._link_entities_for_memory(memory_id, data, session_filters)
return memory_id
def _delete_memory(self, memory_id, existing_memory=None):
@@ -1535,6 +1720,8 @@ class Memory(MemoryBase):
prev_value = existing_memory.payload.get("data", "")
created_at = _normalize_iso_timestamp_to_utc(existing_memory.payload.get("created_at"))
updated_at = datetime.now(timezone.utc).isoformat()
payload = existing_memory.payload or {}
session_filters = {k: payload[k] for k in ("user_id", "agent_id", "run_id") if payload.get(k)}
self.vector_store.delete(vector_id=memory_id)
self.db.add_history(
memory_id,
@@ -1547,6 +1734,11 @@ class Memory(MemoryBase):
role=existing_memory.payload.get("role"),
is_deleted=1,
)
# Entity-store cleanup: strip this memory's id from any entity records
# that linked to it. Non-fatal — the helper swallows errors.
self._remove_memory_from_entity_store(memory_id, session_filters)
return memory_id
def reset(self):
@@ -1640,11 +1832,127 @@ class AsyncMemory(MemoryBase):
entity_config.collection_name = entity_collection
elif isinstance(entity_config, dict):
entity_config['collection_name'] = entity_collection
# For Qdrant, share the existing client to avoid RocksDB lock contention
# when using embedded mode (path=...). QdrantConfig.client takes precedence
# over host/port/path.
if self.config.vector_store.provider == "qdrant" and hasattr(self.vector_store, "client"):
if hasattr(entity_config, "client"):
entity_config.client = self.vector_store.client
elif isinstance(entity_config, dict):
entity_config["client"] = self.vector_store.client
self._entity_store = VectorStoreFactory.create(
self.config.vector_store.provider, entity_config
)
return self._entity_store
async def _upsert_entity_async(self, entity_text, entity_type, memory_id, filters):
"""Async variant of `_upsert_entity` — per-entity search-then-update-or-insert."""
try:
entity_embedding = await asyncio.to_thread(self.embedding_model.embed, entity_text, "add")
search_filters = {k: v for k, v in filters.items() if k in ("user_id", "agent_id", "run_id") and v}
existing = await asyncio.to_thread(
self.entity_store.search,
query=entity_text,
vectors=entity_embedding,
top_k=1,
filters=search_filters,
)
if existing and existing[0].score >= 0.95:
match = existing[0]
payload = match.payload or {}
linked_ids = payload.get("linked_memory_ids", [])
if memory_id not in linked_ids:
linked_ids.append(memory_id)
payload["linked_memory_ids"] = linked_ids
await asyncio.to_thread(
self.entity_store.update,
vector_id=match.id,
vector=None,
payload=payload,
)
else:
entity_id = str(uuid.uuid4())
entity_payload = {
"data": entity_text,
"entity_type": entity_type,
"linked_memory_ids": [memory_id],
**{k: v for k, v in search_filters.items()},
}
await asyncio.to_thread(
self.entity_store.insert,
vectors=[entity_embedding],
ids=[entity_id],
payloads=[entity_payload],
)
except Exception as e:
logger.warning(f"Entity upsert failed for '{entity_text}' (async): {e}")
async def _remove_memory_from_entity_store(self, memory_id, filters):
"""Async variant of `Memory._remove_memory_from_entity_store`."""
if self._entity_store is None:
return
search_filters = {k: v for k, v in filters.items() if k in ("user_id", "agent_id", "run_id") and v}
try:
listed = await asyncio.to_thread(self.entity_store.list, filters=search_filters, top_k=10000)
rows = listed[0] if isinstance(listed, (list, tuple)) and listed and isinstance(listed[0], list) else listed
for row in rows or []:
try:
payload = getattr(row, "payload", None) or {}
linked = payload.get("linked_memory_ids", [])
if not isinstance(linked, list) or memory_id not in linked:
continue
remaining = [mid for mid in linked if mid != memory_id]
if not remaining:
try:
await asyncio.to_thread(self.entity_store.delete, vector_id=row.id)
except Exception as e:
logger.debug(f"Entity delete failed for id={row.id} (async): {e}")
else:
entity_text = payload.get("data")
if not isinstance(entity_text, str) or not entity_text:
logger.debug(f"Entity id={row.id} missing 'data'; skipping update during cleanup (async)")
continue
try:
vec = await asyncio.to_thread(self.embedding_model.embed, entity_text, "update")
except Exception as e:
logger.debug(f"Entity re-embed failed for '{entity_text}' (async): {e}")
continue
new_payload = {**payload, "linked_memory_ids": remaining}
try:
await asyncio.to_thread(
self.entity_store.update,
vector_id=row.id,
vector=vec,
payload=new_payload,
)
except Exception as e:
logger.debug(f"Entity update failed for id={row.id} (async): {e}")
except Exception as e:
logger.debug(f"Entity cleanup error (async): {e}")
except Exception as e:
logger.warning(f"Entity store cleanup failed for memory_id={memory_id} (async): {e}")
async def _link_entities_for_memory(self, memory_id, text, filters):
"""Async variant of `Memory._link_entities_for_memory`."""
try:
entities = await asyncio.to_thread(extract_entities, text)
if not entities:
return
seen = set()
for entity_type, entity_text in entities:
key = entity_text.strip().lower()
if not key or key in seen:
continue
seen.add(key)
try:
await self._upsert_entity_async(entity_text, entity_type, memory_id, filters)
except Exception as e:
logger.debug(f"Entity link failed for '{entity_text}' (async): {e}")
except Exception as e:
logger.warning(f"Entity linking failed for memory_id={memory_id} (async): {e}")
@classmethod
def from_config(cls, config_dict: Dict[str, Any]):
try:
@@ -2115,7 +2423,7 @@ class AsyncMemory(MemoryBase):
self,
*,
filters: Optional[Dict[str, Any]] = None,
top_k: int = 100,
top_k: int = 20,
**kwargs,
):
"""
@@ -2125,20 +2433,38 @@ class AsyncMemory(MemoryBase):
filters (dict): Filter dict containing entity IDs and optional metadata filters.
Must contain at least one of: user_id, agent_id, run_id.
Example: filters={"user_id": "u1", "agent_id": "a1"}
top_k (int, optional): The maximum number of memories to return. Defaults to 100.
top_k (int, optional): The maximum number of memories to return. Defaults to 20.
Returns:
dict: A dictionary containing a list of memories under the "results" key.
Example for v1.1+: `{"results": [{"id": "...", "memory": "...", ...}]}`
Raises:
ValueError: If filters doesn't contain at least one of user_id, agent_id, run_id.
ValueError: If filters doesn't contain at least one of user_id, agent_id, run_id,
or if top_k is invalid.
"""
# Reject top-level entity params - must use filters instead
_reject_top_level_entity_params(kwargs, "get_all")
# Validate top_k
_validate_search_params(top_k=top_k)
# Validate and trim entity IDs in filters
effective_filters = dict(filters) if filters else {}
if "user_id" in effective_filters:
effective_filters["user_id"] = _validate_and_trim_entity_id(
effective_filters["user_id"], "user_id"
)
if "agent_id" in effective_filters:
effective_filters["agent_id"] = _validate_and_trim_entity_id(
effective_filters["agent_id"], "agent_id"
)
if "run_id" in effective_filters:
effective_filters["run_id"] = _validate_and_trim_entity_id(
effective_filters["run_id"], "run_id"
)
# Validate filters contains at least one entity ID
effective_filters = filters or {}
if not any(key in effective_filters for key in ("user_id", "agent_id", "run_id")):
raise ValueError(
"filters must contain at least one of: user_id, agent_id, run_id. "
@@ -2207,7 +2533,7 @@ class AsyncMemory(MemoryBase):
self,
query: str,
*,
top_k: int = 100,
top_k: int = 20,
filters: Optional[Dict[str, Any]] = None,
threshold: float = 0.1,
rerank: bool = False,
@@ -2218,7 +2544,7 @@ class AsyncMemory(MemoryBase):
Args:
query (str): Query to search for.
top_k (int, optional): Maximum number of results to return. Defaults to 100.
top_k (int, optional): Maximum number of results to return. Defaults to 20.
filters (dict): Filter dict containing entity IDs and optional metadata filters.
Must contain at least one of: user_id, agent_id, run_id.
Example: filters={"user_id": "u1", "agent_id": "a1"}
@@ -2247,13 +2573,31 @@ class AsyncMemory(MemoryBase):
Example for v1.1+: `{"results": [{"id": "...", "memory": "...", "score": 0.8, ...}]}`
Raises:
ValueError: If filters doesn't contain at least one of user_id, agent_id, run_id.
ValueError: If filters doesn't contain at least one of user_id, agent_id, run_id,
or if threshold/top_k values are invalid.
"""
# Reject top-level entity params - must use filters instead
_reject_top_level_entity_params(kwargs, "search")
# Validate filters contains at least one entity ID
# Validate search parameters (before applying defaults)
_validate_search_params(threshold=threshold, top_k=top_k)
# Validate and trim entity IDs in filters
effective_filters = filters.copy() if filters else {}
if "user_id" in effective_filters:
effective_filters["user_id"] = _validate_and_trim_entity_id(
effective_filters["user_id"], "user_id"
)
if "agent_id" in effective_filters:
effective_filters["agent_id"] = _validate_and_trim_entity_id(
effective_filters["agent_id"], "agent_id"
)
if "run_id" in effective_filters:
effective_filters["run_id"] = _validate_and_trim_entity_id(
effective_filters["run_id"], "run_id"
)
# Validate filters contains at least one entity ID
if not any(key in effective_filters for key in ("user_id", "agent_id", "run_id")):
raise ValueError(
"filters must contain at least one of: user_id, agent_id, run_id. "
@@ -2480,8 +2824,6 @@ class AsyncMemory(MemoryBase):
score=scored["score"],
).model_dump()
memory_item_dict["score_breakdown"] = scored.get("score_breakdown", {})
for key in promoted_payload_keys:
if key in payload:
memory_item_dict[key] = payload[key]
@@ -2784,6 +3126,13 @@ class AsyncMemory(MemoryBase):
actor_id=new_metadata.get("actor_id"),
role=new_metadata.get("role"),
)
# Entity-store cleanup: strip this memory's id from old-text entities,
# then re-extract entities from the new text and link them back.
session_filters = {k: new_metadata[k] for k in ("user_id", "agent_id", "run_id") if new_metadata.get(k)}
await self._remove_memory_from_entity_store(memory_id, session_filters)
await self._link_entities_for_memory(memory_id, data, session_filters)
return memory_id
async def _delete_memory(self, memory_id, existing_memory=None):
@@ -2795,6 +3144,8 @@ class AsyncMemory(MemoryBase):
prev_value = existing_memory.payload.get("data", "")
created_at = _normalize_iso_timestamp_to_utc(existing_memory.payload.get("created_at"))
updated_at = datetime.now(timezone.utc).isoformat()
payload = existing_memory.payload or {}
session_filters = {k: payload[k] for k in ("user_id", "agent_id", "run_id") if payload.get(k)}
await asyncio.to_thread(self.vector_store.delete, vector_id=memory_id)
await asyncio.to_thread(
@@ -2810,6 +3161,10 @@ class AsyncMemory(MemoryBase):
is_deleted=1,
)
# Entity-store cleanup: strip this memory's id from any entity records
# that linked to it. Non-fatal — the helper swallows errors.
await self._remove_memory_from_entity_store(memory_id, session_filters)
return memory_id
async def reset(self):
+14 -3
View File
@@ -32,15 +32,26 @@ def get_user_id():
return "anonymous_user"
def get_or_create_user_id(vector_store):
"""Store user_id in vector store and return it."""
def get_or_create_user_id(vector_store=None):
"""Store user_id in vector store and return it.
If vector_store is None, simply returns the user_id from config.
This ensures telemetry initialization never fails due to missing vector store.
"""
user_id = get_user_id()
# If no vector store provided, just return the user_id
if vector_store is None:
return user_id
# Try to get existing user_id from vector store
try:
existing = vector_store.get(vector_id=user_id)
if existing and hasattr(existing, "payload") and existing.payload and "user_id" in existing.payload:
return existing.payload["user_id"]
stored_id = existing.payload["user_id"]
# Ensure we never return None from vector store
if stored_id is not None:
return stored_id
except Exception:
pass
+47 -23
View File
@@ -88,6 +88,12 @@ class AnonymousTelemetry:
if self.posthog is None:
return
# Determine distinct_id, skip if None to prevent crashes
distinct_id = self.user_id if user_email is None else user_email
if distinct_id is None:
_logger.debug("Skipping telemetry event %r: no distinct_id available", event_name)
return
if properties is None:
properties = {}
properties = {
@@ -101,8 +107,10 @@ class AnonymousTelemetry:
"machine": platform.machine(),
**properties,
}
distinct_id = self.user_id if user_email is None else user_email
self.posthog.capture(distinct_id=distinct_id, event=event_name, properties=properties)
try:
self.posthog.capture(distinct_id=distinct_id, event=event_name, properties=properties)
except Exception as e:
_logger.debug("Failed to capture telemetry event %r: %s", event_name, e)
def close(self):
if self.posthog is not None:
@@ -158,36 +166,52 @@ atexit.register(client_telemetry.close)
def capture_event(event_name, memory_instance, additional_data=None):
"""Capture telemetry event for OSS Memory instances.
This function is designed to never raise exceptions - telemetry failures
should not affect the main application flow.
"""
if not MEM0_TELEMETRY:
return
oss_telemetry = _get_oss_telemetry()
if oss_telemetry is None:
return
try:
oss_telemetry = _get_oss_telemetry()
if oss_telemetry is None:
return
event_data = {
"collection": memory_instance.collection_name,
"vector_size": memory_instance.embedding_model.config.embedding_dims,
"history_store": "sqlite",
"vector_store": f"{memory_instance.vector_store.__class__.__module__}.{memory_instance.vector_store.__class__.__name__}",
"llm": f"{memory_instance.llm.__class__.__module__}.{memory_instance.llm.__class__.__name__}",
"embedding_model": f"{memory_instance.embedding_model.__class__.__module__}.{memory_instance.embedding_model.__class__.__name__}",
"function": f"{memory_instance.__class__.__module__}.{memory_instance.__class__.__name__}.{memory_instance.api_version}",
}
if additional_data:
event_data.update(additional_data)
event_data = {
"collection": memory_instance.collection_name,
"vector_size": memory_instance.embedding_model.config.embedding_dims,
"history_store": "sqlite",
"vector_store": f"{memory_instance.vector_store.__class__.__module__}.{memory_instance.vector_store.__class__.__name__}",
"llm": f"{memory_instance.llm.__class__.__module__}.{memory_instance.llm.__class__.__name__}",
"embedding_model": f"{memory_instance.embedding_model.__class__.__module__}.{memory_instance.embedding_model.__class__.__name__}",
"function": f"{memory_instance.__class__.__module__}.{memory_instance.__class__.__name__}.{memory_instance.api_version}",
}
if additional_data:
event_data.update(additional_data)
oss_telemetry.capture_event(event_name, event_data)
oss_telemetry.capture_event(event_name, event_data)
except Exception as e:
_logger.debug("Failed to capture OSS telemetry event %r: %s", event_name, e)
def capture_client_event(event_name, instance, additional_data=None):
"""Capture telemetry event for hosted MemoryClient instances.
This function is designed to never raise exceptions - telemetry failures
should not affect the main application flow.
"""
if not MEM0_TELEMETRY:
return
event_data = {
"function": f"{instance.__class__.__module__}.{instance.__class__.__name__}",
}
if additional_data:
event_data.update(additional_data)
try:
event_data = {
"function": f"{instance.__class__.__module__}.{instance.__class__.__name__}",
}
if additional_data:
event_data.update(additional_data)
client_telemetry.capture_event(event_name, event_data, instance.user_email)
client_telemetry.capture_event(event_name, event_data, instance.user_email)
except Exception as e:
_logger.debug("Failed to capture client telemetry event %r: %s", event_name, e)
@@ -6,7 +6,7 @@ from mem0.configs.rerankers.base import BaseRerankerConfig
from mem0.configs.rerankers.sentence_transformer import SentenceTransformerRerankerConfig
try:
from sentence_transformers import SentenceTransformer
from sentence_transformers import CrossEncoder
SENTENCE_TRANSFORMERS_AVAILABLE = True
except ImportError:
SENTENCE_TRANSFORMERS_AVAILABLE = False
@@ -41,7 +41,7 @@ class SentenceTransformerReranker(BaseReranker):
)
self.config = config
self.model = SentenceTransformer(self.config.model, device=self.config.device)
self.model = CrossEncoder(self.config.model, device=self.config.device)
def rerank(self, query: str, documents: List[Dict[str, Any]], top_k: int = None) -> List[Dict[str, Any]]:
"""
@@ -74,8 +74,11 @@ class SentenceTransformerReranker(BaseReranker):
# Create query-document pairs
pairs = [[query, doc_text] for doc_text in doc_texts]
# Get similarity scores
scores = self.model.predict(pairs)
scores = self.model.predict(
pairs,
batch_size=self.config.batch_size,
show_progress_bar=self.config.show_progress_bar,
)
if isinstance(scores, np.ndarray):
scores = scores.tolist()
-5
View File
@@ -113,11 +113,6 @@ def score_and_rank(
{
"id": mem_id_str,
"score": combined,
"score_breakdown": {
"semantic": semantic_score,
"bm25": bm25_score,
"entity_boost": entity_boost,
},
"payload": result.get("payload"),
}
)
+155 -17
View File
@@ -1,10 +1,11 @@
import json
import logging
import os
import pickle
import uuid
import warnings
from pathlib import Path
from typing import Dict, List, Optional
from typing import Any, Dict, List, Optional
import numpy as np
from pydantic import BaseModel
@@ -13,7 +14,7 @@ try:
# Suppress SWIG deprecation warnings from FAISS
warnings.filterwarnings("ignore", category=DeprecationWarning, message=".*SwigPy.*")
warnings.filterwarnings("ignore", category=DeprecationWarning, message=".*swigvarlink.*")
logging.getLogger("faiss").setLevel(logging.WARNING)
logging.getLogger("faiss.loader").setLevel(logging.WARNING)
@@ -30,6 +31,93 @@ from mem0.vector_stores.base import VectorStoreBase
logger = logging.getLogger(__name__)
class SafeUnpickler(pickle.Unpickler):
"""
Restricted unpickler that only allows safe built-in types.
This prevents arbitrary code execution via pickle deserialization by only
allowing a whitelist of safe types (dict, list, str, int, float, bool, tuple, None).
"""
# Only allow builtins module
SAFE_MODULES = frozenset({"builtins", "__builtin__"})
# Only allow safe basic types
SAFE_NAMES = frozenset({"dict", "list", "str", "int", "float", "bool", "tuple", "set", "frozenset", "NoneType"})
def find_class(self, module: str, name: str) -> Any:
"""Override find_class to only allow safe types."""
if module in self.SAFE_MODULES and name in self.SAFE_NAMES:
import builtins
if hasattr(builtins, name):
return getattr(builtins, name)
# NoneType special case
if name == "NoneType":
return type(None)
raise pickle.UnpicklingError(
f"Unsafe pickle: attempted to load '{module}.{name}'. "
f"Only basic Python types are allowed for security reasons."
)
def _safe_pickle_load(file_path: str) -> Any:
"""
Safely load a pickle file using restricted unpickler.
Args:
file_path: Path to the pickle file.
Returns:
The deserialized object (only basic Python types allowed).
Raises:
pickle.UnpicklingError: If the pickle contains unsafe types.
"""
with open(file_path, "rb") as f:
return SafeUnpickler(f).load()
def _validate_docstore_structure(data: Any) -> tuple:
"""
Validate that loaded data has the expected structure.
Args:
data: The loaded data to validate.
Returns:
Tuple of (docstore, index_to_id) if valid.
Raises:
ValueError: If the data structure is invalid.
"""
if not isinstance(data, tuple) or len(data) != 2:
raise ValueError("Invalid docstore format: expected tuple of (docstore, index_to_id)")
docstore, index_to_id = data
if not isinstance(docstore, dict):
raise ValueError("Invalid docstore format: docstore must be a dict")
if not isinstance(index_to_id, dict):
raise ValueError("Invalid docstore format: index_to_id must be a dict")
# Validate docstore entries
for key, value in docstore.items():
if not isinstance(key, str):
raise ValueError(f"Invalid docstore key type: {type(key)}, expected str")
if not isinstance(value, dict):
raise ValueError(f"Invalid docstore value type: {type(value)}, expected dict")
# Validate index_to_id entries
for key, value in index_to_id.items():
if not isinstance(key, int):
raise ValueError(f"Invalid index_to_id key type: {type(key)}, expected int")
if not isinstance(value, str):
raise ValueError(f"Invalid index_to_id value type: {type(value)}, expected str")
return docstore, index_to_id
class OutputData(BaseModel):
id: Optional[str] # memory id
score: Optional[float] # distance
@@ -73,9 +161,13 @@ class FAISS(VectorStoreBase):
# Try to load existing index if available
index_path = f"{self.path}/{collection_name}.faiss"
docstore_path = f"{self.path}/{collection_name}.pkl"
if os.path.exists(index_path) and os.path.exists(docstore_path):
self._load(index_path, docstore_path)
json_docstore_path = f"{self.path}/{collection_name}.json"
pkl_docstore_path = f"{self.path}/{collection_name}.pkl"
# Check for index file and either JSON (preferred) or legacy pickle docstore
if os.path.exists(index_path) and (os.path.exists(json_docstore_path) or os.path.exists(pkl_docstore_path)):
# _load will prefer JSON over pickle and auto-migrate
self._load(index_path, pkl_docstore_path)
else:
self.create_col(collection_name)
@@ -83,34 +175,76 @@ class FAISS(VectorStoreBase):
"""
Load FAISS index and docstore from disk.
Supports both JSON (preferred) and legacy pickle formats. Pickle files are loaded
using a restricted unpickler that only allows basic Python types to prevent
arbitrary code execution (CVE mitigation).
Args:
index_path (str): Path to FAISS index file.
docstore_path (str): Path to docstore pickle file.
docstore_path (str): Path to docstore file (.json or legacy .pkl).
"""
try:
self.index = faiss.read_index(index_path)
with open(docstore_path, "rb") as f:
self.docstore, self.index_to_id = pickle.load(f)
logger.info(f"Loaded FAISS index from {index_path} with {self.index.ntotal} vectors")
# Determine docstore format - prefer JSON over pickle
json_docstore_path = docstore_path.replace(".pkl", ".json")
if os.path.exists(json_docstore_path):
# Load from JSON (safe, preferred format)
with open(json_docstore_path, "r", encoding="utf-8") as f:
data = json.load(f)
self.docstore = data.get("docstore", {})
# JSON keys are always strings, convert back to int
self.index_to_id = {int(k): v for k, v in data.get("index_to_id", {}).items()}
logger.info(f"Loaded FAISS index from {index_path} with {self.index.ntotal} vectors (JSON format)")
elif os.path.exists(docstore_path):
# Load from legacy pickle using safe unpickler
# This prevents arbitrary code execution from malicious pickle files
logger.warning(
f"Loading legacy pickle docstore from {docstore_path}. "
f"Consider migrating to JSON format for better security."
)
data = _safe_pickle_load(docstore_path)
self.docstore, self.index_to_id = _validate_docstore_structure(data)
logger.info(f"Loaded FAISS index from {index_path} with {self.index.ntotal} vectors (pickle format)")
# Auto-migrate to JSON format
self._save()
logger.info(f"Migrated docstore to JSON format: {json_docstore_path}")
else:
raise FileNotFoundError(f"No docstore found at {docstore_path} or {json_docstore_path}")
except pickle.UnpicklingError as e:
logger.error(f"Security error loading FAISS docstore: {e}")
raise ValueError(f"Failed to load FAISS docstore: potentially malicious pickle file. {e}") from e
except Exception as e:
logger.warning(f"Failed to load FAISS index: {e}")
self.docstore = {}
self.index_to_id = {}
def _save(self):
"""Save FAISS index and docstore to disk."""
"""Save FAISS index and docstore to disk using JSON format (secure)."""
if not self.path or not self.index:
return
try:
os.makedirs(self.path, exist_ok=True)
index_path = f"{self.path}/{self.collection_name}.faiss"
docstore_path = f"{self.path}/{self.collection_name}.pkl"
json_docstore_path = f"{self.path}/{self.collection_name}.json"
faiss.write_index(self.index, index_path)
with open(docstore_path, "wb") as f:
pickle.dump((self.docstore, self.index_to_id), f)
# Save docstore as JSON (safe format, no code execution risk)
# JSON keys must be strings, so convert int keys to str
data = {
"docstore": self.docstore,
"index_to_id": {str(k): v for k, v in self.index_to_id.items()},
}
with open(json_docstore_path, "w", encoding="utf-8") as f:
json.dump(data, f, indent=2)
except Exception as e:
logger.warning(f"Failed to save FAISS index: {e}")
@@ -417,12 +551,16 @@ class FAISS(VectorStoreBase):
if self.path:
try:
index_path = f"{self.path}/{self.collection_name}.faiss"
docstore_path = f"{self.path}/{self.collection_name}.pkl"
json_docstore_path = f"{self.path}/{self.collection_name}.json"
pkl_docstore_path = f"{self.path}/{self.collection_name}.pkl"
if os.path.exists(index_path):
os.remove(index_path)
if os.path.exists(docstore_path):
os.remove(docstore_path)
if os.path.exists(json_docstore_path):
os.remove(json_docstore_path)
# Also clean up legacy pickle files if they exist
if os.path.exists(pkl_docstore_path):
os.remove(pkl_docstore_path)
logger.info(f"Deleted collection {self.collection_name}")
except Exception as e:
+40 -17
View File
@@ -46,6 +46,9 @@ class MilvusDB(VectorStoreBase):
self.embedding_model_dims = embedding_model_dims
self.metric_type = metric_type
self.client = MilvusClient(uri=url, token=token, db_name=db_name)
# Whether this collection has the `text` + `sparse` fields for v3 BM25.
# Pre-v3 collections lack them; writing a top-level `text` field is rejected.
self._has_bm25_schema = False
self.create_col(
collection_name=self.collection_name,
vector_size=self.embedding_model_dims,
@@ -68,6 +71,15 @@ class MilvusDB(VectorStoreBase):
if self.client.has_collection(collection_name):
logger.info(f"Collection {collection_name} already exists. Skipping creation.")
desc = self.client.describe_collection(collection_name=collection_name)
field_names = {f.get("name") for f in desc.get("fields", [])}
self._has_bm25_schema = "text" in field_names and "sparse" in field_names
if not self._has_bm25_schema:
logger.warning(
f"Collection '{collection_name}' predates v3 hybrid search (no 'text'/'sparse' fields). "
"BM25 keyword scoring will be disabled for this collection; semantic search works normally. "
"To enable hybrid search, use a fresh collection."
)
else:
fields = [
FieldSchema(name="id", dtype=DataType.VARCHAR, is_primary=True, max_length=512),
@@ -101,6 +113,7 @@ class MilvusDB(VectorStoreBase):
index_name="sparse_index",
)
self.client.create_collection(collection_name=collection_name, schema=schema, index_params=index_params)
self._has_bm25_schema = True
def insert(self, ids, vectors, payloads, **kwargs: Optional[dict[str, any]]):
"""Insert vectors into a collection.
@@ -110,17 +123,17 @@ class MilvusDB(VectorStoreBase):
payloads (List[Dict], optional): List of payloads corresponding to vectors.
ids (List[str], optional): List of IDs corresponding to vectors.
"""
# Batch insert all records at once for better performance and consistency
data = [
{
"id": idx,
"vectors": embedding,
"metadata": metadata,
# Batch insert all records at once for better performance and consistency.
# Only include the `text` field when the collection's schema has it — legacy
# collections created pre-v3 reject unknown top-level fields.
def _build_record(idx, embedding, metadata):
record = {"id": idx, "vectors": embedding, "metadata": metadata}
if self._has_bm25_schema:
# Populate the text field for BM25 sparse search; prefer lemmatized text, fall back to raw data
"text": (metadata.get("text_lemmatized") or metadata.get("data", ""))[:65535] if metadata else "",
}
for idx, embedding, metadata in zip(ids, vectors, payloads)
]
record["text"] = (metadata.get("text_lemmatized") or metadata.get("data", ""))[:65535] if metadata else ""
return record
data = [_build_record(idx, embedding, metadata) for idx, embedding, metadata in zip(ids, vectors, payloads)]
self.client.insert(collection_name=self.collection_name, data=data, **kwargs)
def _create_filter(self, filters: dict):
@@ -179,13 +192,21 @@ class MilvusDB(VectorStoreBase):
list: Search results.
"""
query_filter = self._create_filter(filters) if filters else None
hits = self.client.search(
collection_name=self.collection_name,
data=[vectors],
limit=top_k,
filter=query_filter,
output_fields=["*"],
)
# v3 collections carry both a dense `vectors` field and a sparse `sparse`
# field (for BM25), which makes anns_field ambiguous — Milvus rejects the
# query otherwise with "multiple anns_fields exist". Legacy single-vector
# collections don't need the hint, so only pass it when the hybrid schema
# is present.
search_kwargs = {
"collection_name": self.collection_name,
"data": [vectors],
"limit": top_k,
"filter": query_filter,
"output_fields": ["*"],
}
if self._has_bm25_schema:
search_kwargs["anns_field"] = "vectors"
hits = self.client.search(**search_kwargs)
result = self._parse_output(data=hits[0])
return result
@@ -206,6 +227,8 @@ class MilvusDB(VectorStoreBase):
list: Search results in the same format as search(), or None if sparse search
is not supported on this collection.
"""
if not self._has_bm25_schema:
return None
try:
query_filter = self._create_filter(filters) if filters else None
hits = self.client.search(
+31 -13
View File
@@ -79,6 +79,10 @@ class Qdrant(VectorStoreBase):
self.embedding_model_dims = embedding_model_dims
self.on_disk = on_disk
self._bm25_encoder = None
# Whether this collection has the `bm25` named sparse vector slot.
# Pre-v3 collections lack it; writing a `bm25` sparse vector into such a
# collection is rejected by Qdrant ("Not existing vector name error: bm25").
self._has_bm25_slot = False
self.create_col(embedding_model_dims, on_disk)
def _get_bm25_encoder(self):
@@ -127,6 +131,15 @@ class Qdrant(VectorStoreBase):
for collection in response.collections:
if collection.name == self.collection_name:
logger.debug(f"Collection {self.collection_name} already exists. Skipping creation.")
info = self.client.get_collection(self.collection_name)
sparse_cfg = info.config.params.sparse_vectors
self._has_bm25_slot = bool(sparse_cfg and "bm25" in sparse_cfg)
if not self._has_bm25_slot:
logger.warning(
f"Collection '{self.collection_name}' predates v3 hybrid search (no 'bm25' sparse slot). "
"BM25 keyword scoring will be disabled for this collection; semantic search works normally. "
"To enable hybrid search, use a fresh collection."
)
self._create_filter_indexes()
return
@@ -139,6 +152,7 @@ class Qdrant(VectorStoreBase):
),
},
)
self._has_bm25_slot = True
self._create_filter_indexes()
def _create_filter_indexes(self):
@@ -177,13 +191,14 @@ class Qdrant(VectorStoreBase):
payload = payloads[idx] if payloads else {}
point_id = idx if ids is None else ids[idx]
# Build named vectors: dense + optional BM25 sparse
# Build named vectors: dense + optional BM25 sparse (only if collection has the slot).
named_vectors = {"": vector}
text_for_bm25 = payload.get("text_lemmatized") or payload.get("data", "")
if text_for_bm25:
sparse = self._encode_bm25(text_for_bm25)
if sparse is not None:
named_vectors["bm25"] = sparse
if self._has_bm25_slot:
text_for_bm25 = payload.get("text_lemmatized") or payload.get("data", "")
if text_for_bm25:
sparse = self._encode_bm25(text_for_bm25)
if sparse is not None:
named_vectors["bm25"] = sparse
points.append(PointStruct(id=point_id, vector=named_vectors, payload=payload))
@@ -382,7 +397,7 @@ class Qdrant(VectorStoreBase):
"""Batch search using Qdrant's query_batch_points for efficiency."""
query_filter = self._create_filter(filters) if filters else None
requests = [
models.QueryRequest(query=vec, filter=query_filter, limit=top_k)
models.QueryRequest(query=vec, filter=query_filter, limit=top_k, with_payload=True)
for vec in vectors_list
]
try:
@@ -407,6 +422,8 @@ class Qdrant(VectorStoreBase):
Returns:
list: Search results, or None if BM25 is not available.
"""
if not self._has_bm25_slot:
return None
sparse_query = self._encode_bm25(query)
if sparse_query is None:
return None
@@ -449,13 +466,14 @@ class Qdrant(VectorStoreBase):
payload (dict, optional): Updated payload. Defaults to None.
"""
if vector is not None and payload is not None:
# Full update: attach BM25 sparse vector alongside dense vector
# Full update: attach BM25 sparse vector alongside dense vector (only if slot exists).
named_vectors = {"": vector}
text_for_bm25 = payload.get("text_lemmatized") or payload.get("data", "")
if text_for_bm25:
sparse = self._encode_bm25(text_for_bm25)
if sparse is not None:
named_vectors["bm25"] = sparse
if self._has_bm25_slot:
text_for_bm25 = payload.get("text_lemmatized") or payload.get("data", "")
if text_for_bm25:
sparse = self._encode_bm25(text_for_bm25)
if sparse is not None:
named_vectors["bm25"] = sparse
point = PointStruct(id=vector_id, vector=named_vectors, payload=payload)
self.client.upsert(collection_name=self.collection_name, points=[point])
else:
+31 -1
View File
@@ -122,7 +122,37 @@ class S3Vectors(VectorStoreBase):
)
def update(self, vector_id, vector=None, payload=None):
# S3 Vectors uses put_vectors for updates (overwrite)
# S3 Vectors uses put_vectors for updates (overwrite).
# When vector=None (e.g. metadata-only update triggered by event=NONE),
# fetch the existing vector data first to avoid passing None to boto3
# which causes a parameter validation error:
# "Invalid type for parameter vectors[0].data.float32, value: None"
if vector is None:
existing = self.get(vector_id)
if existing is None:
logger.warning(f"update called with vector=None but {vector_id} not found; skipping")
return
try:
response = self.client.get_vectors(
vectorBucketName=self.vector_bucket_name,
indexName=self.collection_name,
keys=[vector_id],
returnData=True,
returnMetadata=True,
)
vectors = response.get("vectors", [])
if not vectors:
logger.warning(f"update: no vector data found for {vector_id}; skipping")
return
vector = vectors[0].get("data", {}).get("float32")
if vector is None:
logger.warning(f"update: float32 data is None for {vector_id}; skipping")
return
if payload is None:
payload = existing.payload
except Exception as e:
logger.error(f"update: failed to fetch existing vector for {vector_id}: {e}")
return
self.insert(vectors=[vector], payloads=[payload], ids=[vector_id])
def get(self, vector_id) -> Optional[OutputData]:
+13 -3
View File
@@ -52,6 +52,7 @@ class ValkeyDB(VectorStoreBase):
hnsw_m: int = 16,
hnsw_ef_construction: int = 200,
hnsw_ef_runtime: int = 10,
cluster_mode: bool = False,
):
"""
Initialize the Valkey vector store.
@@ -65,6 +66,7 @@ class ValkeyDB(VectorStoreBase):
hnsw_m (int, optional): HNSW M parameter (connections per node). Defaults to 16.
hnsw_ef_construction (int, optional): HNSW ef_construction parameter. Defaults to 200.
hnsw_ef_runtime (int, optional): HNSW ef_runtime parameter. Defaults to 10.
cluster_mode (bool, optional): Enable cluster mode for Valkey cluster (CME) deployments. Defaults to False.
"""
self.embedding_model_dims = embedding_model_dims
self.collection_name = collection_name
@@ -74,6 +76,7 @@ class ValkeyDB(VectorStoreBase):
self.hnsw_m = hnsw_m
self.hnsw_ef_construction = hnsw_ef_construction
self.hnsw_ef_runtime = hnsw_ef_runtime
self.cluster_mode = cluster_mode
# Validate index type
if self.index_type not in ["hnsw", "flat"]:
@@ -81,8 +84,13 @@ class ValkeyDB(VectorStoreBase):
# Connect to Valkey
try:
self.client = valkey.from_url(valkey_url)
logger.debug(f"Successfully connected to Valkey at {valkey_url}")
if self.cluster_mode:
from valkey.cluster import ValkeyCluster
self.client = ValkeyCluster.from_url(valkey_url)
else:
self.client = valkey.from_url(valkey_url)
logger.debug(f"Successfully connected to Valkey at {valkey_url} (cluster_mode={cluster_mode})")
except Exception as e:
logger.exception(f"Failed to connect to Valkey at {valkey_url}: {e}")
raise
@@ -185,7 +193,6 @@ class ValkeyDB(VectorStoreBase):
"""
# Check if the search module is available
try:
# Try to execute a search command
self.client.execute_command("FT._LIST")
except ResponseError as e:
if "unknown command" in str(e).lower():
@@ -353,6 +360,9 @@ class ValkeyDB(VectorStoreBase):
"""
Execute a search query.
In cluster mode, the valkey-search module's built-in coordinator handles
fan-out across all shards and aggregates results server-side.
Args:
query (str): The search query to execute.
params (dict): The query parameters.
+5
View File
@@ -147,6 +147,7 @@ export const DEFAULT_CUSTOM_CATEGORIES: Record<string, string> = {
const ALLOWED_KEYS = [
"mode",
"apiKey",
"anonymousTelemetryId",
"baseUrl",
"userId",
"userEmail",
@@ -215,6 +216,10 @@ export const mem0ConfigSchema = {
return {
mode,
apiKey: resolvedApiKey,
anonymousTelemetryId:
typeof cfg.anonymousTelemetryId === "string"
? cfg.anonymousTelemetryId
: undefined,
baseUrl: resolvedBaseUrl,
userId:
typeof cfg.userId === "string" && cfg.userId
+4
View File
@@ -134,6 +134,10 @@
"topK": {
"type": "number"
},
"anonymousTelemetryId": {
"type": "string",
"description": "Persistent anonymous telemetry identifier"
},
"oss": {
"type": "object",
"properties": {
+5
View File
@@ -88,6 +88,11 @@ describe("mem0ConfigSchema.parse() — defaults", () => {
const cfg = mem0ConfigSchema.parse({ apiKey: "test-key" });
expect(cfg.skills).toBeUndefined();
});
it("allows anonymousTelemetryId", () => {
const cfg = mem0ConfigSchema.parse({ apiKey: "test-key", anonymousTelemetryId: "123" });
expect(cfg.anonymousTelemetryId).toBe("123");
});
});
// ---------------------------------------------------------------------------
+1
View File
@@ -8,6 +8,7 @@ export type Mem0Config = {
mode: Mem0Mode;
// Platform-specific
apiKey?: string;
anonymousTelemetryId?: string;
baseUrl?: string;
customInstructions: string;
customCategories: Record<string, string>;
+1 -1
View File
@@ -4,7 +4,7 @@ build-backend = "hatchling.build"
[project]
name = "mem0ai"
version = "1.0.11"
version = "2.0.0b2"
description = "Long-term memory for AI Agents"
authors = [
{ name = "Mem0", email = "support@mem0.ai" }
+2 -2
View File
@@ -286,8 +286,8 @@ def test_init_with_env_vars(monkeypatch):
http_client=None,
default_headers=None,
)
# Should default to "gpt-4.1-nano-2025-04-14" if model is None
assert llm.config.model == "gpt-4.1-nano-2025-04-14"
# Should default to "gpt-5-mini" if model is None
assert llm.config.model == "gpt-5-mini"
def test_init_with_default_azure_credential(monkeypatch):
+2 -2
View File
@@ -57,7 +57,7 @@ def test_init_with_default_credential(mock_credential, mock_token_provider, mock
mock_token_provider.return_value = "token-provider"
llm = AzureOpenAIStructuredLLM(config)
# Should set default model if not provided
assert llm.config.model == "gpt-4.1-nano-2025-04-14"
assert llm.config.model == "gpt-5-mini"
mock_credential.assert_called_once()
mock_token_provider.assert_called_once_with(mock_credential.return_value, SCOPE)
mock_azure_openai.assert_called_once()
@@ -92,7 +92,7 @@ def test_init_with_placeholder_api_key_uses_default_credential(
config = DummyConfig(model=None, azure_kwargs=DummyAzureKwargs(api_key="your-api-key"))
mock_token_provider.return_value = "token-provider"
llm = AzureOpenAIStructuredLLM(config)
assert llm.config.model == "gpt-4.1-nano-2025-04-14"
assert llm.config.model == "gpt-5-mini"
mock_credential.assert_called_once()
mock_token_provider.assert_called_once_with(mock_credential.return_value, SCOPE)
mock_azure_openai.assert_called_once()
+56 -2
View File
@@ -55,7 +55,7 @@ def test_generate_response_without_tools(mock_openai_client):
response = llm.generate_response(messages)
mock_openai_client.chat.completions.create.assert_called_once_with(
model="gpt-4.1-nano-2025-04-14", messages=messages, temperature=0.7, max_tokens=100, top_p=1.0, store=False
model="gpt-4.1-nano-2025-04-14", messages=messages, temperature=0.7, max_tokens=100, top_p=1.0
)
assert response == "I'm doing well, thank you for asking!"
@@ -97,7 +97,7 @@ def test_generate_response_with_tools(mock_openai_client):
response = llm.generate_response(messages, tools=tools)
mock_openai_client.chat.completions.create.assert_called_once_with(
model="gpt-4.1-nano-2025-04-14", messages=messages, temperature=0.7, max_tokens=100, top_p=1.0, tools=tools, tool_choice="auto", store=False
model="gpt-4.1-nano-2025-04-14", messages=messages, temperature=0.7, max_tokens=100, top_p=1.0, tools=tools, tool_choice="auto"
)
assert response["content"] == "I've added the memory for you."
@@ -231,6 +231,60 @@ def test_reasoning_effort_config_values():
assert config.reasoning_effort is None
def test_store_not_sent_by_default(mock_openai_client):
"""`store` must NOT be injected into requests when the user has not
explicitly configured it. Regression test for issue #4709, where
`store=False` was unconditionally sent and rejected by OpenAI-compatible
backends such as Google Gemini."""
config = OpenAIConfig(model="gpt-4.1-nano-2025-04-14", temperature=0.1)
assert config.store is None # new opt-in default
llm = OpenAILLM(config)
messages = [{"role": "user", "content": "Hello"}]
mock_response = Mock()
mock_response.choices = [Mock(message=Mock(content="Response"))]
mock_openai_client.chat.completions.create.return_value = mock_response
llm.generate_response(messages)
call_kwargs = mock_openai_client.chat.completions.create.call_args.kwargs
assert "store" not in call_kwargs
def test_store_sent_when_explicitly_true(mock_openai_client):
"""When the user explicitly sets `store=True`, the field must be forwarded."""
config = OpenAIConfig(model="gpt-4.1-nano-2025-04-14", store=True)
llm = OpenAILLM(config)
messages = [{"role": "user", "content": "Hello"}]
mock_response = Mock()
mock_response.choices = [Mock(message=Mock(content="Response"))]
mock_openai_client.chat.completions.create.return_value = mock_response
llm.generate_response(messages)
call_kwargs = mock_openai_client.chat.completions.create.call_args.kwargs
assert call_kwargs["store"] is True
def test_store_sent_when_explicitly_false(mock_openai_client):
"""When the user explicitly sets `store=False`, the field must still be
forwarded — explicit opt-out is a valid configuration for users who rely on
it for OpenAI's zero-data-retention behavior."""
config = OpenAIConfig(model="gpt-4.1-nano-2025-04-14", store=False)
llm = OpenAILLM(config)
messages = [{"role": "user", "content": "Hello"}]
mock_response = Mock()
mock_response.choices = [Mock(message=Mock(content="Response"))]
mock_openai_client.chat.completions.create.return_value = mock_response
llm.generate_response(messages)
call_kwargs = mock_openai_client.chat.completions.create.call_args.kwargs
assert call_kwargs["store"] is False
def test_callback_with_tools(mock_openai_client):
mock_callback = Mock()
config = OpenAIConfig(model="gpt-4.1-nano-2025-04-14", response_callback=mock_callback)
+59
View File
@@ -0,0 +1,59 @@
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from mem0.client.main import AsyncMemoryClient, MemoryClient
def _build_memory_client(response_payload=None):
client = MemoryClient.__new__(MemoryClient)
client.user_email = "user@example.com"
client.client = MagicMock()
response = client.client.post.return_value
response.json.return_value = response_payload or {"message": "Feedback recorded"}
response.raise_for_status.return_value = None
return client
def test_memory_client_feedback_uses_single_telemetry_payload():
client = _build_memory_client()
with patch("mem0.client.main.capture_client_event") as mock_capture:
result = client.feedback("mem_1", feedback="positive", feedback_reason="accurate")
assert result == {"message": "Feedback recorded"}
client.client.post.assert_called_once_with(
"/v1/feedback/",
json={"memory_id": "mem_1", "feedback": "POSITIVE", "feedback_reason": "accurate"},
)
mock_capture.assert_called_once_with(
"client.feedback",
client,
{"memory_id": "mem_1", "feedback": "POSITIVE", "feedback_reason": "accurate", "sync_type": "sync"},
)
@pytest.mark.asyncio
async def test_async_memory_client_feedback_uses_single_telemetry_payload():
client = AsyncMemoryClient.__new__(AsyncMemoryClient)
client.user_email = "user@example.com"
client.async_client = MagicMock()
response = MagicMock()
response.json.return_value = {"message": "Feedback recorded"}
response.raise_for_status.return_value = None
client.async_client.post = AsyncMock(return_value=response)
with patch("mem0.client.main.capture_client_event") as mock_capture:
result = await client.feedback("mem_1", feedback="negative", feedback_reason="outdated")
assert result == {"message": "Feedback recorded"}
client.async_client.post.assert_awaited_once_with(
"/v1/feedback/",
json={"memory_id": "mem_1", "feedback": "NEGATIVE", "feedback_reason": "outdated"},
)
mock_capture.assert_called_once_with(
"client.feedback",
client,
{"memory_id": "mem_1", "feedback": "NEGATIVE", "feedback_reason": "outdated", "sync_type": "async"},
)
+110 -4
View File
@@ -110,11 +110,10 @@ def test_search(memory_instance):
assert result["results"][0]["user_id"] == "test_user"
# Score is now combined score (semantic only since no BM25/entity), still 0.9
assert result["results"][0]["score"] == pytest.approx(0.9)
assert "score_breakdown" in result["results"][0]
# Hybrid pipeline over-fetches: max(100*4, 60) = 400
# Hybrid pipeline over-fetches: max(20*4, 60) = 80 (top_k default is now 20)
memory_instance.vector_store.search.assert_called_once_with(
query="test query", vectors=[0.1, 0.2, 0.3], top_k=400, filters={"user_id": "test_user"}
query="test query", vectors=[0.1, 0.2, 0.3], top_k=80, filters={"user_id": "test_user"}
)
@@ -201,7 +200,7 @@ def test_get_all(memory_instance):
assert result["results"][0]["memory"] == "Memory 1"
assert result["results"][0]["user_id"] == "test_user"
memory_instance.vector_store.list.assert_called_once_with(filters={"user_id": "test_user"}, top_k=100)
memory_instance.vector_store.list.assert_called_once_with(filters={"user_id": "test_user"}, top_k=20)
def test_no_telemetry_vector_store_when_disabled():
@@ -242,3 +241,110 @@ def test_telemetry_vector_store_created_when_enabled():
# VectorStoreFactory.create should be called twice — user data + telemetry
assert mock_vector_store.create.call_count == 2
# =============================================================================
# Input Validation Tests
# =============================================================================
class TestEntityIdValidation:
"""Tests for entity ID validation (whitespace rejection and trimming)."""
def test_search_rejects_whitespace_only_user_id(self, memory_instance):
"""Search should reject whitespace-only user_id in filters."""
with pytest.raises(ValueError, match="Invalid user_id.*cannot be empty"):
memory_instance.search("test query", filters={"user_id": " "})
def test_search_rejects_internal_whitespace_user_id(self, memory_instance):
"""Search should reject user_id with internal whitespace."""
with pytest.raises(ValueError, match="Invalid user_id.*cannot contain whitespace"):
memory_instance.search("test query", filters={"user_id": "user 123"})
def test_search_rejects_tab_in_user_id(self, memory_instance):
"""Search should reject user_id with tab character."""
with pytest.raises(ValueError, match="Invalid user_id.*cannot contain whitespace"):
memory_instance.search("test query", filters={"user_id": "user\t123"})
def test_get_all_rejects_whitespace_only_user_id(self, memory_instance):
"""get_all should reject whitespace-only user_id in filters."""
with pytest.raises(ValueError, match="Invalid user_id.*cannot be empty"):
memory_instance.get_all(filters={"user_id": " "})
def test_get_all_rejects_internal_whitespace_user_id(self, memory_instance):
"""get_all should reject user_id with internal whitespace."""
with pytest.raises(ValueError, match="Invalid user_id.*cannot contain whitespace"):
memory_instance.get_all(filters={"user_id": "user 123"})
def test_add_rejects_whitespace_only_user_id(self, memory_instance):
"""add should reject whitespace-only user_id."""
with pytest.raises(ValueError, match="Invalid user_id.*cannot be empty"):
memory_instance.add("test message", user_id=" ")
def test_add_rejects_internal_whitespace_user_id(self, memory_instance):
"""add should reject user_id with internal whitespace."""
with pytest.raises(ValueError, match="Invalid user_id.*cannot contain whitespace"):
memory_instance.add("test message", user_id="user 123")
class TestSearchParamValidation:
"""Tests for search parameter validation (threshold and top_k)."""
def test_search_rejects_threshold_above_1(self, memory_instance):
"""Search should reject threshold > 1."""
with pytest.raises(ValueError, match="Invalid threshold.*Must be between 0 and 1"):
memory_instance.search("test query", filters={"user_id": "test"}, threshold=1.5)
def test_search_rejects_negative_threshold(self, memory_instance):
"""Search should reject negative threshold."""
with pytest.raises(ValueError, match="Invalid threshold.*Must be between 0 and 1"):
memory_instance.search("test query", filters={"user_id": "test"}, threshold=-0.5)
def test_search_rejects_negative_top_k(self, memory_instance):
"""Search should reject negative top_k."""
with pytest.raises(ValueError, match="Invalid top_k.*Must be a non-negative"):
memory_instance.search("test query", filters={"user_id": "test"}, top_k=-5)
def test_get_all_rejects_negative_top_k(self, memory_instance):
"""get_all should reject negative top_k."""
with pytest.raises(ValueError, match="Invalid top_k.*Must be a non-negative"):
memory_instance.get_all(filters={"user_id": "test"}, top_k=-1)
def test_search_accepts_threshold_zero(self, memory_instance):
"""Search should accept threshold=0 (edge case)."""
mock_memories = []
memory_instance.vector_store.search = Mock(return_value=mock_memories)
memory_instance.vector_store.keyword_search = Mock(return_value=None)
memory_instance.embedding_model.embed = Mock(return_value=[0.1, 0.2, 0.3])
with patch("mem0.memory.main.lemmatize_for_bm25", return_value="test"), \
patch("mem0.memory.main.extract_entities", return_value=[]):
result = memory_instance.search("test", filters={"user_id": "test"}, threshold=0)
assert "results" in result
def test_search_accepts_threshold_one(self, memory_instance):
"""Search should accept threshold=1.0 (edge case)."""
mock_memories = []
memory_instance.vector_store.search = Mock(return_value=mock_memories)
memory_instance.vector_store.keyword_search = Mock(return_value=None)
memory_instance.embedding_model.embed = Mock(return_value=[0.1, 0.2, 0.3])
with patch("mem0.memory.main.lemmatize_for_bm25", return_value="test"), \
patch("mem0.memory.main.extract_entities", return_value=[]):
result = memory_instance.search("test", filters={"user_id": "test"}, threshold=1.0)
assert "results" in result
def test_search_accepts_top_k_zero(self, memory_instance):
"""Search should accept top_k=0."""
mock_memories = []
memory_instance.vector_store.search = Mock(return_value=mock_memories)
memory_instance.vector_store.keyword_search = Mock(return_value=None)
memory_instance.embedding_model.embed = Mock(return_value=[0.1, 0.2, 0.3])
with patch("mem0.memory.main.lemmatize_for_bm25", return_value="test"), \
patch("mem0.memory.main.extract_entities", return_value=[]):
result = memory_instance.search("test", filters={"user_id": "test"}, top_k=0)
assert "results" in result
+20 -20
View File
@@ -2,7 +2,7 @@
Verifies that the Pydantic request models in server/main.py correctly accept
and forward all parameters supported by the underlying Memory class methods,
including limit, threshold, infer, memory_type, and prompt — which were
including top_k, threshold, infer, memory_type, and prompt — which were
previously silently dropped by Pydantic v2's default extra='ignore' behavior.
"""
@@ -55,31 +55,31 @@ def mock_memory(_mock_memory):
# ===========================================================================
# SearchRequest: limit parameter
# SearchRequest: top_k parameter
# ===========================================================================
class TestSearchLimit:
"""Verify that the limit parameter is accepted and forwarded to Memory.search()."""
"""Verify that the top_k parameter is accepted and forwarded to Memory.search()."""
def test_limit_forwarded(self, client, mock_memory):
resp = client.post("/search", json={"query": "food", "user_id": "u1", "limit": 5})
resp = client.post("/search", json={"query": "food", "user_id": "u1", "top_k": 5})
assert resp.status_code == 200
_, kwargs = mock_memory.search.call_args
assert kwargs["limit"] == 5
assert kwargs["top_k"] == 5
def test_limit_one(self, client, mock_memory):
resp = client.post("/search", json={"query": "food", "user_id": "u1", "limit": 1})
resp = client.post("/search", json={"query": "food", "user_id": "u1", "top_k": 1})
assert resp.status_code == 200
_, kwargs = mock_memory.search.call_args
assert kwargs["limit"] == 1
assert kwargs["top_k"] == 1
def test_limit_omitted_uses_memory_default(self, client, mock_memory):
"""When limit is not sent, it should not appear in the kwargs,
"""When top_k is not sent, it should not appear in the kwargs,
allowing Memory.search() to use its own default (100)."""
resp = client.post("/search", json={"query": "food", "user_id": "u1"})
assert resp.status_code == 200
_, kwargs = mock_memory.search.call_args
assert "limit" not in kwargs
assert "top_k" not in kwargs
# ===========================================================================
@@ -110,18 +110,18 @@ class TestSearchThreshold:
# ===========================================================================
# SearchRequest: limit + threshold together
# SearchRequest: top_k + threshold together
# ===========================================================================
class TestSearchLimitAndThreshold:
def test_both_forwarded(self, client, mock_memory):
resp = client.post("/search", json={
"query": "food", "user_id": "u1", "limit": 10, "threshold": 0.5
"query": "food", "user_id": "u1", "top_k": 10, "threshold": 0.5
})
assert resp.status_code == 200
_, kwargs = mock_memory.search.call_args
assert kwargs["limit"] == 10
assert kwargs["top_k"] == 10
assert kwargs["threshold"] == 0.5
@@ -334,8 +334,8 @@ class TestOpenAPISchema:
def test_search_schema_includes_limit(self, client):
schema = client.get("/openapi.json").json()
search_props = schema["components"]["schemas"]["SearchRequest"]["properties"]
assert "limit" in search_props
assert search_props["limit"]["description"] == "Maximum number of results to return."
assert "top_k" in search_props
assert search_props["top_k"]["description"] == "Maximum number of results to return."
def test_search_schema_includes_threshold(self, client):
schema = client.get("/openapi.json").json()
@@ -367,7 +367,7 @@ class TestTypeValidation:
def test_limit_string_rejected(self, client):
resp = client.post("/search", json={
"query": "food", "user_id": "u1", "limit": "not_a_number",
"query": "food", "user_id": "u1", "top_k": "not_a_number",
})
assert resp.status_code == 422
@@ -399,7 +399,7 @@ class TestTypeValidation:
def test_limit_float_rejected(self, client):
resp = client.post("/search", json={
"query": "food", "user_id": "u1", "limit": 5.7,
"query": "food", "user_id": "u1", "top_k": 5.7,
})
assert resp.status_code == 422
@@ -422,11 +422,11 @@ class TestExplicitNull:
def test_limit_null_uses_memory_default(self, client, mock_memory):
resp = client.post("/search", json={
"query": "food", "user_id": "u1", "limit": None,
"query": "food", "user_id": "u1", "top_k": None,
})
assert resp.status_code == 200
_, kwargs = mock_memory.search.call_args
assert "limit" not in kwargs
assert "top_k" not in kwargs
def test_infer_null_uses_memory_default(self, client, mock_memory):
resp = client.post("/memories", json={
@@ -462,12 +462,12 @@ class TestCallSignatureMatch:
resp = client.post("/search", json={
"query": "food", "user_id": "u1", "agent_id": "a1",
"run_id": "r1", "filters": {"k": "v"},
"limit": 10, "threshold": 0.5,
"top_k": 10, "threshold": 0.5,
})
assert resp.status_code == 200
# The handler passes query= as a keyword arg, so it appears in kwargs too
_, kwargs = mock_memory.search.call_args
valid_params = {"query", "user_id", "agent_id", "run_id", "limit", "filters", "threshold", "rerank"}
valid_params = {"query", "user_id", "agent_id", "run_id", "top_k", "filters", "threshold", "rerank"}
for key in kwargs:
assert key in valid_params, f"Unexpected kwarg '{key}' forwarded to Memory.search()"
+60
View File
@@ -359,3 +359,63 @@ class TestTelemetryEnvVar:
def test_env_var_parsing(self, value, expected):
result = value.lower() in ("true", "1", "yes")
assert result == expected
class TestTelemetryNullUserIdHandling:
"""Verify telemetry doesn't crash when user_id is None.
This is a regression test for the bug where Memory.from_config() crashed
with AssertionError because distinct_id was None in PostHog.capture().
"""
def test_capture_event_skips_when_user_id_is_none(self):
"""AnonymousTelemetry.capture_event should not crash when user_id is None."""
with patch.object(telemetry_module, "MEM0_TELEMETRY", True):
with patch("mem0.memory.telemetry.Posthog") as mock_posthog_cls:
at = telemetry_module.AnonymousTelemetry()
at.user_id = None # Simulate the bug condition
# This should not raise, even with user_id=None
at.capture_event("test.event", {"key": "value"})
# PostHog.capture should NOT have been called since user_id is None
mock_posthog_cls.return_value.capture.assert_not_called()
def test_capture_event_does_not_crash_on_posthog_error(self):
"""AnonymousTelemetry.capture_event should catch PostHog exceptions."""
with patch.object(telemetry_module, "MEM0_TELEMETRY", True):
with patch("mem0.memory.telemetry.Posthog") as mock_posthog_cls:
with patch("mem0.memory.telemetry.get_or_create_user_id", return_value="test-user"):
at = telemetry_module.AnonymousTelemetry()
mock_posthog_cls.return_value.capture.side_effect = Exception("PostHog error")
# This should not raise, even when PostHog.capture fails
at.capture_event("test.event", {"key": "value"})
def test_oss_capture_event_does_not_crash_memory_init(self):
"""capture_event() should never raise, even if everything inside fails."""
with patch.object(telemetry_module, "MEM0_TELEMETRY", True):
# Make _get_oss_telemetry return a broken telemetry object
mock_at = MagicMock()
mock_at.capture_event.side_effect = Exception("Telemetry is broken")
with patch.object(telemetry_module, "_oss_telemetry_instance", mock_at):
mock_memory = MagicMock()
mock_memory.config.graph_store.config = None
mock_memory.api_version = "v1"
# This should not raise, even when telemetry fails
telemetry_module.capture_event("mem0.init", mock_memory)
def test_client_capture_event_does_not_crash(self):
"""capture_client_event() should never raise exceptions."""
with patch.object(telemetry_module, "MEM0_TELEMETRY", True):
mock_client_telemetry = MagicMock()
mock_client_telemetry.capture_event.side_effect = Exception("Client telemetry broken")
with patch.object(telemetry_module, "client_telemetry", mock_client_telemetry):
mock_instance = MagicMock()
mock_instance.user_email = "test@example.com"
# This should not raise
telemetry_module.capture_client_event("test.event", mock_instance)
-10
View File
@@ -101,16 +101,6 @@ class TestScoreAndRank:
scored = score_and_rank(results, {}, {}, threshold=0.1, top_k=5)
assert len(scored) == 5
def test_score_breakdown_present(self):
results = [{"id": "a", "score": 0.8, "payload": {"data": "x"}}]
bm25 = {"a": 0.4}
entity = {"a": 0.2}
scored = score_and_rank(results, bm25, entity, threshold=0.1, top_k=10)
breakdown = scored[0]["score_breakdown"]
assert breakdown["semantic"] == 0.8
assert breakdown["bm25"] == 0.4
assert breakdown["entity_boost"] == 0.2
def test_adaptive_divisor_semantic_only(self):
results = [{"id": "a", "score": 0.8, "payload": {}}]
scored = score_and_rank(results, {}, {}, threshold=0.1, top_k=10)
+368 -3
View File
@@ -1,4 +1,6 @@
import json
import os
import pickle
import tempfile
from unittest.mock import Mock, patch
@@ -6,7 +8,13 @@ import faiss
import numpy as np
import pytest
from mem0.vector_stores.faiss import FAISS, OutputData
from mem0.vector_stores.faiss import (
FAISS,
OutputData,
SafeUnpickler,
_safe_pickle_load,
_validate_docstore_structure,
)
@pytest.fixture
@@ -273,8 +281,8 @@ def test_delete_col(faiss_instance):
# Call delete_col
faiss_instance.delete_col()
# Verify os.remove was called twice (for index and docstore files)
assert mock_remove.call_count == 2
# Verify os.remove was called for index, json docstore, and legacy pkl files
assert mock_remove.call_count == 3
# Verify the internal state was reset
assert faiss_instance.index is None
@@ -299,3 +307,360 @@ def test_normalize_L2(faiss_instance, mock_faiss_index):
# Verify faiss.normalize_L2 was called
mock_normalize.assert_called_once()
# =============================================================================
# Security Tests for Pickle Deserialization Vulnerability Fix
# =============================================================================
class TestSafeUnpickler:
"""Tests for the SafeUnpickler class that prevents arbitrary code execution."""
def test_safe_unpickler_allows_basic_types(self):
"""SafeUnpickler should allow basic Python types."""
# Create a legitimate pickle with basic types
data = (
{"key1": "value1", "key2": {"nested": "dict"}},
{0: "id1", 1: "id2"},
)
pickled = pickle.dumps(data)
# Should load successfully
import io
result = SafeUnpickler(io.BytesIO(pickled)).load()
assert result == data
def test_safe_unpickler_blocks_os_system(self):
"""SafeUnpickler should block os.system execution attempts."""
# Generate the malicious payload dynamically to ensure correct format
import io
class Evil:
def __reduce__(self):
return (os.system, ("echo pwned",))
malicious_payload = pickle.dumps(Evil())
with pytest.raises(pickle.UnpicklingError) as exc_info:
SafeUnpickler(io.BytesIO(malicious_payload)).load()
assert "Unsafe pickle" in str(exc_info.value)
assert "posix.system" in str(exc_info.value)
def test_safe_unpickler_blocks_subprocess(self):
"""SafeUnpickler should block subprocess execution attempts."""
import subprocess
# Create a malicious pickle that tries to use subprocess
class MaliciousSubprocess:
def __reduce__(self):
return (subprocess.call, (["echo", "pwned"],))
malicious_payload = pickle.dumps(MaliciousSubprocess())
import io
with pytest.raises(pickle.UnpicklingError) as exc_info:
SafeUnpickler(io.BytesIO(malicious_payload)).load()
assert "Unsafe pickle" in str(exc_info.value)
def test_safe_unpickler_blocks_eval(self):
"""SafeUnpickler should block eval/exec attempts."""
# Create a malicious pickle that tries to use eval
class MaliciousEval:
def __reduce__(self):
return (eval, ("__import__('os').system('touch pwned')",))
malicious_payload = pickle.dumps(MaliciousEval())
import io
with pytest.raises(pickle.UnpicklingError) as exc_info:
SafeUnpickler(io.BytesIO(malicious_payload)).load()
assert "Unsafe pickle" in str(exc_info.value)
def test_safe_unpickler_blocks_arbitrary_modules(self):
"""SafeUnpickler should block imports from arbitrary modules."""
# Create a pickle that tries to load a class from a non-builtins module
class ArbitraryClass:
def __reduce__(self):
return (type, ("Evil", (), {}))
malicious_payload = pickle.dumps(ArbitraryClass())
import io
# This should either work (type is a builtin) or fail safely
# The key is it shouldn't execute arbitrary code
try:
result = SafeUnpickler(io.BytesIO(malicious_payload)).load()
# If it loads, verify it's just a benign type object
assert isinstance(result, type)
except pickle.UnpicklingError:
# This is also acceptable - blocking unknown patterns
pass
class TestSafePickleLoad:
"""Tests for the _safe_pickle_load function."""
def test_safe_pickle_load_with_valid_file(self):
"""_safe_pickle_load should load valid pickle files."""
with tempfile.NamedTemporaryFile(mode="wb", suffix=".pkl", delete=False) as f:
data = ({"id1": {"data": "test"}}, {0: "id1"})
pickle.dump(data, f)
temp_path = f.name
try:
result = _safe_pickle_load(temp_path)
assert result == data
finally:
os.unlink(temp_path)
def test_safe_pickle_load_blocks_malicious_file(self):
"""_safe_pickle_load should block malicious pickle files."""
# Generate the malicious payload dynamically
class Evil:
def __reduce__(self):
return (os.system, ("echo pwned",))
malicious_payload = pickle.dumps(Evil())
with tempfile.NamedTemporaryFile(mode="wb", suffix=".pkl", delete=False) as f:
f.write(malicious_payload)
temp_path = f.name
try:
with pytest.raises(pickle.UnpicklingError) as exc_info:
_safe_pickle_load(temp_path)
assert "Unsafe pickle" in str(exc_info.value)
finally:
os.unlink(temp_path)
class TestValidateDocstoreStructure:
"""Tests for the _validate_docstore_structure function."""
def test_valid_structure(self):
"""Should accept valid docstore structure."""
data = ({"id1": {"data": "test"}}, {0: "id1"})
docstore, index_to_id = _validate_docstore_structure(data)
assert docstore == {"id1": {"data": "test"}}
assert index_to_id == {0: "id1"}
def test_invalid_tuple_length(self):
"""Should reject tuples with wrong length."""
with pytest.raises(ValueError, match="expected tuple"):
_validate_docstore_structure(({}, {}, {}))
def test_invalid_docstore_type(self):
"""Should reject non-dict docstore."""
with pytest.raises(ValueError, match="docstore must be a dict"):
_validate_docstore_structure(("not a dict", {}))
def test_invalid_index_to_id_type(self):
"""Should reject non-dict index_to_id."""
with pytest.raises(ValueError, match="index_to_id must be a dict"):
_validate_docstore_structure(({}, "not a dict"))
def test_invalid_docstore_key_type(self):
"""Should reject non-string docstore keys."""
with pytest.raises(ValueError, match="Invalid docstore key type"):
_validate_docstore_structure(({123: {"data": "test"}}, {0: "id1"}))
def test_invalid_index_to_id_key_type(self):
"""Should reject non-int index_to_id keys."""
with pytest.raises(ValueError, match="Invalid index_to_id key type"):
_validate_docstore_structure(({"id1": {"data": "test"}}, {"0": "id1"}))
class TestFAISSSecurityIntegration:
"""Integration tests for FAISS security fixes."""
def test_faiss_saves_as_json(self):
"""FAISS should save docstore as JSON, not pickle."""
with tempfile.TemporaryDirectory() as temp_dir:
mock_index = Mock()
mock_index.d = 128
mock_index.ntotal = 0
with patch("mem0.vector_stores.faiss.faiss.IndexFlatL2", return_value=mock_index):
with patch("mem0.vector_stores.faiss.faiss.write_index"):
faiss_store = FAISS(
collection_name="test_security",
path=os.path.join(temp_dir, "test_faiss"),
distance_strategy="euclidean",
)
faiss_store.index = mock_index
# Insert some data
faiss_store.docstore = {"id1": {"data": "test"}}
faiss_store.index_to_id = {0: "id1"}
faiss_store._save()
# Verify JSON file was created
json_path = os.path.join(temp_dir, "test_faiss", "test_security.json")
pkl_path = os.path.join(temp_dir, "test_faiss", "test_security.pkl")
assert os.path.exists(json_path), "JSON docstore file should be created"
assert not os.path.exists(pkl_path), "Pickle file should NOT be created"
# Verify JSON content
with open(json_path, "r") as f:
data = json.load(f)
assert data["docstore"] == {"id1": {"data": "test"}}
assert data["index_to_id"] == {"0": "id1"}
def test_faiss_loads_json_preferentially(self):
"""FAISS should prefer JSON over pickle when both exist."""
with tempfile.TemporaryDirectory() as temp_dir:
faiss_path = os.path.join(temp_dir, "test_faiss")
os.makedirs(faiss_path)
# Create both JSON and pickle files with different data
json_data = {"docstore": {"id1": {"source": "json"}}, "index_to_id": {"0": "id1"}}
pkl_data = ({"id1": {"source": "pickle"}}, {0: "id1"})
with open(os.path.join(faiss_path, "test_pref.json"), "w") as f:
json.dump(json_data, f)
with open(os.path.join(faiss_path, "test_pref.pkl"), "wb") as f:
pickle.dump(pkl_data, f)
mock_index = Mock()
mock_index.d = 128
mock_index.ntotal = 1
with patch("mem0.vector_stores.faiss.faiss.read_index", return_value=mock_index):
with patch("mem0.vector_stores.faiss.faiss.write_index"):
faiss_store = FAISS.__new__(FAISS)
faiss_store.collection_name = "test_pref"
faiss_store.path = faiss_path
faiss_store.index = None
faiss_store.docstore = {}
faiss_store.index_to_id = {}
faiss_store._load(
os.path.join(faiss_path, "test_pref.faiss"),
os.path.join(faiss_path, "test_pref.pkl"),
)
# Should have loaded from JSON, not pickle
assert faiss_store.docstore == {"id1": {"source": "json"}}
def test_faiss_blocks_malicious_pickle_on_load(self):
"""FAISS should block loading of malicious pickle files."""
with tempfile.TemporaryDirectory() as temp_dir:
faiss_path = os.path.join(temp_dir, "test_faiss")
os.makedirs(faiss_path)
# Create a malicious pickle file (RCE payload)
class Evil:
def __reduce__(self):
return (os.system, (f"touch {temp_dir}/pwned",))
malicious_payload = pickle.dumps(Evil())
with open(os.path.join(faiss_path, "malicious.pkl"), "wb") as f:
f.write(malicious_payload)
mock_index = Mock()
mock_index.ntotal = 1
with patch("mem0.vector_stores.faiss.faiss.read_index", return_value=mock_index):
faiss_store = FAISS.__new__(FAISS)
faiss_store.collection_name = "malicious"
faiss_store.path = faiss_path
faiss_store.index = None
faiss_store.docstore = {}
faiss_store.index_to_id = {}
# Should raise an error, not execute the malicious payload
with pytest.raises(ValueError) as exc_info:
faiss_store._load(
os.path.join(faiss_path, "malicious.faiss"),
os.path.join(faiss_path, "malicious.pkl"),
)
assert "malicious pickle" in str(exc_info.value).lower() or "unsafe" in str(exc_info.value).lower()
# Verify the malicious command was NOT executed
pwned_file = os.path.join(temp_dir, "pwned")
assert not os.path.exists(pwned_file), "Malicious payload should NOT have been executed!"
def test_faiss_migrates_legacy_pickle_to_json(self):
"""FAISS should auto-migrate valid pickle files to JSON format."""
with tempfile.TemporaryDirectory() as temp_dir:
faiss_path = os.path.join(temp_dir, "test_faiss")
os.makedirs(faiss_path)
# Create a legitimate legacy pickle file
pkl_data = ({"id1": {"data": "legacy"}}, {0: "id1"})
with open(os.path.join(faiss_path, "legacy.pkl"), "wb") as f:
pickle.dump(pkl_data, f)
mock_index = Mock()
mock_index.d = 128
mock_index.ntotal = 1
with patch("mem0.vector_stores.faiss.faiss.read_index", return_value=mock_index):
with patch("mem0.vector_stores.faiss.faiss.write_index"):
faiss_store = FAISS.__new__(FAISS)
faiss_store.collection_name = "legacy"
faiss_store.path = faiss_path
faiss_store.index = None
faiss_store.docstore = {}
faiss_store.index_to_id = {}
faiss_store._load(
os.path.join(faiss_path, "legacy.faiss"),
os.path.join(faiss_path, "legacy.pkl"),
)
# Data should be loaded correctly
assert faiss_store.docstore == {"id1": {"data": "legacy"}}
assert faiss_store.index_to_id == {0: "id1"}
# JSON file should now exist (auto-migrated)
json_path = os.path.join(faiss_path, "legacy.json")
assert os.path.exists(json_path), "JSON file should be created during migration"
def test_delete_col_removes_json_and_pkl(self):
"""delete_col should remove both JSON and legacy pickle files."""
with tempfile.TemporaryDirectory() as temp_dir:
faiss_path = os.path.join(temp_dir, "test_faiss")
os.makedirs(faiss_path)
# Create both file types
json_path = os.path.join(faiss_path, "test_del.json")
pkl_path = os.path.join(faiss_path, "test_del.pkl")
faiss_index_path = os.path.join(faiss_path, "test_del.faiss")
with open(json_path, "w") as f:
json.dump({"docstore": {}, "index_to_id": {}}, f)
with open(pkl_path, "wb") as f:
pickle.dump(({}, {}), f)
with open(faiss_index_path, "w") as f:
f.write("dummy")
with patch("faiss.IndexFlatL2"):
faiss_store = FAISS.__new__(FAISS)
faiss_store.collection_name = "test_del"
faiss_store.path = faiss_path
faiss_store.index = Mock()
faiss_store.docstore = {}
faiss_store.index_to_id = {}
faiss_store.delete_col()
# Both files should be deleted
assert not os.path.exists(json_path), "JSON file should be deleted"
assert not os.path.exists(pkl_path), "PKL file should be deleted"
assert not os.path.exists(faiss_index_path), "FAISS index should be deleted"
+120 -6
View File
@@ -222,9 +222,7 @@ def test_update_with_none_vector_preserves_embedding(valkey_db, mock_valkey_clie
mock_valkey_client.hset.assert_called_once()
args, kwargs = mock_valkey_client.hset.call_args
assert "embedding" not in kwargs["mapping"], (
"embedding should not be in hash_data when vector is None"
)
assert "embedding" not in kwargs["mapping"], "embedding should not be in hash_data when vector is None"
assert kwargs["mapping"]["memory_id"] == "test_id"
assert kwargs["mapping"]["memory"] == "updated_data"
@@ -243,9 +241,7 @@ def test_update_with_vector_includes_embedding(valkey_db, mock_valkey_client):
mock_valkey_client.hset.assert_called_once()
args, kwargs = mock_valkey_client.hset.call_args
assert "embedding" in kwargs["mapping"], (
"embedding should be in hash_data when vector is provided"
)
assert "embedding" in kwargs["mapping"], "embedding should be in hash_data when vector is provided"
expected_bytes = np.array(vector, dtype=np.float32).tobytes()
assert kwargs["mapping"]["embedding"] == expected_bytes
@@ -906,3 +902,121 @@ def test_list_with_missing_fields_and_defaults(valkey_db, mock_valkey_client):
assert result.id == "fallback_id"
assert "hash" in result.payload
assert "data" in result.payload # memory is renamed to data
# Cluster mode tests
@pytest.fixture
def mock_valkey_cluster_client():
"""Create a mock ValkeyCluster client."""
with patch("valkey.cluster.ValkeyCluster.from_url") as mock_from_url:
mock_client = MagicMock()
mock_ft = MagicMock()
mock_client.ft = MagicMock(return_value=mock_ft)
mock_client.execute_command = MagicMock()
mock_client.hset = MagicMock()
mock_client.hgetall = MagicMock()
mock_client.delete = MagicMock()
mock_from_url.return_value = mock_client
yield mock_client
@pytest.fixture
def valkey_db_cluster(mock_valkey_cluster_client):
"""Create a ValkeyDB instance in cluster mode with a mock client."""
valkey_db = ValkeyDB(
valkey_url="valkey://localhost:7000",
collection_name="test_cluster",
embedding_model_dims=1536,
cluster_mode=True,
)
valkey_db.client = mock_valkey_cluster_client
return valkey_db
def test_cluster_mode_init(mock_valkey_cluster_client):
"""Test that cluster_mode=True uses ValkeyCluster client."""
db = ValkeyDB(
valkey_url="valkey://localhost:7000",
collection_name="test_cluster",
embedding_model_dims=1536,
cluster_mode=True,
)
assert db.cluster_mode is True
def test_cluster_mode_create_index(valkey_db_cluster, mock_valkey_cluster_client):
"""Test that index creation works in cluster mode (server handles propagation)."""
mock_valkey_cluster_client.execute_command.reset_mock()
mock_valkey_cluster_client.ft.return_value.info.side_effect = ResponseError("not found")
valkey_db_cluster._create_index(1536)
call_args = mock_valkey_cluster_client.execute_command.call_args
assert call_args is not None
assert "FT.CREATE" in call_args.args
def test_cluster_mode_drop_index(valkey_db_cluster, mock_valkey_cluster_client):
"""Test that dropping index works in cluster mode."""
mock_valkey_cluster_client.execute_command.reset_mock()
valkey_db_cluster._drop_index("test_cluster")
call_args = mock_valkey_cluster_client.execute_command.call_args
assert "FT.DROPINDEX" in call_args.args
def test_cluster_mode_search(valkey_db_cluster, mock_valkey_cluster_client):
"""Test that search in cluster mode uses ft().search() (server handles cross-shard fan-out)."""
ts = str(int(datetime.now().timestamp()))
mock_doc = MagicMock()
mock_doc.memory_id = "id1"
mock_doc.hash = "h1"
mock_doc.memory = "data1"
mock_doc.created_at = ts
mock_doc.metadata = "{}"
mock_doc.vector_score = "0.1"
mock_results = MagicMock()
mock_results.docs = [mock_doc]
mock_valkey_cluster_client.ft.return_value.search.return_value = mock_results
results = valkey_db_cluster.search("test", np.random.rand(1536).tolist(), top_k=5)
assert len(results) == 1
assert results[0].id == "id1"
mock_valkey_cluster_client.ft.return_value.search.assert_called_once()
def test_cluster_mode_insert(valkey_db_cluster, mock_valkey_cluster_client):
"""Test that insert works in cluster mode (ValkeyCluster handles routing)."""
vectors = [np.random.rand(1536).tolist()]
payloads = [{"hash": "h1", "data": "test", "user_id": "u1"}]
ids = ["id1"]
valkey_db_cluster.insert(vectors=vectors, payloads=payloads, ids=ids)
mock_valkey_cluster_client.hset.assert_called_once()
def test_cluster_mode_get(valkey_db_cluster, mock_valkey_cluster_client):
"""Test that get works in cluster mode."""
mock_valkey_cluster_client.hgetall.return_value = {
"memory_id": "id1",
"hash": "h1",
"memory": "test_data",
"created_at": str(int(datetime.now().timestamp())),
"metadata": "{}",
}
result = valkey_db_cluster.get("id1")
assert result.id == "id1"
assert result.payload["data"] == "test_data"
def test_cluster_mode_delete(valkey_db_cluster, mock_valkey_cluster_client):
"""Test that delete works in cluster mode."""
valkey_db_cluster.delete("id1")
mock_valkey_cluster_client.delete.assert_called_once_with("mem0:test_cluster:id1")