refactor: update MemoryClient to use filters for entity parameters
This commit is contained in:
@@ -3,7 +3,6 @@ import type * as MemoryTypes from "./mem0.types";
|
||||
|
||||
// Re-export all types from mem0.types
|
||||
export type {
|
||||
EntityOptions,
|
||||
AddMemoryOptions,
|
||||
SearchMemoryOptions,
|
||||
GetAllMemoryOptions,
|
||||
|
||||
@@ -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}`,
|
||||
{
|
||||
|
||||
@@ -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();
|
||||
});
|
||||
});
|
||||
|
||||
@@ -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);
|
||||
|
||||
|
||||
@@ -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;
|
||||
}
|
||||
|
||||
@@ -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);
|
||||
|
||||
|
||||
@@ -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;
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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;
|
||||
}
|
||||
|
||||
@@ -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;
|
||||
}
|
||||
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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",
|
||||
);
|
||||
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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);
|
||||
|
||||
@@ -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
@@ -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
@@ -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
@@ -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)
|
||||
|
||||
@@ -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
@@ -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
@@ -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")
|
||||
|
||||
Reference in New Issue
Block a user