refactor: update MemoryClient to use filters for entity parameters

This commit is contained in:
kartik-mem0
2026-04-14 16:17:50 +05:30
parent 133625bbc6
commit 28cc8fcee4
29 changed files with 1440 additions and 453 deletions
-1
View File
@@ -3,7 +3,6 @@ import type * as MemoryTypes from "./mem0.types";
// Re-export all types from mem0.types
export type {
EntityOptions,
AddMemoryOptions,
SearchMemoryOptions,
GetAllMemoryOptions,
+56 -4
View File
@@ -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<string, any> | 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<Array<Memory>> {
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<string, any> = {
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<Array<Memory>> {
// Reject top-level entity params - must use filters instead
rejectTopLevelEntityParams(options as Record<string, any>, "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<Memory> }> {
// Reject top-level entity params - must use filters instead
rejectTopLevelEntityParams(options as Record<string, any>, "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<string, string> = {};
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}`,
{
+5 -10
View File
@@ -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<string, any>;
metadata?: Record<string, any>;
infer?: boolean;
customCategories?: custom_categories[];
@@ -35,7 +28,9 @@ export interface GetAllMemoryOptions {
categories?: string[];
}
export interface DeleteAllMemoryOptions extends EntityOptions {}
export interface DeleteAllMemoryOptions {
filters?: Record<string, any>;
}
// ─── Project Options ────────────────────────────────────────
export interface ProjectOptions {
@@ -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");
});
@@ -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
@@ -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
@@ -187,7 +187,7 @@ export async function cleanupTestUser(
userId: string,
): Promise<void> {
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<void> {
// 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
@@ -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<string, { status: number; body: unknown }>();
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]) =>
@@ -83,6 +83,89 @@ describe("MemoryClient - search()", () => {
});
});
test("passes complex AND filters through to the API body", async () => {
const extra = new Map<string, { status: number; body: unknown }>();
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<string, { status: number; body: unknown }>();
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<string, { status: number; body: unknown }>();
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<string, { status: number; body: unknown }>();
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<string, { status: number; body: unknown }>();
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<string, { status: number; body: unknown }>();
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();
});
});
+5 -5
View File
@@ -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);
+4 -2
View File
@@ -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;
}
+5 -5
View File
@@ -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);
+276 -97
View File
@@ -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<string, any>,
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<string, any>;
}> = [];
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<SearchResult> {
// Reject top-level entity params - must use filters instead
rejectTopLevelEntityParams(config as Record<string, any>, "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<string, any> = { ...(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<SearchResult> {
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<string, any>, "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<string, any>): 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<string, any>,
): Record<string, any> {
const processedFilters: Record<string, any> = {};
const processCondition = (
key: string,
condition: any,
): Record<string, any> => {
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<string, any> = {};
const operatorMap: Record<string, string> = {
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<string, any> = {};
for (const [subKey, subValue] of Object.entries(
condition as Record<string, any>,
)) {
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<string, any> = {};
for (const [subKey, subValue] of Object.entries(
condition as Record<string, any>,
)) {
Object.assign(notCondition, processCondition(subKey, subValue));
}
processedFilters["$not"].push(notCondition);
}
} else {
Object.assign(processedFilters, processCondition(key, value));
}
}
return processedFilters;
}
}
+7 -10
View File
@@ -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<string, any>;
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;
}
+3 -3
View File
@@ -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;
}
+132 -4
View File
@@ -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<string, any>,
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<string, string> = {
$and: "AND",
$or: "OR",
$not: "NOT",
};
const normalized: Record<string, any> = {};
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(
+173 -29
View File
@@ -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<string, string> = {
$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<string, number | string> = {};
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<string, any> = {};
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(
+5 -5
View File
@@ -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();
@@ -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",
);
+21 -17
View File
@@ -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(
+44 -26
View File
@@ -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);
+7 -3
View File
@@ -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);
});
});
@@ -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",
);
+68 -11
View File
@@ -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()
+20 -13
View File
@@ -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):
+184 -148
View File
@@ -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)
+149
View File
@@ -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
+7 -4
View File
@@ -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
+48 -11
View File
@@ -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")