From 28cc8fcee48807318d3c393fb183dbbe7bcb2a56 Mon Sep 17 00:00:00 2001 From: kartik-mem0 Date: Tue, 14 Apr 2026 16:17:50 +0530 Subject: [PATCH] refactor: update MemoryClient to use filters for entity parameters --- mem0-ts/src/client/index.ts | 1 - mem0-ts/src/client/mem0.ts | 60 ++- mem0-ts/src/client/mem0.types.ts | 15 +- .../src/client/tests/integration/crud.test.ts | 14 +- .../client/tests/integration/global-setup.ts | 10 +- .../tests/integration/global-teardown.ts | 10 +- .../src/client/tests/integration/helpers.ts | 12 +- .../client/tests/memoryClient.crud.test.ts | 20 +- .../client/tests/memoryClient.search.test.ts | 154 ++++++++ mem0-ts/src/oss/examples/basic.ts | 10 +- mem0-ts/src/oss/examples/local-llms.ts | 6 +- mem0-ts/src/oss/examples/utils/test-utils.ts | 10 +- mem0-ts/src/oss/src/memory/index.ts | 373 +++++++++++++----- mem0-ts/src/oss/src/memory/memory.types.ts | 17 +- mem0-ts/src/oss/src/types/index.ts | 6 +- mem0-ts/src/oss/src/vector_stores/memory.ts | 136 ++++++- mem0-ts/src/oss/src/vector_stores/qdrant.ts | 202 ++++++++-- mem0-ts/src/oss/tests/config-manager.test.ts | 10 +- .../oss/tests/dimension-autodetect.test.ts | 18 +- mem0-ts/src/oss/tests/memory.add.test.ts | 38 +- mem0-ts/src/oss/tests/memory.crud.test.ts | 70 ++-- mem0-ts/src/oss/tests/memory.init.test.ts | 10 +- .../oss/tests/vector-stores-compat.test.ts | 28 +- mem0/client/main.py | 79 +++- mem0/client/types.py | 33 +- mem0/memory/main.py | 332 +++++++++------- tests/test_client.py | 149 +++++++ tests/test_main.py | 11 +- tests/test_memory.py | 59 ++- 29 files changed, 1440 insertions(+), 453 deletions(-) create mode 100644 tests/test_client.py diff --git a/mem0-ts/src/client/index.ts b/mem0-ts/src/client/index.ts index b925909e2..24ec5b8f4 100644 --- a/mem0-ts/src/client/index.ts +++ b/mem0-ts/src/client/index.ts @@ -3,7 +3,6 @@ import type * as MemoryTypes from "./mem0.types"; // Re-export all types from mem0.types export type { - EntityOptions, AddMemoryOptions, SearchMemoryOptions, GetAllMemoryOptions, diff --git a/mem0-ts/src/client/mem0.ts b/mem0-ts/src/client/mem0.ts index d4bcc6f98..65f71c5c5 100644 --- a/mem0-ts/src/client/mem0.ts +++ b/mem0-ts/src/client/mem0.ts @@ -23,6 +23,28 @@ import { captureClientEvent, generateHash } from "./telemetry"; import { camelToSnake, camelToSnakeKeys, snakeToCamelKeys } from "./utils"; import { createExceptionFromResponse, MemoryError } from "../common/exceptions"; +// Entity parameters that must be passed via filters, not top-level (snake_case only - matches API) +const ENTITY_PARAMS = ["user_id", "agent_id", "app_id", "run_id"]; + +/** + * Validates that no top-level entity parameters are passed. + * @throws Error if entity params are found at top level + */ +function rejectTopLevelEntityParams( + options: Record | undefined, + methodName: string, +): void { + const invalidKeys = Object.keys(options ?? {}).filter((k) => + ENTITY_PARAMS.includes(k), + ); + if (invalidKeys.length > 0) { + throw new Error( + `Top-level entity parameters [${invalidKeys.join(", ")}] are not supported in ${methodName}(). ` + + `Use filters: { user_id: "..." } instead.`, + ); + } +} + class APIError extends Error { constructor(message: string) { super(message); @@ -185,7 +207,13 @@ export default class MemoryClient { ): Promise> { if (this.telemetryId === "") await this.ping(); - const payload = this._preparePayload(messages, options); + // Extract filters and spread entity IDs into payload (API expects top-level entity IDs) + const { filters, ...rest } = options; + const payload: Record = { + messages, + ...camelToSnakeKeys(rest), + ...(filters && filters), // Spread filters content into payload + }; const payloadKeys = Object.keys(payload); this._captureEvent("add", [payloadKeys]); @@ -254,6 +282,9 @@ export default class MemoryClient { } async getAll(options?: GetAllMemoryOptions): Promise> { + // Reject top-level entity params - must use filters instead + rejectTopLevelEntityParams(options as Record, "getAll"); + if (this.telemetryId === "") await this.ping(); const payloadKeys = Object.keys(options || {}); this._captureEvent("get_all", [payloadKeys]); @@ -281,6 +312,9 @@ export default class MemoryClient { query: string, options?: SearchMemoryOptions, ): Promise<{ results: Array }> { + // Reject top-level entity params - must use filters instead + rejectTopLevelEntityParams(options as Record, "search"); + if (this.telemetryId === "") await this.ping(); const payloadKeys = Object.keys(options || {}); this._captureEvent("search", [payloadKeys]); @@ -321,9 +355,27 @@ export default class MemoryClient { if (this.telemetryId === "") await this.ping(); const payloadKeys = Object.keys(options || {}); this._captureEvent("delete_all", [payloadKeys]); - const snakeOptions = camelToSnakeKeys(this._prepareParams(options)); - // @ts-ignore - const params = new URLSearchParams(snakeOptions); + + // Extract filters and build query params from filters (snake_case keys) + const { filters, ...rest } = options; + const queryParams: Record = {}; + if (filters) { + // Pass filter keys directly as query params (already snake_case) + for (const [key, value] of Object.entries(filters)) { + if (value != null) { + queryParams[key] = String(value); + } + } + } + // Add any other non-filter params + const otherParams = camelToSnakeKeys(this._prepareParams(rest)); + for (const [key, value] of Object.entries(otherParams)) { + if (value != null) { + queryParams[key] = String(value); + } + } + + const params = new URLSearchParams(queryParams); const response = await this._fetchWithErrorHandling( `${this.host}/v1/memories/?${params}`, { diff --git a/mem0-ts/src/client/mem0.types.ts b/mem0-ts/src/client/mem0.types.ts index af1f67a72..4ed4f7c8d 100644 --- a/mem0-ts/src/client/mem0.types.ts +++ b/mem0-ts/src/client/mem0.types.ts @@ -1,13 +1,6 @@ -// ─── Entity Options (for add/delete — top-level identity) ─── -export interface EntityOptions { - userId?: string; - agentId?: string; - appId?: string; - runId?: string; -} - // ─── Per-Method Options ───────────────────────────────────── -export interface AddMemoryOptions extends EntityOptions { +export interface AddMemoryOptions { + filters?: Record; metadata?: Record; infer?: boolean; customCategories?: custom_categories[]; @@ -35,7 +28,9 @@ export interface GetAllMemoryOptions { categories?: string[]; } -export interface DeleteAllMemoryOptions extends EntityOptions {} +export interface DeleteAllMemoryOptions { + filters?: Record; +} // ─── Project Options ──────────────────────────────────────── export interface ProjectOptions { diff --git a/mem0-ts/src/client/tests/integration/crud.test.ts b/mem0-ts/src/client/tests/integration/crud.test.ts index 78b51edb4..1d9698002 100644 --- a/mem0-ts/src/client/tests/integration/crud.test.ts +++ b/mem0-ts/src/client/tests/integration/crud.test.ts @@ -51,7 +51,9 @@ describeIntegration("MemoryClient Integration — CRUD", () => { }, ]; - const result = await client.add(messages, { userId: TEST_USER_ID }); + const result = await client.add(messages, { + filters: { user_id: TEST_USER_ID }, + }); // v3 API processes memories asynchronously — returns PENDING expect(result).toHaveProperty("eventId"); @@ -70,7 +72,9 @@ describeIntegration("MemoryClient Integration — CRUD", () => { }, ]; - const result = await client.add(messages, { userId: TEST_USER_ID }); + const result = await client.add(messages, { + filters: { user_id: TEST_USER_ID }, + }); expect(result).toHaveProperty("eventId"); }); @@ -204,7 +208,7 @@ describeIntegration("MemoryClient Integration — CRUD", () => { test("deleteAll for non-existent user does not throw", async () => { const result = await client.deleteAll({ - userId: `nonexistent-user-${randomUUID()}`, + filters: { user_id: `nonexistent-user-${randomUUID()}` }, }); expect(result).toBeDefined(); @@ -234,7 +238,9 @@ describeIntegration("MemoryClient Integration — CRUD", () => { // ─── Delete all + delete user ───────────────────────────── describe("cleanup operations", () => { test("deletes all memories for test user", async () => { - const result = await client.deleteAll({ userId: TEST_USER_ID }); + const result = await client.deleteAll({ + filters: { user_id: TEST_USER_ID }, + }); expect(result).toBeDefined(); expect(typeof result.message).toBe("string"); }); diff --git a/mem0-ts/src/client/tests/integration/global-setup.ts b/mem0-ts/src/client/tests/integration/global-setup.ts index 7adae5831..c1cf0f082 100644 --- a/mem0-ts/src/client/tests/integration/global-setup.ts +++ b/mem0-ts/src/client/tests/integration/global-setup.ts @@ -18,10 +18,12 @@ export default async function globalSetup() { // Full project wipe — all four filters set explicitly try { await client.deleteAll({ - userId: "*", - agentId: "*", - appId: "*", - runId: "*", + filters: { + user_id: "*", + agent_id: "*", + app_id: "*", + run_id: "*", + }, }); } catch { // ignore — may 404 if no data exists diff --git a/mem0-ts/src/client/tests/integration/global-teardown.ts b/mem0-ts/src/client/tests/integration/global-teardown.ts index 56815d556..dd4d1b16e 100644 --- a/mem0-ts/src/client/tests/integration/global-teardown.ts +++ b/mem0-ts/src/client/tests/integration/global-teardown.ts @@ -17,10 +17,12 @@ export default async function globalTeardown() { try { await client.deleteAll({ - userId: "*", - agentId: "*", - appId: "*", - runId: "*", + filters: { + user_id: "*", + agent_id: "*", + app_id: "*", + run_id: "*", + }, }); } catch { // ignore diff --git a/mem0-ts/src/client/tests/integration/helpers.ts b/mem0-ts/src/client/tests/integration/helpers.ts index b08274cac..1adfc9401 100644 --- a/mem0-ts/src/client/tests/integration/helpers.ts +++ b/mem0-ts/src/client/tests/integration/helpers.ts @@ -187,7 +187,7 @@ export async function cleanupTestUser( userId: string, ): Promise { try { - await client.deleteAll({ userId }); + await client.deleteAll({ filters: { user_id: userId } }); } catch { // ignore } @@ -210,10 +210,12 @@ export async function fullProjectCleanup(client: MemoryClient): Promise { // Delete all memories — all four filters set explicitly try { await client.deleteAll({ - userId: "*", - agentId: "*", - appId: "*", - runId: "*", + filters: { + user_id: "*", + agent_id: "*", + app_id: "*", + run_id: "*", + }, }); } catch { // ignore — may 404 if no data exists diff --git a/mem0-ts/src/client/tests/memoryClient.crud.test.ts b/mem0-ts/src/client/tests/memoryClient.crud.test.ts index ecbbf92c3..702cc14bd 100644 --- a/mem0-ts/src/client/tests/memoryClient.crud.test.ts +++ b/mem0-ts/src/client/tests/memoryClient.crud.test.ts @@ -27,7 +27,9 @@ describe("MemoryClient - add()", () => { const mock = setupMockFetch(extra); const client = new MemoryClient({ apiKey: TEST_API_KEY }); - await client.add([{ role: "user", content: "Hello" }], { userId: "u1" }); + await client.add([{ role: "user", content: "Hello" }], { + filters: { user_id: "u1" }, + }); expect(findFetchCall(mock, "/v3/memories/", "POST")).toBeDefined(); }); @@ -39,24 +41,26 @@ describe("MemoryClient - add()", () => { const mock = setupMockFetch(extra); const client = new MemoryClient({ apiKey: TEST_API_KEY }); - await client.add(messages, { userId: "u1" }); + await client.add(messages, { filters: { user_id: "u1" } }); const call = findFetchCall(mock, "/v3/memories/", "POST"); expect(getFetchBody(call!).messages).toEqual(messages); }); - test("includes user_id in request body", async () => { + test("spreads filters into top-level request body (API expects flat entity IDs)", async () => { const extra = new Map(); extra.set("/v3/memories/", { status: 200, body: [createMockMemory()] }); const mock = setupMockFetch(extra); const client = new MemoryClient({ apiKey: TEST_API_KEY }); await client.add([{ role: "user", content: "test" }], { - user_id: "user_1", + filters: { user_id: "user_1" }, }); const call = findFetchCall(mock, "/v3/memories/", "POST"); - expect(getFetchBody(call!).user_id).toBe("user_1"); + const body = getFetchBody(call!) as { user_id: string }; + // Entity IDs are spread into top level, not nested under filters + expect(body.user_id).toBe("user_1"); }); test("sends empty messages array without crashing", async () => { @@ -65,7 +69,7 @@ describe("MemoryClient - add()", () => { const mock = setupMockFetch(extra); const client = new MemoryClient({ apiKey: TEST_API_KEY }); - await client.add([], { userId: "u1" }); + await client.add([], { filters: { user_id: "u1" } }); const call = findFetchCall(mock, "/v3/memories/", "POST"); expect(getFetchBody(call!).messages).toEqual([]); @@ -215,7 +219,7 @@ describe("MemoryClient - deleteAll()", () => { const mock = setupMockFetch(extra); const client = new MemoryClient({ apiKey: TEST_API_KEY }); - await client.deleteAll({ userId: "u1" }); + await client.deleteAll({ filters: { user_id: "u1" } }); const call = mock.mock.calls.find( (c: [string, RequestInit]) => @@ -231,7 +235,7 @@ describe("MemoryClient - deleteAll()", () => { const mock = setupMockFetch(extra); const client = new MemoryClient({ apiKey: TEST_API_KEY }); - await client.deleteAll({ userId: "user@email.com" }); + await client.deleteAll({ filters: { user_id: "user@email.com" } }); const call = mock.mock.calls.find( (c: [string, RequestInit]) => diff --git a/mem0-ts/src/client/tests/memoryClient.search.test.ts b/mem0-ts/src/client/tests/memoryClient.search.test.ts index 0933faece..c29bbf88e 100644 --- a/mem0-ts/src/client/tests/memoryClient.search.test.ts +++ b/mem0-ts/src/client/tests/memoryClient.search.test.ts @@ -83,6 +83,89 @@ describe("MemoryClient - search()", () => { }); }); + test("passes complex AND filters through to the API body", async () => { + const extra = new Map(); + extra.set("/v3/memories/search/", { + status: 200, + body: { results: [] }, + }); + const mock = setupMockFetch(extra); + + const client = new MemoryClient({ apiKey: TEST_API_KEY }); + await client.search("query", { + filters: { + AND: [ + { user_id: "u1" }, + { created_at: { gte: "2024-01-01T00:00:00Z" } }, + ], + }, + }); + + const call = findFetchCall(mock, "/v3/memories/search/", "POST"); + const body = getFetchBody(call!); + expect(body.filters).toEqual({ + AND: [{ user_id: "u1" }, { created_at: { gte: "2024-01-01T00:00:00Z" } }], + }); + }); + + test("passes NOT filters through to the API body", async () => { + const extra = new Map(); + extra.set("/v3/memories/search/", { + status: 200, + body: { results: [] }, + }); + const mock = setupMockFetch(extra); + + const client = new MemoryClient({ apiKey: TEST_API_KEY }); + await client.search("query", { + filters: { + AND: [ + { user_id: "u1" }, + { NOT: { categories: { in: ["spam", "test"] } } }, + ], + }, + }); + + const call = findFetchCall(mock, "/v3/memories/search/", "POST"); + const body = getFetchBody(call!); + expect(body.filters).toEqual({ + AND: [ + { user_id: "u1" }, + { NOT: { categories: { in: ["spam", "test"] } } }, + ], + }); + }); + + test("passes complex nested AND/OR/NOT filters through to the API body", async () => { + const extra = new Map(); + extra.set("/v3/memories/search/", { + status: 200, + body: { results: [] }, + }); + const mock = setupMockFetch(extra); + + const client = new MemoryClient({ apiKey: TEST_API_KEY }); + const complexFilter = { + AND: [ + { user_id: "u1" }, + { created_at: { gte: "2024-01-01T00:00:00Z" } }, + { + NOT: { + OR: [ + { categories: { in: ["spam"] } }, + { categories: { in: ["test"] } }, + ], + }, + }, + ], + }; + await client.search("query", { filters: complexFilter }); + + const call = findFetchCall(mock, "/v3/memories/search/", "POST"); + const body = getFetchBody(call!); + expect(body.filters).toEqual(complexFilter); + }); + test("does not crash when called without options", async () => { const extra = new Map(); extra.set("/v3/memories/search/", { @@ -111,3 +194,74 @@ describe("MemoryClient - search()", () => { expect(result.results).toHaveLength(0); }); }); + +describe("MemoryClient - search() entity param rejection", () => { + test("rejects user_id at top level", async () => { + setupMockFetch(); + const client = new MemoryClient({ apiKey: TEST_API_KEY }); + await expect( + client.search("query", { user_id: "u1" } as any), + ).rejects.toThrow(/filters/); + }); + + test("rejects agent_id at top level", async () => { + setupMockFetch(); + const client = new MemoryClient({ apiKey: TEST_API_KEY }); + await expect( + client.search("query", { agent_id: "a1" } as any), + ).rejects.toThrow(/filters/); + }); + + test("rejects app_id at top level", async () => { + setupMockFetch(); + const client = new MemoryClient({ apiKey: TEST_API_KEY }); + await expect( + client.search("query", { app_id: "app1" } as any), + ).rejects.toThrow(/filters/); + }); + + test("accepts filters with user_id", async () => { + const extra = new Map(); + extra.set("/v3/memories/search/", { + status: 200, + body: { results: [] }, + }); + const mock = setupMockFetch(extra); + + const client = new MemoryClient({ apiKey: TEST_API_KEY }); + // Should not throw + await client.search("query", { filters: { user_id: "u1" } }); + expect(findFetchCall(mock, "/v3/memories/search/", "POST")).toBeDefined(); + }); +}); + +describe("MemoryClient - getAll() entity param rejection", () => { + test("rejects user_id at top level", async () => { + setupMockFetch(); + const client = new MemoryClient({ apiKey: TEST_API_KEY }); + await expect(client.getAll({ user_id: "u1" } as any)).rejects.toThrow( + /filters/, + ); + }); + + test("rejects agent_id at top level", async () => { + setupMockFetch(); + const client = new MemoryClient({ apiKey: TEST_API_KEY }); + await expect(client.getAll({ agent_id: "a1" } as any)).rejects.toThrow( + /filters/, + ); + }); + + test("accepts filters with user_id", async () => { + const extra = new Map(); + extra.set("/v2/memories/", { + status: 200, + body: { results: [] }, + }); + const mock = setupMockFetch(extra); + + const client = new MemoryClient({ apiKey: TEST_API_KEY }); + await client.getAll({ filters: { user_id: "u1" } }); + expect(findFetchCall(mock, "/v2/memories/", "POST")).toBeDefined(); + }); +}); diff --git a/mem0-ts/src/oss/examples/basic.ts b/mem0-ts/src/oss/examples/basic.ts index fdfd101cb..73f4f203c 100644 --- a/mem0-ts/src/oss/examples/basic.ts +++ b/mem0-ts/src/oss/examples/basic.ts @@ -30,7 +30,7 @@ async function runTests(memory: Memory) { const result1 = await memory.add( "Hi, my name is John and I am a software engineer.", { - userId: "john", + filters: { user_id: "john" }, }, ); console.log("Added memory:", result1); @@ -43,7 +43,7 @@ async function runTests(memory: Memory) { { role: "assistant", content: "I love Paris, it is my favorite city." }, ], { - userId: "john", + filters: { user_id: "john" }, }, ); console.log("Added messages:", result2); @@ -58,7 +58,7 @@ async function runTests(memory: Memory) { }, ], { - userId: "john", + filters: { user_id: "john" }, }, ); console.log("Updated messages:", result3); @@ -82,14 +82,14 @@ async function runTests(memory: Memory) { // Get all memories console.log("\nGetting all memories..."); const allMemories = await memory.getAll({ - userId: "john", + filters: { user_id: "john" }, }); console.log("All memories:", allMemories); // Search for memories console.log("\nSearching memories..."); const searchResult = await memory.search("What do you know about Paris?", { - userId: "john", + filters: { user_id: "john" }, }); console.log("Search results:", searchResult); diff --git a/mem0-ts/src/oss/examples/local-llms.ts b/mem0-ts/src/oss/examples/local-llms.ts index 29a8812c4..f86bf8adf 100644 --- a/mem0-ts/src/oss/examples/local-llms.ts +++ b/mem0-ts/src/oss/examples/local-llms.ts @@ -26,7 +26,9 @@ const memory = new Memory({ }); async function chatWithMemories(message: string, userId = "default_user") { - const relevantMemories = await memory.search(message, { userId: userId }); + const relevantMemories = await memory.search(message, { + filters: { user_id: userId }, + }); const memoriesStr = relevantMemories.results .map((entry) => `- ${entry.memory}`) @@ -50,7 +52,7 @@ ${memoriesStr}`; const assistantResponse = response.message.content || ""; messages.push({ role: "assistant", content: assistantResponse }); - await memory.add(messages, { userId: userId }); + await memory.add(messages, { filters: { user_id: userId } }); return assistantResponse; } diff --git a/mem0-ts/src/oss/examples/utils/test-utils.ts b/mem0-ts/src/oss/examples/utils/test-utils.ts index a89399dcf..3eddba249 100644 --- a/mem0-ts/src/oss/examples/utils/test-utils.ts +++ b/mem0-ts/src/oss/examples/utils/test-utils.ts @@ -12,7 +12,7 @@ export async function runTests(memory: Memory) { const result1 = await memory.add( "Hi, my name is John and I am a software engineer.", { - userId: "john", + filters: { user_id: "john" }, }, ); console.log("Added memory:", result1); @@ -25,7 +25,7 @@ export async function runTests(memory: Memory) { { role: "assistant", content: "I love Paris, it is my favorite city." }, ], { - userId: "john", + filters: { user_id: "john" }, }, ); console.log("Added messages:", result2); @@ -40,7 +40,7 @@ export async function runTests(memory: Memory) { }, ], { - userId: "john", + filters: { user_id: "john" }, }, ); console.log("Updated messages:", result3); @@ -64,14 +64,14 @@ export async function runTests(memory: Memory) { // Get all memories console.log("\nGetting all memories..."); const allMemories = await memory.getAll({ - userId: "john", + filters: { user_id: "john" }, }); console.log("All memories:", allMemories); // Search for memories console.log("\nSearching memories..."); const searchResult = await memory.search("What do you know about Paris?", { - userId: "john", + filters: { user_id: "john" }, }); console.log("Search results:", searchResult); diff --git a/mem0-ts/src/oss/src/memory/index.ts b/mem0-ts/src/oss/src/memory/index.ts index 08555f89d..69daeb6eb 100644 --- a/mem0-ts/src/oss/src/memory/index.ts +++ b/mem0-ts/src/oss/src/memory/index.ts @@ -53,6 +53,28 @@ import { ScoredResult, } from "../utils/scoring"; +// Entity parameters that must be passed via filters, not top-level (snake_case only - matches API) +const ENTITY_PARAMS = ["user_id", "agent_id", "run_id"]; + +/** + * Validates that no top-level entity parameters are passed in config. + * @throws Error if entity params are found at top level + */ +function rejectTopLevelEntityParams( + config: Record, + methodName: string, +): void { + const invalidKeys = Object.keys(config).filter((k) => + ENTITY_PARAMS.includes(k), + ); + if (invalidKeys.length > 0) { + throw new Error( + `Top-level entity parameters [${invalidKeys.join(", ")}] are not supported in ${methodName}(). ` + + `Use filters: { user_id: "..." } instead.`, + ); + } +} + export class Memory { private config: MemoryConfig; private customInstructions: string | undefined; @@ -184,7 +206,7 @@ export class Memory { private buildSessionScope(filters: SearchFilters): string { const parts: string[] = []; - for (const key of ["agentId", "runId", "userId"].sort()) { + for (const key of ["agent_id", "run_id", "user_id"].sort()) { const val = (filters as any)[key]; if (val) parts.push(`${key}=${val}`); } @@ -254,25 +276,21 @@ 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; - if (userId) filters.userId = metadata.userId = userId; - if (agentId) filters.agentId = metadata.agentId = agentId; - if (runId) filters.runId = metadata.runId = runId; - - if (!filters.userId && !filters.agentId && !filters.runId) { + // Validate filters contains at least one entity ID (snake_case) + if (!filters.user_id && !filters.agent_id && !filters.run_id) { throw new Error( - "One of the filters: userId, agentId or runId is required!", + "filters must contain at least one of: user_id, agent_id, run_id. " + + "Example: { filters: { user_id: 'u1' } }", ); } + // Copy entity IDs to metadata for storage + if (filters.user_id) metadata.user_id = filters.user_id; + if (filters.agent_id) metadata.agent_id = filters.agent_id; + if (filters.run_id) metadata.run_id = filters.run_id; + const parsedMessages = Array.isArray(messages) ? (messages as Message[]) : [{ role: "user", content: messages }]; @@ -337,16 +355,11 @@ export class Memory { const parsedMessages = messages.map((m) => m.content).join("\n"); // Phase 1: Existing memory retrieval - const searchFilters: SearchFilters = {}; - if (filters.userId) searchFilters.userId = filters.userId; - if (filters.agentId) searchFilters.agentId = filters.agentId; - if (filters.runId) searchFilters.runId = filters.runId; - const queryEmbedding = await this.embedder.embed(parsedMessages); const existingResults = await this.vectorStore.search( queryEmbedding, 10, - searchFilters, + filters, ); // Map UUIDs to integers (anti-hallucination) @@ -362,7 +375,7 @@ export class Memory { } // Phase 2: LLM extraction (single call) - const isAgentScoped = !!filters.agentId && !filters.userId; + const isAgentScoped = !!filters.agent_id && !filters.user_id; let systemPrompt = ADDITIVE_EXTRACTION_PROMPT; if (isAgentScoped) { systemPrompt += AGENT_CONTEXT_SUFFIX; @@ -491,9 +504,9 @@ export class Memory { if (mem.attributed_to) { memPayload.attributedTo = mem.attributed_to; } - if (filters.userId) memPayload.userId = filters.userId; - if (filters.agentId) memPayload.agentId = filters.agentId; - if (filters.runId) memPayload.runId = filters.runId; + if (filters.user_id) memPayload.user_id = filters.user_id; + if (filters.agent_id) memPayload.agent_id = filters.agent_id; + if (filters.run_id) memPayload.run_id = filters.run_id; records.push({ memoryId, @@ -661,7 +674,7 @@ export class Memory { payload: Record; }> = []; try { - matches = await entityStore.search(entityVec, 1, searchFilters); + matches = await entityStore.search(entityVec, 1, filters); } catch {} if (matches.length > 0 && (matches[0].score ?? 0) >= 0.95) { @@ -683,12 +696,9 @@ export class Memory { entityType, linkedMemoryIds: Array.from(memoryIds).sort(), }; - if (searchFilters.userId) - entityPayload.userId = searchFilters.userId; - if (searchFilters.agentId) - entityPayload.agentId = searchFilters.agentId; - if (searchFilters.runId) - entityPayload.runId = searchFilters.runId; + 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; toInsertVectors.push(entityVec); toInsertIds.push(uuidv4()); @@ -740,9 +750,9 @@ export class Memory { if (!memory) return null; const filters = { - ...(memory.payload.userId && { userId: memory.payload.userId }), - ...(memory.payload.agentId && { agentId: memory.payload.agentId }), - ...(memory.payload.runId && { runId: memory.payload.runId }), + ...(memory.payload.user_id && { user_id: memory.payload.user_id }), + ...(memory.payload.agent_id && { agent_id: memory.payload.agent_id }), + ...(memory.payload.run_id && { run_id: memory.payload.run_id }), }; const memoryItem: MemoryItem = { @@ -779,28 +789,46 @@ export class Memory { query: string, config: SearchMemoryOptions, ): Promise { + // Reject top-level entity params - must use filters instead + rejectTopLevelEntityParams(config as Record, "search"); + await this._ensureInitialized(); await this._captureEvent("search", { query_length: query.length, topK: config.topK, has_filters: !!config.filters, }); - const { - userId, - agentId, - runId, - topK = 100, - filters = {}, - threshold = 0.1, - } = config; + const { topK = 100, threshold = 0.1 } = config; + let effectiveFilters: Record = { ...(config.filters || {}) }; - if (userId) filters.userId = userId; - if (agentId) filters.agentId = agentId; - if (runId) filters.runId = runId; + // Apply enhanced metadata filtering if advanced operators are detected + if (this._hasAdvancedOperators(effectiveFilters)) { + const processedFilters = this._processMetadataFilters(effectiveFilters); + // Remove logical/operator keys that have been reprocessed + for (const logicalKey of ["AND", "OR", "NOT"]) { + delete effectiveFilters[logicalKey]; + } + for (const fk of Object.keys(effectiveFilters)) { + if ( + !["AND", "OR", "NOT", "user_id", "agent_id", "run_id"].includes(fk) && + typeof effectiveFilters[fk] === "object" && + effectiveFilters[fk] !== null + ) { + delete effectiveFilters[fk]; + } + } + effectiveFilters = { ...effectiveFilters, ...processedFilters }; + } - if (!filters.userId && !filters.agentId && !filters.runId) { + // Validate filters contains at least one entity ID (snake_case) + if ( + !effectiveFilters.user_id && + !effectiveFilters.agent_id && + !effectiveFilters.run_id + ) { throw new Error( - "One of the filters: userId, agentId or runId is required!", + "filters must contain at least one of: user_id, agent_id, run_id. " + + "Example: filters: { user_id: 'u1' }", ); } @@ -816,7 +844,7 @@ export class Memory { const semanticResults = await this.vectorStore.search( queryEmbedding, internalLimit, - filters, + effectiveFilters, ); // Step 4: Keyword search (if store supports it) @@ -831,7 +859,7 @@ export class Memory { (await this.vectorStore.keywordSearch( queryLemmatized, internalLimit, - filters, + effectiveFilters, )) ?? null; } catch { keywordResults = null; @@ -867,11 +895,6 @@ export class Memory { } if (deduped.length > 0) { - const searchFilters: SearchFilters = {}; - if (filters.userId) searchFilters.userId = filters.userId; - if (filters.agentId) searchFilters.agentId = filters.agentId; - if (filters.runId) searchFilters.runId = filters.runId; - const entityStore = await this.getEntityStore(); for (const entity of deduped) { @@ -880,7 +903,7 @@ export class Memory { const matches = await entityStore.search( entityEmbedding, 500, - searchFilters, + effectiveFilters, ); for (const match of matches) { @@ -936,9 +959,9 @@ export class Memory { // Step 9: Format results const excludedKeys = new Set([ - "userId", - "agentId", - "runId", + "user_id", + "agent_id", + "run_id", "hash", "data", "createdAt", @@ -964,9 +987,9 @@ export class Memory { .reduce((acc, [key, value]) => ({ ...acc, [key]: value }), {}), scoreBreakdown: scored.scoreBreakdown, }, - ...(payload.userId && { userId: payload.userId }), - ...(payload.agentId && { agentId: payload.agentId }), - ...(payload.runId && { runId: payload.runId }), + ...(payload.user_id && { user_id: payload.user_id }), + ...(payload.agent_id && { agent_id: payload.agent_id }), + ...(payload.run_id && { run_id: payload.run_id }), }; }); @@ -994,21 +1017,18 @@ export class Memory { config: DeleteAllMemoryOptions, ): Promise<{ message: string }> { await this._ensureInitialized(); + const { filters = {} } = config; + await this._captureEvent("delete_all", { - has_user_id: !!config.userId, - has_agent_id: !!config.agentId, - has_run_id: !!config.runId, + has_user_id: !!filters.user_id, + has_agent_id: !!filters.agent_id, + has_run_id: !!filters.run_id, }); - const { userId, agentId, runId } = config; - const filters: SearchFilters = {}; - if (userId) filters.userId = userId; - if (agentId) filters.agentId = agentId; - if (runId) filters.runId = runId; - - if (!Object.keys(filters).length) { + if (!filters.user_id && !filters.agent_id && !filters.run_id) { throw new Error( - "At least one filter is required to delete all memories. If you want to delete all memories, use the `reset()` method.", + "filters must contain at least one of: user_id, agent_id, run_id. " + + "If you want to delete all memories, use the `reset()` method.", ); } @@ -1077,26 +1097,33 @@ export class Memory { } async getAll(config: GetAllMemoryOptions): Promise { - await this._ensureInitialized(); - await this._captureEvent("get_all", { - topK: config.topK, - has_user_id: !!config.userId, - has_agent_id: !!config.agentId, - has_run_id: !!config.runId, - }); - const { userId, agentId, runId, topK = 100 } = config; + // Reject top-level entity params - must use filters instead + rejectTopLevelEntityParams(config as Record, "getAll"); - const filters: SearchFilters = {}; - if (userId) filters.userId = userId; - if (agentId) filters.agentId = agentId; - if (runId) filters.runId = runId; + await this._ensureInitialized(); + const { topK = 100, filters = {} } = config; + + await this._captureEvent("get_all", { + topK: topK, + has_user_id: !!filters.user_id, + has_agent_id: !!filters.agent_id, + has_run_id: !!filters.run_id, + }); + + // Validate filters contains at least one entity ID (snake_case) + if (!filters.user_id && !filters.agent_id && !filters.run_id) { + throw new Error( + "filters must contain at least one of: user_id, agent_id, run_id. " + + "Example: filters: { user_id: 'u1' }", + ); + } const [memories] = await this.vectorStore.list(filters, topK); const excludedKeys = new Set([ - "userId", - "agentId", - "runId", + "user_id", + "agent_id", + "run_id", "hash", "data", "createdAt", @@ -1113,9 +1140,9 @@ export class Memory { metadata: Object.entries(mem.payload) .filter(([key]) => !excludedKeys.has(key)) .reduce((acc, [key, value]) => ({ ...acc, [key]: value }), {}), - ...(mem.payload.userId && { userId: mem.payload.userId }), - ...(mem.payload.agentId && { agentId: mem.payload.agentId }), - ...(mem.payload.runId && { runId: mem.payload.runId }), + ...(mem.payload.user_id && { user_id: mem.payload.user_id }), + ...(mem.payload.agent_id && { agent_id: mem.payload.agent_id }), + ...(mem.payload.run_id && { run_id: mem.payload.run_id }), })); return { results }; @@ -1170,14 +1197,14 @@ export class Memory { hash: createHash("md5").update(data).digest("hex"), createdAt: existingMemory.payload.createdAt, updatedAt: new Date().toISOString(), - ...(existingMemory.payload.userId && { - userId: existingMemory.payload.userId, + ...(existingMemory.payload.user_id && { + user_id: existingMemory.payload.user_id, }), - ...(existingMemory.payload.agentId && { - agentId: existingMemory.payload.agentId, + ...(existingMemory.payload.agent_id && { + agent_id: existingMemory.payload.agent_id, }), - ...(existingMemory.payload.runId && { - runId: existingMemory.payload.runId, + ...(existingMemory.payload.run_id && { + run_id: existingMemory.payload.run_id, }), }; @@ -1214,4 +1241,156 @@ export class Memory { return memoryId; } + + /** + * Check if filters contain advanced operators that need special processing. + */ + private _hasAdvancedOperators(filters: Record): boolean { + if (!filters || typeof filters !== "object") { + return false; + } + + for (const [key, value] of Object.entries(filters)) { + // Check for platform-style logical operators + if (key === "AND" || key === "OR" || key === "NOT") { + return true; + } + // Check for comparison operators + if ( + typeof value === "object" && + value !== null && + !Array.isArray(value) + ) { + for (const op of Object.keys(value)) { + if ( + [ + "eq", + "ne", + "gt", + "gte", + "lt", + "lte", + "in", + "nin", + "contains", + "icontains", + ].includes(op) + ) { + return true; + } + } + } + // Check for wildcard values + if (value === "*") { + return true; + } + } + return false; + } + + /** + * Process enhanced metadata filters and convert them to vector store compatible format. + * Converts AND/OR/NOT to $or/$not format that vector stores can interpret. + */ + private _processMetadataFilters( + metadataFilters: Record, + ): Record { + const processedFilters: Record = {}; + + const processCondition = ( + key: string, + condition: any, + ): Record => { + if (typeof condition !== "object" || condition === null) { + // Simple equality: {"key": "value"} or wildcard + if (condition === "*") { + return { [key]: "*" }; + } + return { [key]: condition }; + } + + if (Array.isArray(condition)) { + // Array shorthand for "in" operator + return { [key]: { in: condition } }; + } + + const result: Record = {}; + const operatorMap: Record = { + eq: "eq", + ne: "ne", + gt: "gt", + gte: "gte", + lt: "lt", + lte: "lte", + in: "in", + nin: "nin", + contains: "contains", + icontains: "icontains", + }; + + for (const [operator, value] of Object.entries(condition)) { + if (operator in operatorMap) { + if (!result[key]) { + result[key] = {}; + } + result[key][operatorMap[operator]] = value; + } else { + throw new Error(`Unsupported metadata filter operator: ${operator}`); + } + } + return result; + }; + + for (const [key, value] of Object.entries(metadataFilters)) { + if (key === "AND") { + // Logical AND: combine multiple conditions + if (!Array.isArray(value)) { + throw new Error("AND operator requires a list of conditions"); + } + for (const condition of value) { + for (const [subKey, subValue] of Object.entries(condition)) { + Object.assign(processedFilters, processCondition(subKey, subValue)); + } + } + } else if (key === "OR") { + // Logical OR: Pass through to vector store for implementation-specific handling + if (!Array.isArray(value) || value.length === 0) { + throw new Error( + "OR operator requires a non-empty list of conditions", + ); + } + processedFilters["$or"] = []; + for (const condition of value) { + const orCondition: Record = {}; + for (const [subKey, subValue] of Object.entries( + condition as Record, + )) { + Object.assign(orCondition, processCondition(subKey, subValue)); + } + processedFilters["$or"].push(orCondition); + } + } else if (key === "NOT") { + // Logical NOT: Pass through to vector store for implementation-specific handling + if (!Array.isArray(value) || value.length === 0) { + throw new Error( + "NOT operator requires a non-empty list of conditions", + ); + } + processedFilters["$not"] = []; + for (const condition of value) { + const notCondition: Record = {}; + for (const [subKey, subValue] of Object.entries( + condition as Record, + )) { + Object.assign(notCondition, processCondition(subKey, subValue)); + } + processedFilters["$not"].push(notCondition); + } + } else { + Object.assign(processedFilters, processCondition(key, value)); + } + } + + return processedFilters; + } } diff --git a/mem0-ts/src/oss/src/memory/memory.types.ts b/mem0-ts/src/oss/src/memory/memory.types.ts index d1946cc3f..fd1bc4a54 100644 --- a/mem0-ts/src/oss/src/memory/memory.types.ts +++ b/mem0-ts/src/oss/src/memory/memory.types.ts @@ -1,26 +1,23 @@ import { Message } from "../types"; import { SearchFilters } from "../types"; -export interface Entity { - userId?: string; - agentId?: string; - runId?: string; -} - -export interface AddMemoryOptions extends Entity { +export interface AddMemoryOptions { metadata?: Record; filters?: SearchFilters; infer?: boolean; } -export interface SearchMemoryOptions extends Entity { +export interface SearchMemoryOptions { topK?: number; filters?: SearchFilters; threshold?: number; } -export interface GetAllMemoryOptions extends Entity { +export interface GetAllMemoryOptions { topK?: number; + filters?: SearchFilters; } -export interface DeleteAllMemoryOptions extends Entity {} +export interface DeleteAllMemoryOptions { + filters?: SearchFilters; +} diff --git a/mem0-ts/src/oss/src/types/index.ts b/mem0-ts/src/oss/src/types/index.ts index c2c2eaad3..d93061d33 100644 --- a/mem0-ts/src/oss/src/types/index.ts +++ b/mem0-ts/src/oss/src/types/index.ts @@ -81,9 +81,9 @@ export interface MemoryItem { } export interface SearchFilters { - userId?: string; - agentId?: string; - runId?: string; + user_id?: string; + agent_id?: string; + run_id?: string; [key: string]: any; } diff --git a/mem0-ts/src/oss/src/vector_stores/memory.ts b/mem0-ts/src/oss/src/vector_stores/memory.ts index 342256e39..1a4f2ec21 100644 --- a/mem0-ts/src/oss/src/vector_stores/memory.ts +++ b/mem0-ts/src/oss/src/vector_stores/memory.ts @@ -67,11 +67,139 @@ export class MemoryVectorStore implements VectorStore { return dotProduct / (Math.sqrt(normA) * Math.sqrt(normB)); } + /** + * Check if a single field condition matches the payload. + * Supports comparison operators: eq, ne, gt, gte, lt, lte, in, nin, contains, icontains + */ + private matchFieldCondition( + payload: Record, + key: string, + value: any, + ): boolean { + const payloadValue = payload[key]; + + // Handle non-dict values + if (typeof value !== "object" || value === null) { + // Wildcard: match any value + if (value === "*") { + return true; + } + // Simple equality + return payloadValue === value; + } + + // Handle array shorthand: {"field": ["a", "b"]} treated as "in" operator + if (Array.isArray(value)) { + return value.includes(payloadValue); + } + + // Handle comparison operators + if ("eq" in value) { + return payloadValue === value.eq; + } + if ("ne" in value) { + return payloadValue !== value.ne; + } + if ("gt" in value) { + return payloadValue > value.gt; + } + if ("gte" in value) { + return payloadValue >= value.gte; + } + if ("lt" in value) { + return payloadValue < value.lt; + } + if ("lte" in value) { + return payloadValue <= value.lte; + } + if ("in" in value) { + return Array.isArray(value.in) && value.in.includes(payloadValue); + } + if ("nin" in value) { + return !Array.isArray(value.nin) || !value.nin.includes(payloadValue); + } + if ("contains" in value) { + return ( + typeof payloadValue === "string" && + payloadValue.includes(value.contains) + ); + } + if ("icontains" in value) { + return ( + typeof payloadValue === "string" && + payloadValue.toLowerCase().includes(value.icontains.toLowerCase()) + ); + } + + // Unknown operator - treat as nested object for equality (shouldn't happen normally) + return payloadValue === value; + } + + /** + * Filter a vector by the given filters. + * Supports logical operators (AND, OR, NOT) and comparison operators. + */ private filterVector(vector: MemoryVector, filters?: SearchFilters): boolean { - if (!filters) return true; - return Object.entries(filters).every( - ([key, value]) => vector.payload[key] === value, - ); + if (!filters || Object.keys(filters).length === 0) return true; + + // Normalize $or/$not/$and → OR/NOT/AND + const keyMap: Record = { + $and: "AND", + $or: "OR", + $not: "NOT", + }; + const normalized: Record = {}; + for (const [key, value] of Object.entries(filters)) { + const normKey = keyMap[key] || key; + if (!(normKey in normalized)) { + normalized[normKey] = value; + } + } + + for (const [key, value] of Object.entries(normalized)) { + // Handle logical operators + if (key === "AND") { + if (!Array.isArray(value)) { + throw new Error( + `AND filter value must be a list of filter dicts, got ${typeof value}`, + ); + } + // All conditions must match + const allMatch = value.every((sub: SearchFilters) => + this.filterVector(vector, sub), + ); + if (!allMatch) return false; + } else if (key === "OR") { + if (!Array.isArray(value)) { + throw new Error( + `OR filter value must be a list of filter dicts, got ${typeof value}`, + ); + } + // At least one condition must match + const anyMatch = value.some((sub: SearchFilters) => + this.filterVector(vector, sub), + ); + if (!anyMatch) return false; + } else if (key === "NOT") { + if (!Array.isArray(value)) { + throw new Error( + `NOT filter value must be a list of filter dicts, got ${typeof value}`, + ); + } + // None of the conditions should match + const noneMatch = value.every( + (sub: SearchFilters) => !this.filterVector(vector, sub), + ); + if (!noneMatch) return false; + } else { + // Regular field condition + if (!this.matchFieldCondition(vector.payload, key, value)) { + return false; + } + } + } + + return true; } async insert( diff --git a/mem0-ts/src/oss/src/vector_stores/qdrant.ts b/mem0-ts/src/oss/src/vector_stores/qdrant.ts index 84eadfbf8..43a93031c 100644 --- a/mem0-ts/src/oss/src/vector_stores/qdrant.ts +++ b/mem0-ts/src/oss/src/vector_stores/qdrant.ts @@ -31,17 +31,29 @@ interface QdrantConfig extends VectorStoreConfig { } interface QdrantFilter { - must?: QdrantCondition[]; - must_not?: QdrantCondition[]; - should?: QdrantCondition[]; + must?: (QdrantCondition | QdrantFilter)[]; + must_not?: (QdrantCondition | QdrantFilter)[]; + should?: (QdrantCondition | QdrantFilter)[]; } interface QdrantCondition { key: string; - match?: { value: any }; - range?: { gte?: number; gt?: number; lte?: number; lt?: number }; + match?: { value?: any; any?: any[]; except?: any[]; text?: string }; + range?: { + gte?: number | string; + gt?: number | string; + lte?: number | string; + lt?: number | string; + }; } +// Normalize $and/$or/$not to AND/OR/NOT +const KEY_MAP: Record = { + $and: "AND", + $or: "OR", + $not: "NOT", +}; + export class Qdrant implements VectorStore { private client: QdrantClient; private readonly collectionName: string; @@ -90,35 +102,167 @@ export class Qdrant implements VectorStore { this.initialize().catch(console.error); } - private createFilter(filters?: SearchFilters): QdrantFilter | undefined { - if (!filters) return undefined; + /** + * Build a single field condition from a key-value filter pair. + * Supports enhanced filter syntax with comparison operators. + */ + private buildFieldCondition(key: string, value: any): QdrantCondition | null { + // Handle non-dict values + if (typeof value !== "object" || value === null) { + // Wildcard: match any value - skip this filter + if (value === "*") { + return null; + } + // Simple equality + return { key, match: { value } }; + } - const conditions: QdrantCondition[] = []; + // Handle array shorthand: {"field": ["a", "b"]} treated as "in" operator + if (Array.isArray(value)) { + return { key, match: { any: value } }; + } + + const ops = Object.keys(value); + const rangeOps = ["gt", "gte", "lt", "lte"]; + const hasRangeOps = ops.some((op) => rangeOps.includes(op)); + const nonRangeOps = ops.filter((op) => !rangeOps.includes(op)); + + // Handle range operators + if (hasRangeOps) { + if (nonRangeOps.length > 0) { + throw new Error( + `Cannot mix range operators (${ops.filter((o) => rangeOps.includes(o)).join(", ")}) ` + + `with non-range operators (${nonRangeOps.join(", ")}) for field '${key}'. ` + + `Use AND to combine them as separate conditions.`, + ); + } + const range: Record = {}; + for (const op of rangeOps) { + if (op in value) { + range[op] = value[op]; + } + } + return { key, range }; + } + + // Handle comparison operators + if ("eq" in value) { + return { key, match: { value: value.eq } }; + } + if ("ne" in value) { + return { key, match: { except: [value.ne] } }; + } + if ("in" in value) { + return { key, match: { any: value.in } }; + } + if ("nin" in value) { + return { key, match: { except: value.nin } }; + } + if ("contains" in value || "icontains" in value) { + const text = value.contains || value.icontains; + return { key, match: { text } }; + } + + // Unknown operator - treat as nested object for simple match + const supportedOps = [ + "eq", + "ne", + "gt", + "gte", + "lt", + "lte", + "in", + "nin", + "contains", + "icontains", + ]; + throw new Error( + `Unsupported filter operator(s) for field '${key}': ${ops.join(", ")}. ` + + `Supported operators: ${supportedOps.join(", ")}`, + ); + } + + /** + * Create a Filter object from the provided filters. + * Supports logical operators (AND, OR, NOT) and comparison operators. + */ + private createFilter(filters?: SearchFilters): QdrantFilter | undefined { + if (!filters || Object.keys(filters).length === 0) return undefined; + + // Normalize $or/$not/$and → OR/NOT/AND and deduplicate + const normalized: Record = {}; for (const [key, value] of Object.entries(filters)) { - if ( - typeof value === "object" && - value !== null && - "gte" in value && - "lte" in value - ) { - conditions.push({ - key, - range: { - gte: value.gte, - lte: value.lte, - }, - }); - } else { - conditions.push({ - key, - match: { - value, - }, - }); + const normKey = KEY_MAP[key] || key; + if (!(normKey in normalized)) { + normalized[normKey] = value; } } - return conditions.length ? { must: conditions } : undefined; + const must: (QdrantCondition | QdrantFilter)[] = []; + const should: (QdrantCondition | QdrantFilter)[] = []; + const mustNot: (QdrantCondition | QdrantFilter)[] = []; + + for (const [key, value] of Object.entries(normalized)) { + // Handle logical operators + if (key === "AND" || key === "OR" || key === "NOT") { + if (!Array.isArray(value)) { + throw new Error( + `${key} filter value must be a list of filter dicts, got ${typeof value}`, + ); + } + for (let i = 0; i < value.length; i++) { + const item = value[i]; + if ( + typeof item !== "object" || + item === null || + Array.isArray(item) + ) { + throw new Error( + `${key} filter list item at index ${i} must be a dict, got ${typeof item}`, + ); + } + } + + if (key === "AND") { + for (const sub of value) { + const built = this.createFilter(sub); + if (built) { + must.push(built); + } + } + } else if (key === "OR") { + for (const sub of value) { + const built = this.createFilter(sub); + if (built) { + should.push(built); + } + } + } else if (key === "NOT") { + for (const sub of value) { + const built = this.createFilter(sub); + if (built) { + mustNot.push(built); + } + } + } + } else { + // Regular field condition + const condition = this.buildFieldCondition(key, value); + if (condition !== null) { + must.push(condition); + } + } + } + + if (must.length === 0 && should.length === 0 && mustNot.length === 0) { + return undefined; + } + + return { + must: must.length > 0 ? must : undefined, + should: should.length > 0 ? should : undefined, + must_not: mustNot.length > 0 ? mustNot : undefined, + }; } async insert( diff --git a/mem0-ts/src/oss/tests/config-manager.test.ts b/mem0-ts/src/oss/tests/config-manager.test.ts index 8444564db..2eb01c410 100644 --- a/mem0-ts/src/oss/tests/config-manager.test.ts +++ b/mem0-ts/src/oss/tests/config-manager.test.ts @@ -434,7 +434,7 @@ describe("Memory – LM Studio end-to-end flow", () => { disableHistory: true, }); - await mem.getAll({ userId: "u1" }); + await mem.getAll({ filters: { user_id: "u1" } }); expect(mockEmbedderFactory.create).toHaveBeenCalledWith( "lmstudio", @@ -469,7 +469,7 @@ describe("Memory – LM Studio end-to-end flow", () => { disableHistory: true, }); - await mem.getAll({ userId: "u1" }); + await mem.getAll({ filters: { user_id: "u1" } }); expect(mockEmbedder.embed).toHaveBeenCalledWith("dimension probe"); const vsCall = mockVectorStoreFactory.create.mock.calls[0]; @@ -497,7 +497,7 @@ describe("Memory – LM Studio end-to-end flow", () => { disableHistory: true, }); - await mem.getAll({ userId: "u1" }); + await mem.getAll({ filters: { user_id: "u1" } }); expect(mockEmbedderFactory.create).toHaveBeenCalledWith( "lmstudio", @@ -550,7 +550,7 @@ describe("Memory – LM Studio end-to-end flow", () => { }); const result = await mem.search("What does the user like?", { - userId: "u1", + filters: { user_id: "u1" }, }); expect(mockEmbedder.embed).toHaveBeenCalledWith("What does the user like?"); @@ -589,7 +589,7 @@ describe("Memory – LM Studio end-to-end flow", () => { disableHistory: true, }); - await mem.add("I love sushi", { userId: "u1" }); + await mem.add("I love sushi", { filters: { user_id: "u1" } }); expect(mockLlm.generateResponse).toHaveBeenCalled(); expect(mockEmbedder.embed).toHaveBeenCalled(); diff --git a/mem0-ts/src/oss/tests/dimension-autodetect.test.ts b/mem0-ts/src/oss/tests/dimension-autodetect.test.ts index 94013dd00..c63c92c3e 100644 --- a/mem0-ts/src/oss/tests/dimension-autodetect.test.ts +++ b/mem0-ts/src/oss/tests/dimension-autodetect.test.ts @@ -313,7 +313,7 @@ describe("Memory – auto-initialization", () => { disableHistory: true, }); - await mem.getAll({ userId: "u1" }); + await mem.getAll({ filters: { user_id: "u1" } }); // Should have called embed("dimension probe") to detect dimension expect(mockEmbedder.embed).toHaveBeenCalledWith("dimension probe"); @@ -339,7 +339,7 @@ describe("Memory – auto-initialization", () => { disableHistory: true, }); - await mem.getAll({ userId: "u1" }); + await mem.getAll({ filters: { user_id: "u1" } }); // embed should NOT have been called for probing expect(mockEmbedder.embed).not.toHaveBeenCalledWith("dimension probe"); @@ -365,7 +365,7 @@ describe("Memory – auto-initialization", () => { disableHistory: true, }); - await mem.getAll({ userId: "u1" }); + await mem.getAll({ filters: { user_id: "u1" } }); // ConfigManager resolves dimension from embeddingDims → no probe needed expect(mockEmbedder.embed).not.toHaveBeenCalledWith("dimension probe"); @@ -403,9 +403,11 @@ describe("Memory – auto-initialization", () => { let searchDone = false; let getDone = false; - const getAllP = mem.getAll({ userId: "u" }).then(() => (getAllDone = true)); + const getAllP = mem + .getAll({ filters: { user_id: "u" } }) + .then(() => (getAllDone = true)); const searchP = mem - .search("q", { userId: "u" }) + .search("q", { filters: { user_id: "u" } }) .then(() => (searchDone = true)); const getP = mem.get("id").then(() => (getDone = true)); @@ -435,7 +437,7 @@ describe("Memory – auto-initialization", () => { disableHistory: true, }); - await mem.getAll({ userId: "u1" }); + await mem.getAll({ filters: { user_id: "u1" } }); expect(mockVectorStoreFactory.create).toHaveBeenCalledTimes(1); // Reset should re-create vector store @@ -473,7 +475,7 @@ describe("Memory – auto-initialization", () => { disableHistory: true, }); - await mem.getAll({ userId: "u1" }); + await mem.getAll({ filters: { user_id: "u1" } }); expect(mockEmbedder.embed).not.toHaveBeenCalledWith("dimension probe"); }); @@ -497,7 +499,7 @@ describe("Memory – auto-initialization", () => { }); // getAll should reject with the init error - await expect(mem.getAll({ userId: "u1" })).rejects.toThrow( + await expect(mem.getAll({ filters: { user_id: "u1" } })).rejects.toThrow( "auto-detect embedding dimension", ); diff --git a/mem0-ts/src/oss/tests/memory.add.test.ts b/mem0-ts/src/oss/tests/memory.add.test.ts index c1a54c30a..6251681ae 100644 --- a/mem0-ts/src/oss/tests/memory.add.test.ts +++ b/mem0-ts/src/oss/tests/memory.add.test.ts @@ -96,7 +96,7 @@ describe("Memory - add()", () => { test("returns SearchResult with results array for string input", async () => { const result: SearchResult = await memory.add("I am a software engineer", { - userId, + filters: { user_id: userId }, }); expect(Array.isArray(result.results)).toBe(true); }); @@ -104,7 +104,7 @@ describe("Memory - add()", () => { test("returns at least one result with an id", async () => { const result: SearchResult = await memory.add( "I enjoy hiking in the mountains", - { userId }, + { filters: { user_id: userId } }, ); expect(result.results.length).toBeGreaterThan(0); expect(result.results[0].id).toBeDefined(); @@ -112,7 +112,7 @@ describe("Memory - add()", () => { test("result item has a memory string field", async () => { const result: SearchResult = await memory.add("My favorite color is blue", { - userId, + filters: { user_id: userId }, }); expect(typeof result.results[0].memory).toBe("string"); }); @@ -122,31 +122,35 @@ describe("Memory - add()", () => { { role: "user", content: "What is your favorite city?" }, { role: "assistant", content: "I love Paris." }, ]; - const result: SearchResult = await memory.add(messages, { userId }); - expect(result.results.length).toBeGreaterThan(0); - }); - - test("works with agentId filter instead of userId", async () => { - const result: SearchResult = await memory.add("test", { - agentId: "agent_1", + const result: SearchResult = await memory.add(messages, { + filters: { user_id: userId }, }); expect(result.results.length).toBeGreaterThan(0); }); - test("works with runId filter instead of userId", async () => { - const result: SearchResult = await memory.add("test", { runId: "run_1" }); + test("works with agent_id filter instead of user_id", async () => { + const result: SearchResult = await memory.add("test", { + filters: { agent_id: "agent_1" }, + }); expect(result.results.length).toBeGreaterThan(0); }); - test("throws when no userId/agentId/runId provided", async () => { + test("works with run_id filter instead of user_id", async () => { + const result: SearchResult = await memory.add("test", { + filters: { run_id: "run_1" }, + }); + expect(result.results.length).toBeGreaterThan(0); + }); + + test("throws when no user_id/agent_id/run_id provided in filters", async () => { await expect(memory.add("test", {} as any)).rejects.toThrow( - "One of the filters: userId, agentId or runId is required!", + "filters must contain at least one of: user_id, agent_id, run_id", ); }); test("passes metadata through to stored memory", async () => { const result: SearchResult = await memory.add("I love TypeScript", { - userId, + filters: { user_id: userId }, metadata: { source: "chat", tag: "programming" }, }); const stored: MemoryItem | null = await memory.get(result.results[0].id); @@ -158,7 +162,7 @@ describe("Memory - add()", () => { test("with infer=false skips LLM and stores messages directly", async () => { const result: SearchResult = await memory.add("Direct storage content", { - userId, + filters: { user_id: userId }, infer: false, }); expect(result.results.length).toBeGreaterThan(0); @@ -168,7 +172,7 @@ describe("Memory - add()", () => { test("with infer=false marks event as ADD in metadata", async () => { const result: SearchResult = await memory.add("Direct fact", { - userId, + filters: { user_id: userId }, infer: false, }); expect(result.results[0].metadata).toEqual( diff --git a/mem0-ts/src/oss/tests/memory.crud.test.ts b/mem0-ts/src/oss/tests/memory.crud.test.ts index bc055220c..d813de60f 100644 --- a/mem0-ts/src/oss/tests/memory.crud.test.ts +++ b/mem0-ts/src/oss/tests/memory.crud.test.ts @@ -96,7 +96,9 @@ describe("Memory - get()", () => { }); test("returns the memory matching the ID from add()", async () => { - const addResult: SearchResult = await memory.add("I love AI", { userId }); + const addResult: SearchResult = await memory.add("I love AI", { + filters: { user_id: userId }, + }); const id = addResult.results[0].id; const item: MemoryItem | null = await memory.get(id); expect(item).not.toBeNull(); @@ -105,7 +107,7 @@ describe("Memory - get()", () => { test("returns a string for the memory field", async () => { const addResult: SearchResult = await memory.add("Testing get", { - userId, + filters: { user_id: userId }, }); const item: MemoryItem | null = await memory.get(addResult.results[0].id); expect(typeof item!.memory).toBe("string"); @@ -117,7 +119,9 @@ describe("Memory - get()", () => { }); test("returns hash and createdAt on stored memory", async () => { - const addResult: SearchResult = await memory.add("Hash test", { userId }); + const addResult: SearchResult = await memory.add("Hash test", { + filters: { user_id: userId }, + }); const item: MemoryItem | null = await memory.get(addResult.results[0].id); expect(typeof item!.hash).toBe("string"); expect(item!.createdAt).toBeDefined(); @@ -142,7 +146,7 @@ describe("Memory - update()", () => { // Use infer: false for update tests — bypasses LLM, gives us a stable ID test("returns success message", async () => { const addResult: SearchResult = await memory.add("Original", { - userId, + filters: { user_id: userId }, infer: false, }); const id = addResult.results[0].id; @@ -152,7 +156,7 @@ describe("Memory - update()", () => { test("persists the updated text", async () => { const addResult: SearchResult = await memory.add("Before update", { - userId, + filters: { user_id: userId }, infer: false, }); const id = addResult.results[0].id; @@ -163,7 +167,7 @@ describe("Memory - update()", () => { test("preserves createdAt and sets updatedAt", async () => { const addResult: SearchResult = await memory.add("Timestamp test", { - userId, + filters: { user_id: userId }, infer: false, }); const id = addResult.results[0].id; @@ -178,7 +182,7 @@ describe("Memory - update()", () => { test("updates the hash", async () => { const addResult: SearchResult = await memory.add("Hash change", { - userId, + filters: { user_id: userId }, infer: false, }); const id = addResult.results[0].id; @@ -205,7 +209,7 @@ describe("Memory - delete()", () => { test("returns success message", async () => { const addResult: SearchResult = await memory.add("Delete me", { - userId, + filters: { user_id: userId }, infer: false, }); const result = await memory.delete(addResult.results[0].id); @@ -214,7 +218,7 @@ describe("Memory - delete()", () => { test("get() returns null after deletion", async () => { const addResult: SearchResult = await memory.add("Temporary", { - userId, + filters: { user_id: userId }, infer: false, }); const id = addResult.results[0].id; @@ -238,17 +242,19 @@ describe("Memory - deleteAll()", () => { }); test("removes all memories for the user and returns success", async () => { - await memory.add("Fact A", { userId }); - await memory.add("Fact B", { userId }); - const result = await memory.deleteAll({ userId }); + await memory.add("Fact A", { filters: { user_id: userId } }); + await memory.add("Fact B", { filters: { user_id: userId } }); + const result = await memory.deleteAll({ filters: { user_id: userId } }); expect(result.message).toBe("Memories deleted successfully!"); - const remaining: SearchResult = await memory.getAll({ userId }); + const remaining: SearchResult = await memory.getAll({ + filters: { user_id: userId }, + }); expect(remaining.results).toHaveLength(0); }); test("throws when no filter is provided", async () => { await expect(memory.deleteAll({} as any)).rejects.toThrow( - "At least one filter is required", + "filters must contain at least one of: user_id, agent_id, run_id", ); }); }); @@ -268,15 +274,19 @@ describe("Memory - getAll()", () => { }); test("returns all stored memories for the user", async () => { - await memory.add("First", { userId }); - await memory.add("Second", { userId }); - const result: SearchResult = await memory.getAll({ userId }); + await memory.add("First", { filters: { user_id: userId } }); + await memory.add("Second", { filters: { user_id: userId } }); + const result: SearchResult = await memory.getAll({ + filters: { user_id: userId }, + }); expect(Array.isArray(result.results)).toBe(true); expect(result.results.length).toBeGreaterThanOrEqual(2); }); test("each result has id and memory fields", async () => { - const result: SearchResult = await memory.getAll({ userId }); + const result: SearchResult = await memory.getAll({ + filters: { user_id: userId }, + }); for (const item of result.results) { expect(item.id).toBeDefined(); expect(typeof item.memory).toBe("string"); @@ -285,7 +295,7 @@ describe("Memory - getAll()", () => { test("returns empty array when no memories exist", async () => { const result: SearchResult = await memory.getAll({ - userId: "no_such_user", + filters: { user_id: "no_such_user" }, }); expect(result.results).toHaveLength(0); }); @@ -299,7 +309,7 @@ describe("Memory - search()", () => { beforeAll(async () => { memory = createMemory(); - await memory.add("I love TypeScript", { userId }); + await memory.add("I love TypeScript", { filters: { user_id: userId } }); }); afterAll(async () => { @@ -307,12 +317,16 @@ describe("Memory - search()", () => { }); test("returns SearchResult with results array", async () => { - const result: SearchResult = await memory.search("TypeScript", { userId }); + const result: SearchResult = await memory.search("TypeScript", { + filters: { user_id: userId }, + }); expect(Array.isArray(result.results)).toBe(true); }); test("returns results with score field", async () => { - const result: SearchResult = await memory.search("content", { userId }); + const result: SearchResult = await memory.search("content", { + filters: { user_id: userId }, + }); if (result.results.length > 0) { expect(typeof result.results[0].score).toBe("number"); } @@ -320,13 +334,13 @@ describe("Memory - search()", () => { test("throws when no userId/agentId/runId provided", async () => { await expect(memory.search("query", {} as any)).rejects.toThrow( - "One of the filters: userId, agentId or runId is required!", + "filters must contain at least one of: user_id, agent_id, run_id", ); }); test("returns empty results for user with no memories", async () => { const result: SearchResult = await memory.search("query", { - userId: "empty_user", + filters: { user_id: "empty_user" }, }); expect(result.results).toHaveLength(0); }); @@ -347,14 +361,18 @@ describe("Memory - history()", () => { }); test("records ADD event after add()", async () => { - const addResult: SearchResult = await memory.add("New fact", { userId }); + const addResult: SearchResult = await memory.add("New fact", { + filters: { user_id: userId }, + }); const history = await memory.history(addResult.results[0].id); expect(Array.isArray(history)).toBe(true); expect(history.length).toBeGreaterThan(0); }); test("records additional entry after update()", async () => { - const addResult: SearchResult = await memory.add("Before", { userId }); + const addResult: SearchResult = await memory.add("Before", { + filters: { user_id: userId }, + }); const id = addResult.results[0].id; await memory.update(id, "After"); const history = await memory.history(id); diff --git a/mem0-ts/src/oss/tests/memory.init.test.ts b/mem0-ts/src/oss/tests/memory.init.test.ts index 4e429ee5e..c9d7ee9b1 100644 --- a/mem0-ts/src/oss/tests/memory.init.test.ts +++ b/mem0-ts/src/oss/tests/memory.init.test.ts @@ -118,13 +118,17 @@ describe("Memory - reset()", () => { const mem = createMemory(); const userId = `reset_test_${Date.now()}`; - await mem.add("Remember this fact", { userId }); - const before: SearchResult = await mem.getAll({ userId }); + await mem.add("Remember this fact", { filters: { user_id: userId } }); + const before: SearchResult = await mem.getAll({ + filters: { user_id: userId }, + }); expect(before.results.length).toBeGreaterThan(0); await mem.reset(); - const after: SearchResult = await mem.getAll({ userId }); + const after: SearchResult = await mem.getAll({ + filters: { user_id: userId }, + }); expect(after.results).toHaveLength(0); }); }); diff --git a/mem0-ts/src/oss/tests/vector-stores-compat.test.ts b/mem0-ts/src/oss/tests/vector-stores-compat.test.ts index 36fb64192..21bf9eb6e 100644 --- a/mem0-ts/src/oss/tests/vector-stores-compat.test.ts +++ b/mem0-ts/src/oss/tests/vector-stores-compat.test.ts @@ -944,7 +944,7 @@ describe("Memory class – backward compat with all providers", () => { disableHistory: true, }); - await mem.getAll({ userId: "u1" }); + await mem.getAll({ filters: { user_id: "u1" } }); const embedder = mockEmbedderFactory.create.mock.results[0].value; expect(embedder.embed).not.toHaveBeenCalledWith("dimension probe"); @@ -967,7 +967,7 @@ describe("Memory class – backward compat with all providers", () => { const mockEmbedder768 = createMockEmbedder(768); mockEmbedderFactory.create.mockReturnValue(mockEmbedder768); - await mem.getAll({ userId: "u1" }); + await mem.getAll({ filters: { user_id: "u1" } }); expect(mockEmbedder768.embed).not.toHaveBeenCalledWith("dimension probe"); }); @@ -982,7 +982,7 @@ describe("Memory class – backward compat with all providers", () => { disableHistory: true, }); - await mem.getAll({ userId: "u1" }); + await mem.getAll({ filters: { user_id: "u1" } }); expect(mockEmbedder768.embed).toHaveBeenCalledWith("dimension probe"); const vsCreateCall = mockVectorStoreFactory.create.mock.calls[0]; @@ -1003,7 +1003,7 @@ describe("Memory class – backward compat with all providers", () => { disableHistory: true, }); - await mem.getAll({ userId: "u1" }); + await mem.getAll({ filters: { user_id: "u1" } }); expect(mockVStore.initialize).toHaveBeenCalled(); }); @@ -1048,11 +1048,13 @@ describe("Memory class – backward compat with all providers", () => { }); // getAll - const all = await mem.getAll({ userId: "u1" }); + const all = await mem.getAll({ filters: { user_id: "u1" } }); expect(all).toBeDefined(); // search - const searchResult = await mem.search("query", { userId: "u1" }); + const searchResult = await mem.search("query", { + filters: { user_id: "u1" }, + }); expect(searchResult).toBeDefined(); // get @@ -1068,7 +1070,9 @@ describe("Memory class – backward compat with all providers", () => { expect(deleteResult.message).toBe("Memory deleted successfully!"); // deleteAll - const deleteAllResult = await mem.deleteAll({ userId: "u1" }); + const deleteAllResult = await mem.deleteAll({ + filters: { user_id: "u1" }, + }); expect(deleteAllResult.message).toBe("Memories deleted successfully!"); // history @@ -1093,7 +1097,7 @@ describe("Memory class – backward compat with all providers", () => { disableHistory: true, }); - await mem.getAll({ userId: "u1" }); + await mem.getAll({ filters: { user_id: "u1" } }); expect(mockVectorStoreFactory.create).toHaveBeenCalledTimes(1); await mem.reset(); @@ -1120,12 +1124,12 @@ describe("Memory class – backward compat with all providers", () => { disableHistory: true, }); - await expect(mem.getAll({ userId: "u1" })).rejects.toThrow( - "auto-detect embedding dimension", - ); - await expect(mem.search("q", { userId: "u1" })).rejects.toThrow( + await expect(mem.getAll({ filters: { user_id: "u1" } })).rejects.toThrow( "auto-detect embedding dimension", ); + await expect( + mem.search("q", { filters: { user_id: "u1" } }), + ).rejects.toThrow("auto-detect embedding dimension"); await expect(mem.get("id")).rejects.toThrow( "auto-detect embedding dimension", ); diff --git a/mem0/client/main.py b/mem0/client/main.py index 4426fdb0f..9502f4402 100644 --- a/mem0/client/main.py +++ b/mem0/client/main.py @@ -29,6 +29,9 @@ warnings.filterwarnings("default", category=DeprecationWarning) # Setup user config setup_config() +# Entity parameters that must be passed via filters, not top-level +ENTITY_PARAMS = frozenset({"user_id", "agent_id", "app_id", "run_id"}) + class MemoryClient: """Client for interacting with the Mem0 API. @@ -139,8 +142,7 @@ class MemoryClient: or a string. If a string is provided, it will be converted to a user message. options: Typed options for the add operation (AddMemoryOptions). - **kwargs: Additional parameters such as user_id, agent_id, app_id, - metadata, filters. + **kwargs: Additional parameters such as metadata, filters. Returns: A dictionary containing the API response in v1.1 format. @@ -163,7 +165,11 @@ class MemoryClient: raise ValueError(f"messages must be str, dict, or list[dict], got {type(messages).__name__}") kwargs = self._prepare_params(kwargs) + # Extract filters and spread entity IDs into payload (API expects top-level entity IDs) + filters = kwargs.pop("filters", None) payload = self._prepare_payload(messages, kwargs) + if filters: + payload.update(filters) response = self.client.post("/v3/memories/", json=payload) response.raise_for_status() if "metadata" in kwargs: @@ -201,8 +207,7 @@ class MemoryClient: Args: options: Typed options for the get_all operation (GetAllMemoryOptions). - **kwargs: Optional parameters for filtering (user_id, agent_id, - app_id, top_k, page, page_size). + **kwargs: Optional parameters for filtering (filters, page, page_size). Returns: A dictionary containing memories in v1.1 format: {"results": [...]} @@ -215,6 +220,14 @@ class MemoryClient: NetworkError: If network connectivity issues occur. MemoryNotFoundError: If the memory doesn't exist (for updates/deletes). """ + # Reject top-level entity params - must use filters instead + invalid_keys = ENTITY_PARAMS & set(kwargs.keys()) + if invalid_keys: + raise ValueError( + f"Top-level entity parameters {invalid_keys} are not supported in get_all(). " + f"Use filters={{'user_id': '...'}} instead." + ) + kwargs = {**(options.model_dump(exclude_unset=True) if options else {}), **kwargs} params = self._prepare_params(kwargs) @@ -264,6 +277,14 @@ class MemoryClient: NetworkError: If network connectivity issues occur. MemoryNotFoundError: If the memory doesn't exist (for updates/deletes). """ + # Reject top-level entity params - must use filters instead + invalid_keys = ENTITY_PARAMS & set(kwargs.keys()) + if invalid_keys: + raise ValueError( + f"Top-level entity parameters {invalid_keys} are not supported in search(). " + f"Use filters={{'user_id': '...'}} instead." + ) + kwargs = {**(options.model_dump(exclude_unset=True) if options else {}), **kwargs} params = self._prepare_params(kwargs) payload = {"query": query, **params} @@ -349,8 +370,7 @@ class MemoryClient: Args: options: Typed options for the delete_all operation (DeleteAllMemoryOptions). - **kwargs: Optional parameters for filtering (user_id, agent_id, - app_id). + **kwargs: Optional parameters for filtering (filters). Returns: A dictionary containing the API response. @@ -365,7 +385,16 @@ class MemoryClient: """ kwargs = {**(options.model_dump(exclude_unset=True) if options else {}), **kwargs} params = self._prepare_params(kwargs) - response = self.client.delete("/v1/memories/", params=params) + + # Extract filters and pass as query params (snake_case keys) + query_params = {} + if "filters" in params: + filters = params.pop("filters") + if isinstance(filters, dict): + query_params.update(filters) + query_params.update(params) + + response = self.client.delete("/v1/memories/", params=query_params) response.raise_for_status() capture_client_event( "client.delete_all", @@ -1052,8 +1081,7 @@ class AsyncMemoryClient: or a string. If a string is provided, it will be converted to a user message. options: Typed options for the add operation (AddMemoryOptions). - **kwargs: Additional parameters such as user_id, agent_id, app_id, - metadata, filters. + **kwargs: Additional parameters such as metadata, filters. Returns: A dictionary containing the API response in v1.1 format. @@ -1076,7 +1104,11 @@ class AsyncMemoryClient: raise ValueError(f"messages must be str, dict, or list[dict], got {type(messages).__name__}") kwargs = self._prepare_params(kwargs) + # Extract filters and spread entity IDs into payload (API expects top-level entity IDs) + filters = kwargs.pop("filters", None) payload = self._prepare_payload(messages, kwargs) + if filters: + payload.update(filters) response = await self.async_client.post("/v3/memories/", json=payload) response.raise_for_status() if "metadata" in kwargs: @@ -1111,6 +1143,14 @@ class AsyncMemoryClient: NetworkError: If network connectivity issues occur. MemoryNotFoundError: If the memory doesn't exist (for updates/deletes). """ + # Reject top-level entity params - must use filters instead + invalid_keys = ENTITY_PARAMS & set(kwargs.keys()) + if invalid_keys: + raise ValueError( + f"Top-level entity parameters {invalid_keys} are not supported in get_all(). " + f"Use filters={{'user_id': '...'}} instead." + ) + kwargs = {**(options.model_dump(exclude_unset=True) if options else {}), **kwargs} params = self._prepare_params(kwargs) @@ -1160,6 +1200,14 @@ class AsyncMemoryClient: NetworkError: If network connectivity issues occur. MemoryNotFoundError: If the memory doesn't exist (for updates/deletes). """ + # Reject top-level entity params - must use filters instead + invalid_keys = ENTITY_PARAMS & set(kwargs.keys()) + if invalid_keys: + raise ValueError( + f"Top-level entity parameters {invalid_keys} are not supported in search(). " + f"Use filters={{'user_id': '...'}} instead." + ) + kwargs = {**(options.model_dump(exclude_unset=True) if options else {}), **kwargs} params = self._prepare_params(kwargs) payload = {"query": query, **params} @@ -1245,7 +1293,7 @@ class AsyncMemoryClient: Args: options: Typed options for the delete_all operation (DeleteAllMemoryOptions). - **kwargs: Optional parameters for filtering (user_id, agent_id, app_id). + **kwargs: Optional parameters for filtering (filters). Returns: A dictionary containing the API response. @@ -1260,7 +1308,16 @@ class AsyncMemoryClient: """ kwargs = {**(options.model_dump(exclude_unset=True) if options else {}), **kwargs} params = self._prepare_params(kwargs) - response = await self.async_client.delete("/v1/memories/", params=params) + + # Extract filters and pass as query params (snake_case keys) + query_params = {} + if "filters" in params: + filters = params.pop("filters") + if isinstance(filters, dict): + query_params.update(filters) + query_params.update(params) + + response = await self.async_client.delete("/v1/memories/", params=query_params) response.raise_for_status() capture_client_event("client.delete_all", self, {"keys": list(kwargs.keys()), "sync_type": "async"}) return response.json() diff --git a/mem0/client/types.py b/mem0/client/types.py index 4053d69de..04d72c01d 100644 --- a/mem0/client/types.py +++ b/mem0/client/types.py @@ -2,6 +2,9 @@ These models provide IDE autocompletion, runtime validation, and type safety. Methods accept both typed options and **kwargs for backward compatibility. + +Identity fields (user_id, agent_id, app_id, run_id) must be passed inside +the ``filters`` dict — the v3 API does not accept them at the top level. """ from typing import Any, Dict, List, Optional, Union @@ -9,18 +12,16 @@ from typing import Any, Dict, List, Optional, Union from pydantic import BaseModel, Field -class EntityOptions(BaseModel): - """Identity options for add/delete operations (top-level entity IDs).""" +class AddMemoryOptions(BaseModel): + """Options for the add() method. - user_id: Optional[str] = Field(default=None, description="The user ID to associate with the memory") - agent_id: Optional[str] = Field(default=None, description="The agent ID to associate with the memory") - app_id: Optional[str] = Field(default=None, description="The app ID to associate with the memory") - run_id: Optional[str] = Field(default=None, description="The run ID to associate with the memory") - - -class AddMemoryOptions(EntityOptions): - """Options for the add() method.""" + Identity fields (user_id, agent_id, app_id, run_id) must be passed inside + the ``filters`` dict — the v3 API does not accept them at the top level. + """ + filters: Optional[Dict[str, Any]] = Field( + default=None, description="Filters containing entity IDs (e.g. {'user_id': '...'})" + ) metadata: Optional[Dict[str, Any]] = Field(default=None, description="Additional metadata for the memory") infer: Optional[bool] = Field(default=None, description="Whether to infer memories from the input") custom_categories: Optional[List[Dict[str, Any]]] = Field( @@ -72,10 +73,16 @@ class GetAllMemoryOptions(BaseModel): categories: Optional[List[str]] = Field(default=None, description="Categories to filter by") -class DeleteAllMemoryOptions(EntityOptions): - """Options for the delete_all() method.""" +class DeleteAllMemoryOptions(BaseModel): + """Options for the delete_all() method. - pass + Identity fields (user_id, agent_id, app_id, run_id) must be passed inside + the ``filters`` dict — the API does not accept them at the top level. + """ + + filters: Optional[Dict[str, Any]] = Field( + default=None, description="Filters containing entity IDs (e.g. {'user_id': '...'})" + ) class UpdateMemoryOptions(BaseModel): diff --git a/mem0/memory/main.py b/mem0/memory/main.py index 1dd5ac4da..2e66d5001 100644 --- a/mem0/memory/main.py +++ b/mem0/memory/main.py @@ -10,7 +10,6 @@ from copy import deepcopy from datetime import datetime, timezone from typing import Any, Dict, Optional - from pydantic import ValidationError from mem0.configs.base import MemoryConfig, MemoryItem @@ -18,12 +17,9 @@ from mem0.configs.enums import MemoryType from mem0.configs.prompts import ( ADDITIVE_EXTRACTION_PROMPT, AGENT_CONTEXT_SUFFIX, - generate_additive_extraction_prompt, PROCEDURAL_MEMORY_SYSTEM_PROMPT, + generate_additive_extraction_prompt, ) -from mem0.utils.lemmatization import lemmatize_for_bm25 -from mem0.utils.entity_extraction import extract_entities, extract_entities_batch -from mem0.utils.scoring import ENTITY_BOOST_WEIGHT, get_bm25_params, normalize_bm25, score_and_rank from mem0.exceptions import ValidationError as Mem0ValidationError from mem0.memory.base import MemoryBase from mem0.memory.setup import mem0_dir, setup_config @@ -36,11 +32,19 @@ from mem0.memory.utils import ( process_telemetry_filters, remove_code_blocks, ) +from mem0.utils.entity_extraction import extract_entities, extract_entities_batch from mem0.utils.factory import ( EmbedderFactory, LlmFactory, - VectorStoreFactory, RerankerFactory, + VectorStoreFactory, +) +from mem0.utils.lemmatization import lemmatize_for_bm25 +from mem0.utils.scoring import ( + ENTITY_BOOST_WEIGHT, + get_bm25_params, + normalize_bm25, + score_and_rank, ) # Suppress SWIG deprecation warnings globally @@ -92,6 +96,19 @@ _SENSITIVE_SUFFIXES = ( "_credentials", ) +# Entity parameters that must be passed via filters, not top-level kwargs +ENTITY_PARAMS = frozenset({"user_id", "agent_id", "run_id"}) + + +def _reject_top_level_entity_params(kwargs: Dict[str, Any], method_name: str) -> None: + """Reject top-level entity parameters - must use filters instead.""" + invalid_keys = ENTITY_PARAMS & set(kwargs.keys()) + if invalid_keys: + raise ValueError( + f"Top-level entity parameters {invalid_keys} are not supported in {method_name}(). " + f"Use filters={{'user_id': '...'}} instead." + ) + def _is_sensitive_field(field_name: str) -> bool: """Check if a field should be redacted for telemetry safety. @@ -408,26 +425,27 @@ class Memory(MemoryBase): self, messages, *, - user_id: Optional[str] = None, - agent_id: Optional[str] = None, - run_id: Optional[str] = None, + filters: Optional[Dict[str, Any]] = None, metadata: Optional[Dict[str, Any]] = None, infer: bool = True, memory_type: Optional[str] = None, prompt: Optional[str] = None, + **kwargs, ): """ Create a new memory. - Adds new memories scoped to a single session id (e.g. `user_id`, `agent_id`, or `run_id`). One of those ids is required. + Adds new memories scoped to a single session id (e.g. `user_id`, `agent_id`, or `run_id`). One of those ids is required in filters. Args: messages (str or List[Dict[str, str]]): The message content or list of messages (e.g., `[{"role": "user", "content": "Hello"}, {"role": "assistant", "content": "Hi"}]`) to be processed and stored. - user_id (str, optional): ID of the user creating the memory. Defaults to None. - agent_id (str, optional): ID of the agent creating the memory. Defaults to None. - run_id (str, optional): ID of the run creating the memory. Defaults to None. + filters (dict, optional): Dictionary containing entity identifiers. Must contain at least one of: + - user_id: ID of the user creating the memory + - agent_id: ID of the agent creating the memory + - run_id: ID of the run creating the memory + Example: `{"user_id": "user123"}` or `{"agent_id": "agent456"}` metadata (dict, optional): Metadata to store with the memory. Defaults to None. infer (bool, optional): If True (default), an LLM is used to extract key facts from 'messages' and decide whether to add, update, or delete related memories. @@ -435,7 +453,7 @@ class Memory(MemoryBase): memory_type (str, optional): Specifies the type of memory. Currently, only `MemoryType.PROCEDURAL.value` ("procedural_memory") is explicitly handled for creating procedural memories (typically requires 'agent_id'). Otherwise, memories - are treated as general conversational/factual memories.memory_type (str, optional): Type of memory to create. Defaults to None. By default, it creates the short term memories and long term (semantic and episodic) memories. Pass "procedural_memory" to create procedural memories. + are treated as general conversational/factual memories. prompt (str, optional): Prompt to use for the memory creation. Defaults to None. @@ -451,11 +469,19 @@ class Memory(MemoryBase): LLMError: If LLM operations fail. DatabaseError: If database operations fail. """ + filters = filters or {} + + # Validate filters contains at least one entity ID + if not any(key in 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. " + "Example: filters={'user_id': 'u1'}" + ) processed_metadata, effective_filters = _build_filters_and_metadata( - user_id=user_id, - agent_id=agent_id, - run_id=run_id, + user_id=filters.get("user_id"), + agent_id=filters.get("agent_id"), + run_id=filters.get("run_id"), input_metadata=metadata, ) @@ -481,7 +507,7 @@ class Memory(MemoryBase): suggestion="Convert your input to a string, dictionary, or list of dictionaries." ) - if agent_id is not None and memory_type == MemoryType.PROCEDURAL.value: + if filters.get("agent_id") is not None and memory_type == MemoryType.PROCEDURAL.value: results = self._create_procedural_memory(messages, metadata=processed_metadata, prompt=prompt) return results @@ -850,38 +876,39 @@ class Memory(MemoryBase): def get_all( self, *, - user_id: Optional[str] = None, - agent_id: Optional[str] = None, - run_id: Optional[str] = None, filters: Optional[Dict[str, Any]] = None, top_k: int = 100, + **kwargs, ): """ List all memories. Args: - user_id (str, optional): user id - agent_id (str, optional): agent id - run_id (str, optional): run id - filters (dict, optional): Additional custom key-value filters to apply to the search. - These are merged with the ID-based scoping filters. For example, - `filters={"actor_id": "some_user"}`. - limit (int, optional): The maximum number of memories to return. Defaults to 100. + 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. 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. """ + # Reject top-level entity params - must use filters instead + _reject_top_level_entity_params(kwargs, "get_all") + + # 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. " + "Example: filters={'user_id': 'u1'}" + ) limit = top_k - _, effective_filters = _build_filters_and_metadata( - user_id=user_id, agent_id=agent_id, run_id=run_id, input_filters=filters - ) - - if not any(key in effective_filters for key in ("user_id", "agent_id", "run_id")): - raise ValueError("At least one of 'user_id', 'agent_id', or 'run_id' must be specified.") - keys, encoded_ids = process_telemetry_filters(effective_filters) capture_event( "mem0.get_all", self, {"limit": limit, "keys": keys, "encoded_ids": encoded_ids, "sync_type": "sync"} @@ -942,28 +969,26 @@ class Memory(MemoryBase): self, query: str, *, - user_id: Optional[str] = None, - agent_id: Optional[str] = None, - run_id: Optional[str] = None, top_k: int = 100, filters: Optional[Dict[str, Any]] = None, threshold: float = 0.1, rerank: bool = False, + **kwargs, ): """ - Searches for memories based on a query + Searches for memories based on a query. + Args: query (str): Query to search for. - user_id (str, optional): ID of the user to search for. Defaults to None. - agent_id (str, optional): ID of the agent to search for. Defaults to None. - run_id (str, optional): ID of the run to search for. Defaults to None. - limit (int, optional): Limit the number of results. Defaults to 100. - filters (dict, optional): Legacy filters to apply to the search. Defaults to None. - threshold (float, optional): Minimum score for a memory to be included in the results. Defaults to 0.1. - filters (dict, optional): Enhanced metadata filtering with operators: + top_k (int, optional): Maximum number of results to return. Defaults to 100. + 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"} + + Enhanced metadata filtering with operators: - {"key": "value"} - exact match - {"key": {"eq": "value"}} - equals - - {"key": {"ne": "value"}} - not equals + - {"key": {"ne": "value"}} - not equals - {"key": {"in": ["val1", "val2"]}} - in list - {"key": {"nin": ["val1", "val2"]}} - not in list - {"key": {"gt": 10}} - greater than @@ -976,34 +1001,39 @@ class Memory(MemoryBase): - {"AND": [filter1, filter2]} - logical AND - {"OR": [filter1, filter2]} - logical OR - {"NOT": [filter1]} - logical NOT + threshold (float, optional): Minimum score for a memory to be included. Defaults to 0.1. + rerank (bool, optional): Whether to rerank results. Defaults to False. Returns: dict: A dictionary containing the search results under a "results" key. Example for v1.1+: `{"results": [{"id": "...", "memory": "...", "score": 0.8, ...}]}` - """ - limit = top_k - - _, effective_filters = _build_filters_and_metadata( - user_id=user_id, agent_id=agent_id, run_id=run_id, input_filters=filters - ) + Raises: + ValueError: If filters doesn't contain at least one of user_id, agent_id, run_id. + """ + # Reject top-level entity params - must use filters instead + _reject_top_level_entity_params(kwargs, "search") + + # Validate filters contains at least one entity ID + effective_filters = filters.copy() if filters else {} if not any(key in effective_filters for key in ("user_id", "agent_id", "run_id")): - raise ValueError("At least one of 'user_id', 'agent_id', or 'run_id' must be specified.") + raise ValueError( + "filters must contain at least one of: user_id, agent_id, run_id. " + "Example: filters={'user_id': 'u1'}" + ) + + limit = top_k # Apply enhanced metadata filtering if advanced operators are detected - if filters and self._has_advanced_operators(filters): - processed_filters = self._process_metadata_filters(filters) - # Remove original logical/operator keys that _build_filters_and_metadata - # copied verbatim from input_filters — they have now been reprocessed. + if self._has_advanced_operators(effective_filters): + processed_filters = self._process_metadata_filters(effective_filters) + # Remove logical/operator keys that have been reprocessed for logical_key in ("AND", "OR", "NOT"): effective_filters.pop(logical_key, None) - for fk in list(filters.keys()): - if fk not in ("AND", "OR", "NOT") and fk in effective_filters and isinstance(filters[fk], dict): + for fk in list(effective_filters.keys()): + if fk not in ("AND", "OR", "NOT", "user_id", "agent_id", "run_id") and isinstance(effective_filters.get(fk), dict): effective_filters.pop(fk, None) effective_filters.update(processed_filters) - elif filters: - # Simple filters, merge directly - effective_filters.update(filters) keys, encoded_ids = process_telemetry_filters(effective_filters) capture_event( @@ -1330,26 +1360,23 @@ class Memory(MemoryBase): self._delete_memory(memory_id, existing_memory) return {"message": "Memory deleted successfully!"} - def delete_all(self, user_id: Optional[str] = None, agent_id: Optional[str] = None, run_id: Optional[str] = None): + def delete_all(self, *, filters: Optional[Dict[str, Any]] = None, **kwargs): """ Delete all memories. Args: - user_id (str, optional): ID of the user to delete memories for. Defaults to None. - agent_id (str, optional): ID of the agent to delete memories for. Defaults to None. - run_id (str, optional): ID of the run to delete memories for. Defaults to None. + filters (dict, optional): Dictionary containing entity identifiers. Must contain at least one of: + - user_id: ID of the user to delete memories for + - agent_id: ID of the agent to delete memories for + - run_id: ID of the run to delete memories for + Example: `{"user_id": "user123"}` """ - filters: Dict[str, Any] = {} - if user_id: - filters["user_id"] = user_id - if agent_id: - filters["agent_id"] = agent_id - if run_id: - filters["run_id"] = run_id + filters = filters or {} - if not filters: + if not any(key in filters for key in ("user_id", "agent_id", "run_id")): raise ValueError( - "At least one filter is required to delete all memories. If you want to delete all memories, use the `reset()` method." + "filters must contain at least one of: user_id, agent_id, run_id. " + "If you want to delete all memories, use the `reset()` method." ) keys, encoded_ids = process_telemetry_filters(filters) @@ -1667,23 +1694,24 @@ class AsyncMemory(MemoryBase): self, messages, *, - user_id: Optional[str] = None, - agent_id: Optional[str] = None, - run_id: Optional[str] = None, + filters: Optional[Dict[str, Any]] = None, metadata: Optional[Dict[str, Any]] = None, infer: bool = True, memory_type: Optional[str] = None, prompt: Optional[str] = None, llm=None, + **kwargs, ): """ Create a new memory asynchronously. Args: messages (str or List[Dict[str, str]]): Messages to store in the memory. - user_id (str, optional): ID of the user creating the memory. - agent_id (str, optional): ID of the agent creating the memory. Defaults to None. - run_id (str, optional): ID of the run creating the memory. Defaults to None. + filters (dict, optional): Dictionary containing entity identifiers. Must contain at least one of: + - user_id: ID of the user creating the memory + - agent_id: ID of the agent creating the memory + - run_id: ID of the run creating the memory + Example: `{"user_id": "user123"}` or `{"agent_id": "agent456"}` metadata (dict, optional): Metadata to store with the memory. Defaults to None. infer (bool, optional): Whether to infer the memories. Defaults to True. memory_type (str, optional): Type of memory to create. Defaults to None. @@ -1693,8 +1721,20 @@ class AsyncMemory(MemoryBase): Returns: dict: A dictionary containing the result of the memory addition operation. """ + filters = filters or {} + + # Validate filters contains at least one entity ID + if not any(key in 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. " + "Example: filters={'user_id': 'u1'}" + ) + processed_metadata, effective_filters = _build_filters_and_metadata( - user_id=user_id, agent_id=agent_id, run_id=run_id, input_metadata=metadata + user_id=filters.get("user_id"), + agent_id=filters.get("agent_id"), + run_id=filters.get("run_id"), + input_metadata=metadata, ) if memory_type is not None and memory_type != MemoryType.PROCEDURAL.value: @@ -1716,7 +1756,7 @@ class AsyncMemory(MemoryBase): suggestion="Convert your input to a string, dictionary, or list of dictionaries." ) - if agent_id is not None and memory_type == MemoryType.PROCEDURAL.value: + if filters.get("agent_id") is not None and memory_type == MemoryType.PROCEDURAL.value: results = await self._create_procedural_memory( messages, metadata=processed_metadata, prompt=prompt, llm=llm ) @@ -2093,41 +2133,39 @@ class AsyncMemory(MemoryBase): async def get_all( self, *, - user_id: Optional[str] = None, - agent_id: Optional[str] = None, - run_id: Optional[str] = None, filters: Optional[Dict[str, Any]] = None, top_k: int = 100, + **kwargs, ): """ List all memories. - Args: - user_id (str, optional): user id - agent_id (str, optional): agent id - run_id (str, optional): run id - filters (dict, optional): Additional custom key-value filters to apply to the search. - These are merged with the ID-based scoping filters. For example, - `filters={"actor_id": "some_user"}`. - limit (int, optional): The maximum number of memories to return. Defaults to 100. + Args: + 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. - Returns: - dict: A dictionary containing a list of memories under the "results" key. - Example for v1.1+: `{"results": [{"id": "...", "memory": "...", ...}]}` + 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. """ + # Reject top-level entity params - must use filters instead + _reject_top_level_entity_params(kwargs, "get_all") - limit = top_k - - _, effective_filters = _build_filters_and_metadata( - user_id=user_id, agent_id=agent_id, run_id=run_id, input_filters=filters - ) - + # 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( - "When 'conversation_id' is not provided (classic mode), " - "at least one of 'user_id', 'agent_id', or 'run_id' must be specified for get_all." + "filters must contain at least one of: user_id, agent_id, run_id. " + "Example: filters={'user_id': 'u1'}" ) + limit = top_k + keys, encoded_ids = process_telemetry_filters(effective_filters) capture_event( "mem0.get_all", self, {"limit": limit, "keys": keys, "encoded_ids": encoded_ids, "sync_type": "async"} @@ -2188,29 +2226,26 @@ class AsyncMemory(MemoryBase): self, query: str, *, - user_id: Optional[str] = None, - agent_id: Optional[str] = None, - run_id: Optional[str] = None, top_k: int = 100, filters: Optional[Dict[str, Any]] = None, threshold: float = 0.1, - metadata_filters: Optional[Dict[str, Any]] = None, rerank: bool = False, + **kwargs, ): """ - Searches for memories based on a query + Searches for memories based on a query. + Args: query (str): Query to search for. - user_id (str, optional): ID of the user to search for. Defaults to None. - agent_id (str, optional): ID of the agent to search for. Defaults to None. - run_id (str, optional): ID of the run to search for. Defaults to None. - limit (int, optional): Limit the number of results. Defaults to 100. - filters (dict, optional): Legacy filters to apply to the search. Defaults to None. - threshold (float, optional): Minimum score for a memory to be included in the results. Defaults to None. - filters (dict, optional): Enhanced metadata filtering with operators: + top_k (int, optional): Maximum number of results to return. Defaults to 100. + 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"} + + Enhanced metadata filtering with operators: - {"key": "value"} - exact match - {"key": {"eq": "value"}} - equals - - {"key": {"ne": "value"}} - not equals + - {"key": {"ne": "value"}} - not equals - {"key": {"in": ["val1", "val2"]}} - in list - {"key": {"nin": ["val1", "val2"]}} - not in list - {"key": {"gt": 10}} - greater than @@ -2223,35 +2258,39 @@ class AsyncMemory(MemoryBase): - {"AND": [filter1, filter2]} - logical AND - {"OR": [filter1, filter2]} - logical OR - {"NOT": [filter1]} - logical NOT + threshold (float, optional): Minimum score for a memory to be included. Defaults to 0.1. + rerank (bool, optional): Whether to rerank results. Defaults to False. Returns: dict: A dictionary containing the search results under a "results" key. 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. """ + # Reject top-level entity params - must use filters instead + _reject_top_level_entity_params(kwargs, "search") + + # Validate filters contains at least one entity ID + effective_filters = filters.copy() if filters else {} + 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. " + "Example: filters={'user_id': 'u1'}" + ) limit = top_k - _, effective_filters = _build_filters_and_metadata( - user_id=user_id, agent_id=agent_id, run_id=run_id, input_filters=filters - ) - - if not any(key in effective_filters for key in ("user_id", "agent_id", "run_id")): - raise ValueError("at least one of 'user_id', 'agent_id', or 'run_id' must be specified ") - # Apply enhanced metadata filtering if advanced operators are detected - if filters and self._has_advanced_operators(filters): - processed_filters = self._process_metadata_filters(filters) - # Remove original logical/operator keys that _build_filters_and_metadata - # copied verbatim from input_filters — they have now been reprocessed. + if self._has_advanced_operators(effective_filters): + processed_filters = self._process_metadata_filters(effective_filters) + # Remove logical/operator keys that have been reprocessed for logical_key in ("AND", "OR", "NOT"): effective_filters.pop(logical_key, None) - for fk in list(filters.keys()): - if fk not in ("AND", "OR", "NOT") and fk in effective_filters and isinstance(filters[fk], dict): + for fk in list(effective_filters.keys()): + if fk not in ("AND", "OR", "NOT", "user_id", "agent_id", "run_id") and isinstance(effective_filters.get(fk), dict): effective_filters.pop(fk, None) effective_filters.update(processed_filters) - elif filters: - # Simple filters, merge directly - effective_filters.update(filters) keys, encoded_ids = process_telemetry_filters(effective_filters) capture_event( @@ -2567,26 +2606,23 @@ class AsyncMemory(MemoryBase): await self._delete_memory(memory_id, existing_memory) return {"message": "Memory deleted successfully!"} - async def delete_all(self, user_id=None, agent_id=None, run_id=None): + async def delete_all(self, *, filters: Optional[Dict[str, Any]] = None, **kwargs): """ Delete all memories asynchronously. Args: - user_id (str, optional): ID of the user to delete memories for. Defaults to None. - agent_id (str, optional): ID of the agent to delete memories for. Defaults to None. - run_id (str, optional): ID of the run to delete memories for. Defaults to None. + filters (dict, optional): Dictionary containing entity identifiers. Must contain at least one of: + - user_id: ID of the user to delete memories for + - agent_id: ID of the agent to delete memories for + - run_id: ID of the run to delete memories for + Example: `{"user_id": "user123"}` """ - filters = {} - if user_id: - filters["user_id"] = user_id - if agent_id: - filters["agent_id"] = agent_id - if run_id: - filters["run_id"] = run_id + filters = filters or {} - if not filters: + if not any(key in filters for key in ("user_id", "agent_id", "run_id")): raise ValueError( - "At least one filter is required to delete all memories. If you want to delete all memories, use the `reset()` method." + "filters must contain at least one of: user_id, agent_id, run_id. " + "If you want to delete all memories, use the `reset()` method." ) keys, encoded_ids = process_telemetry_filters(filters) diff --git a/tests/test_client.py b/tests/test_client.py new file mode 100644 index 000000000..fb685c02a --- /dev/null +++ b/tests/test_client.py @@ -0,0 +1,149 @@ +"""Tests for MemoryClient entity parameter rejection.""" + +from unittest.mock import MagicMock, patch + +import pytest + + +@pytest.fixture +def mock_memory_client(): + """Create a mock MemoryClient for testing entity param rejection.""" + with patch("mem0.client.main.httpx.Client") as mock_httpx: + # Create a mock client instance + mock_http_client = MagicMock() + mock_http_client.get.return_value = MagicMock( + json=lambda: {"org_id": "org1", "project_id": "proj1", "user_email": "test@test.com"}, + raise_for_status=lambda: None + ) + mock_httpx.return_value = mock_http_client + + with patch("mem0.client.main.capture_client_event"): + from mem0.client.main import MemoryClient + client = MemoryClient(api_key="test-api-key") + yield client + + +class TestSearchEntityParamRejection: + """Tests that top-level entity params are rejected in search().""" + + def test_search_rejects_user_id_kwarg(self, mock_memory_client): + """search() should reject user_id as top-level kwarg.""" + with pytest.raises(ValueError, match=r"user_id"): + mock_memory_client.search("test query", user_id="u1") + + def test_search_rejects_agent_id_kwarg(self, mock_memory_client): + """search() should reject agent_id as top-level kwarg.""" + with pytest.raises(ValueError, match=r"agent_id"): + mock_memory_client.search("test query", agent_id="a1") + + def test_search_rejects_app_id_kwarg(self, mock_memory_client): + """search() should reject app_id as top-level kwarg.""" + with pytest.raises(ValueError, match=r"app_id"): + mock_memory_client.search("test query", app_id="app1") + + def test_search_rejects_run_id_kwarg(self, mock_memory_client): + """search() should reject run_id as top-level kwarg.""" + with pytest.raises(ValueError, match=r"run_id"): + mock_memory_client.search("test query", run_id="r1") + + def test_search_rejects_multiple_entity_params(self, mock_memory_client): + """search() should reject multiple top-level entity params.""" + with pytest.raises(ValueError, match=r"user_id|agent_id"): + mock_memory_client.search("test query", user_id="u1", agent_id="a1") + + +class TestGetAllEntityParamRejection: + """Tests that top-level entity params are rejected in get_all().""" + + def test_get_all_rejects_user_id_kwarg(self, mock_memory_client): + """get_all() should reject user_id as top-level kwarg.""" + with pytest.raises(ValueError, match=r"user_id"): + mock_memory_client.get_all(user_id="u1") + + def test_get_all_rejects_agent_id_kwarg(self, mock_memory_client): + """get_all() should reject agent_id as top-level kwarg.""" + with pytest.raises(ValueError, match=r"agent_id"): + mock_memory_client.get_all(agent_id="a1") + + def test_get_all_rejects_app_id_kwarg(self, mock_memory_client): + """get_all() should reject app_id as top-level kwarg.""" + with pytest.raises(ValueError, match=r"app_id"): + mock_memory_client.get_all(app_id="app1") + + def test_get_all_rejects_run_id_kwarg(self, mock_memory_client): + """get_all() should reject run_id as top-level kwarg.""" + with pytest.raises(ValueError, match=r"run_id"): + mock_memory_client.get_all(run_id="r1") + + +class TestFilterOperatorPassthrough: + """Tests that AND/OR/NOT filter operators are passed through to the API.""" + + def test_search_passes_and_filters(self, mock_memory_client): + """search() should pass AND filters to the API.""" + mock_memory_client.client.post.return_value = MagicMock( + json=lambda: {"results": []}, + raise_for_status=lambda: None + ) + + mock_memory_client.search( + "test query", + filters={"AND": [{"user_id": "u1"}, {"created_at": {"gte": "2024-01-01"}}]} + ) + + # Verify the POST was called with filters intact + call_args = mock_memory_client.client.post.call_args + payload = call_args.kwargs.get("json", call_args.args[1] if len(call_args.args) > 1 else {}) + assert payload["filters"] == {"AND": [{"user_id": "u1"}, {"created_at": {"gte": "2024-01-01"}}]} + + def test_search_passes_or_filters(self, mock_memory_client): + """search() should pass OR filters to the API.""" + mock_memory_client.client.post.return_value = MagicMock( + json=lambda: {"results": []}, + raise_for_status=lambda: None + ) + + mock_memory_client.search( + "test query", + filters={"OR": [{"user_id": "u1"}, {"agent_id": "a1"}]} + ) + + call_args = mock_memory_client.client.post.call_args + payload = call_args.kwargs.get("json", call_args.args[1] if len(call_args.args) > 1 else {}) + assert payload["filters"] == {"OR": [{"user_id": "u1"}, {"agent_id": "a1"}]} + + def test_search_passes_not_filters(self, mock_memory_client): + """search() should pass NOT filters to the API.""" + mock_memory_client.client.post.return_value = MagicMock( + json=lambda: {"results": []}, + raise_for_status=lambda: None + ) + + mock_memory_client.search( + "test query", + filters={"AND": [{"user_id": "u1"}, {"NOT": {"categories": {"in": ["spam"]}}}]} + ) + + call_args = mock_memory_client.client.post.call_args + payload = call_args.kwargs.get("json", call_args.args[1] if len(call_args.args) > 1 else {}) + assert payload["filters"] == {"AND": [{"user_id": "u1"}, {"NOT": {"categories": {"in": ["spam"]}}}]} + + def test_search_passes_complex_nested_filters(self, mock_memory_client): + """search() should pass complex nested AND/OR/NOT filters to the API.""" + mock_memory_client.client.post.return_value = MagicMock( + json=lambda: {"results": []}, + raise_for_status=lambda: None + ) + + complex_filter = { + "AND": [ + {"user_id": "u1"}, + {"created_at": {"gte": "2024-01-01"}}, + {"NOT": {"OR": [{"categories": {"in": ["spam"]}}, {"categories": {"in": ["test"]}}]}} + ] + } + mock_memory_client.search("test query", filters=complex_filter) + + call_args = mock_memory_client.client.post.call_args + payload = call_args.kwargs.get("json", call_args.args[1] if len(call_args.args) > 1 else {}) + assert payload["filters"] == complex_filter diff --git a/tests/test_main.py b/tests/test_main.py index 06fa8d9a4..2fc49fd93 100644 --- a/tests/test_main.py +++ b/tests/test_main.py @@ -57,7 +57,10 @@ def test_add(memory_instance, version): memory_instance.config.version = version memory_instance._add_to_vector_store = Mock(return_value=[{"memory": "Test memory", "event": "ADD"}]) - result = memory_instance.add(messages=[{"role": "user", "content": "Test message"}], user_id="test_user") + result = memory_instance.add( + messages=[{"role": "user", "content": "Test message"}], + filters={"user_id": "test_user"}, + ) assert "results" in result assert result["results"] == [{"memory": "Test memory", "event": "ADD"}] @@ -105,7 +108,7 @@ def test_search(memory_instance, version): with patch("mem0.memory.main.lemmatize_for_bm25", return_value="test query"), \ patch("mem0.memory.main.extract_entities", return_value=[]): - result = memory_instance.search("test query", user_id="test_user") + result = memory_instance.search("test query", filters={"user_id": "test_user"}) assert "results" in result assert len(result["results"]) == 2 @@ -185,7 +188,7 @@ def test_delete_all(memory_instance, version): memory_instance.vector_store.reset = Mock() memory_instance._delete_memory = Mock() - result = memory_instance.delete_all(user_id="test_user") + result = memory_instance.delete_all(filters={"user_id": "test_user"}) assert memory_instance._delete_memory.call_count == 2 # Ensure the collection is NOT dropped — only matched memories should be removed @@ -206,7 +209,7 @@ def test_get_all(memory_instance, version, expected_result): mock_memories = [Mock(id="1", payload={"data": "Memory 1", "user_id": "test_user"})] memory_instance.vector_store.list = Mock(return_value=(mock_memories, None)) - result = memory_instance.get_all(user_id="test_user") + result = memory_instance.get_all(filters={"user_id": "test_user"}) assert isinstance(result, dict) assert "results" in result diff --git a/tests/test_memory.py b/tests/test_memory.py index 9b7e8bbfd..198fcf449 100644 --- a/tests/test_memory.py +++ b/tests/test_memory.py @@ -33,20 +33,20 @@ def memory_client(): def test_create_memory(memory_client): data = "Name is John Doe." - result = memory_client.add([{"role": "user", "content": data}], user_id="test_user") + result = memory_client.add([{"role": "user", "content": data}], filters={"user_id": "test_user"}) assert result["results"][0]["memory"] == data def test_get_memory(memory_client): data = "Name is John Doe." - memory_client.add([{"role": "user", "content": data}], user_id="test_user") + memory_client.add([{"role": "user", "content": data}], filters={"user_id": "test_user"}) result = memory_client.get("1") assert result["memory"] == data def test_update_memory(memory_client): data = "Name is John Doe." - memory_client.add([{"role": "user", "content": data}], user_id="test_user") + memory_client.add([{"role": "user", "content": data}], filters={"user_id": "test_user"}) new_data = "Name is John Kapoor." update_result = memory_client.update("1", text=new_data) assert update_result["message"] == "Memory updated successfully!" @@ -54,14 +54,14 @@ def test_update_memory(memory_client): def test_delete_memory(memory_client): data = "Name is John Doe." - memory_client.add([{"role": "user", "content": data}], user_id="test_user") + memory_client.add([{"role": "user", "content": data}], filters={"user_id": "test_user"}) delete_result = memory_client.delete("1") assert delete_result["message"] == "Memory deleted successfully!" def test_history(memory_client): data = "I like Indian food." - memory_client.add([{"role": "user", "content": data}], user_id="test_user") + memory_client.add([{"role": "user", "content": data}], filters={"user_id": "test_user"}) memory_client.update("1", text="I like Italian food.") history = memory_client.history("1") assert history[0]["memory"] == "I like Indian food." @@ -71,9 +71,9 @@ def test_history(memory_client): def test_list_memories(memory_client): data1 = "Name is John Doe." data2 = "Name is John Doe. I like to code in Python." - memory_client.add([{"role": "user", "content": data1}], user_id="test_user") - memory_client.add([{"role": "user", "content": data2}], user_id="test_user") - memories = memory_client.get_all(user_id="test_user") + memory_client.add([{"role": "user", "content": data1}], filters={"user_id": "test_user"}) + memory_client.add([{"role": "user", "content": data2}], filters={"user_id": "test_user"}) + memories = memory_client.get_all(filters={"user_id": "test_user"}) assert data1 in memories assert data2 in memories @@ -469,7 +469,7 @@ def test_add_infer_false_embeds_once(mock_sqlite, mock_llm_factory, mock_vector_ from mem0.memory.main import Memory as MemoryClass memory = MemoryClass(MemoryConfig()) - memory.add("foo", user_id="test_user", infer=False) + memory.add("foo", filters={"user_id": "test_user"}, infer=False) assert embedder.embed.call_count == 1 mock_vector_store.insert.assert_called_once() @@ -511,7 +511,7 @@ def test_add_infer_true_caches_embedding_on_llm_rewrite(mock_sqlite, mock_llm_fa from mem0.memory.main import Memory as MemoryClass memory = MemoryClass(MemoryConfig()) - memory.add("I like Python", user_id="test_user", infer=True) + memory.add("I like Python", filters={"user_id": "test_user"}, infer=True) # V3 pipeline: embed called once for search query (Phase 1), # embed_batch called once for extracted memories (Phase 3) @@ -568,7 +568,7 @@ def test_update_infer_true_caches_embedding_on_llm_rewrite(mock_sqlite, mock_llm from mem0.memory.main import Memory as MemoryClass memory = MemoryClass(MemoryConfig()) - memory.add("I love Python now", user_id="test_user", infer=True) + memory.add("I love Python now", filters={"user_id": "test_user"}, infer=True) # V3 pipeline: embed called once for search query (Phase 1), # embed_batch called once for extracted memories (Phase 3) @@ -785,3 +785,40 @@ def test_reset_skips_graph_when_graph_disabled(mock_sqlite, mock_llm_factory, mo # graph should remain None after reset assert memory.graph is None + + +# ─── Entity Param Rejection Tests ───────────────────────────────────────────── +@patch('mem0.utils.factory.EmbedderFactory.create') +@patch('mem0.utils.factory.VectorStoreFactory.create') +@patch('mem0.utils.factory.LlmFactory.create') +@patch('mem0.memory.storage.SQLiteManager') +def test_search_rejects_user_id_kwarg(mock_sqlite, mock_llm_factory, mock_vector_factory, mock_embedder_factory): + """search() should reject user_id as top-level kwarg.""" + mock_embedder_factory.return_value = MagicMock() + mock_vector_factory.return_value = MagicMock() + mock_llm_factory.return_value = MagicMock() + mock_sqlite.return_value = MagicMock() + + config = MemoryConfig() + memory = Memory(config) + + with pytest.raises(ValueError, match=r"user_id.*filters"): + memory.search("test query", user_id="u1") + + +@patch('mem0.utils.factory.EmbedderFactory.create') +@patch('mem0.utils.factory.VectorStoreFactory.create') +@patch('mem0.utils.factory.LlmFactory.create') +@patch('mem0.memory.storage.SQLiteManager') +def test_get_all_rejects_user_id_kwarg(mock_sqlite, mock_llm_factory, mock_vector_factory, mock_embedder_factory): + """get_all() should reject user_id as top-level kwarg.""" + mock_embedder_factory.return_value = MagicMock() + mock_vector_factory.return_value = MagicMock() + mock_llm_factory.return_value = MagicMock() + mock_sqlite.return_value = MagicMock() + + config = MemoryConfig() + memory = Memory(config) + + with pytest.raises(ValueError, match=r"user_id.*filters"): + memory.get_all(user_id="u1")