Compare commits
24 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 50db9e428d | |||
| fb87349664 | |||
| 8827553576 | |||
| c8e20a9bb5 | |||
| 93a51f4763 | |||
| 86fe275f53 | |||
| e6d6276bb9 | |||
| 9692726db4 | |||
| d8d776636f | |||
| a5a688295e | |||
| 5d40592e42 | |||
| a488e19044 | |||
| 57f944e18a | |||
| fe3f7ae618 | |||
| 4a7e166f9a | |||
| 85768e78e7 | |||
| 7b395f3bf7 | |||
| 4180409b09 | |||
| 649e719ce6 | |||
| 1a53852d93 | |||
| ac9cdd4840 | |||
| cf530c4bec | |||
| 92b958c1cc | |||
| c239d8a483 |
@@ -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>
|
||||
@@ -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
@@ -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",
|
||||
|
||||
@@ -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
@@ -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": {
|
||||
|
||||
@@ -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,5 +1,4 @@
|
||||
{
|
||||
"description": "Mem0 memory capture hooks — automatic memory extraction at key lifecycle points",
|
||||
"hooks": {
|
||||
"SessionStart": [
|
||||
{
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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();
|
||||
});
|
||||
});
|
||||
|
||||
@@ -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,
|
||||
},
|
||||
},
|
||||
|
||||
@@ -110,6 +110,7 @@ export class ConfigManager {
|
||||
((userConf as Record<string, unknown>)?.lmstudio_base_url as
|
||||
| string
|
||||
| undefined) ??
|
||||
userConf?.url ??
|
||||
defaultConf.baseURL;
|
||||
|
||||
return {
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -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,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 {
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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:",
|
||||
});
|
||||
|
||||
@@ -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
@@ -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()
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
@@ -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
@@ -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
@@ -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
@@ -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()
|
||||
|
||||
|
||||
@@ -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
@@ -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:
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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]:
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -134,6 +134,10 @@
|
||||
"topK": {
|
||||
"type": "number"
|
||||
},
|
||||
"anonymousTelemetryId": {
|
||||
"type": "string",
|
||||
"description": "Persistent anonymous telemetry identifier"
|
||||
},
|
||||
"oss": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
|
||||
@@ -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");
|
||||
});
|
||||
});
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
@@ -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
@@ -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" }
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
@@ -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
@@ -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()"
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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"
|
||||
|
||||
@@ -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")
|
||||
|
||||
Reference in New Issue
Block a user