refactor: update MemoryClient to use entity options instead of filters
This commit is contained in:
+15
-30
@@ -23,8 +23,17 @@ 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"];
|
||||
// Entity params that must be passed via filters - check both snake_case and camelCase
|
||||
const ENTITY_PARAMS = [
|
||||
"user_id",
|
||||
"agent_id",
|
||||
"app_id",
|
||||
"run_id",
|
||||
"userId",
|
||||
"agentId",
|
||||
"appId",
|
||||
"runId",
|
||||
];
|
||||
|
||||
/**
|
||||
* Validates that no top-level entity parameters are passed.
|
||||
@@ -207,13 +216,7 @@ export default class MemoryClient {
|
||||
): Promise<Array<Memory>> {
|
||||
if (this.telemetryId === "") await this.ping();
|
||||
|
||||
// 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 payload = this._preparePayload(messages, options);
|
||||
const payloadKeys = Object.keys(payload);
|
||||
this._captureEvent("add", [payloadKeys]);
|
||||
|
||||
@@ -355,27 +358,9 @@ export default class MemoryClient {
|
||||
if (this.telemetryId === "") await this.ping();
|
||||
const payloadKeys = Object.keys(options || {});
|
||||
this._captureEvent("delete_all", [payloadKeys]);
|
||||
|
||||
// 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 snakeOptions = camelToSnakeKeys(this._prepareParams(options));
|
||||
// @ts-ignore
|
||||
const params = new URLSearchParams(snakeOptions);
|
||||
const response = await this._fetchWithErrorHandling(
|
||||
`${this.host}/v1/memories/?${params}`,
|
||||
{
|
||||
|
||||
@@ -1,6 +1,13 @@
|
||||
// ─── 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 {
|
||||
filters?: Record<string, any>;
|
||||
export interface AddMemoryOptions extends EntityOptions {
|
||||
metadata?: Record<string, any>;
|
||||
infer?: boolean;
|
||||
customCategories?: custom_categories[];
|
||||
@@ -28,9 +35,7 @@ export interface GetAllMemoryOptions {
|
||||
categories?: string[];
|
||||
}
|
||||
|
||||
export interface DeleteAllMemoryOptions {
|
||||
filters?: Record<string, any>;
|
||||
}
|
||||
export interface DeleteAllMemoryOptions extends EntityOptions {}
|
||||
|
||||
// ─── Project Options ────────────────────────────────────────
|
||||
export interface ProjectOptions {
|
||||
|
||||
@@ -52,7 +52,7 @@ describeIntegration("MemoryClient Integration — CRUD", () => {
|
||||
];
|
||||
|
||||
const result = await client.add(messages, {
|
||||
filters: { user_id: TEST_USER_ID },
|
||||
userId: TEST_USER_ID,
|
||||
});
|
||||
|
||||
// v3 API processes memories asynchronously — returns PENDING
|
||||
@@ -73,7 +73,7 @@ describeIntegration("MemoryClient Integration — CRUD", () => {
|
||||
];
|
||||
|
||||
const result = await client.add(messages, {
|
||||
filters: { user_id: TEST_USER_ID },
|
||||
userId: TEST_USER_ID,
|
||||
});
|
||||
expect(result).toHaveProperty("eventId");
|
||||
});
|
||||
@@ -208,7 +208,7 @@ describeIntegration("MemoryClient Integration — CRUD", () => {
|
||||
|
||||
test("deleteAll for non-existent user does not throw", async () => {
|
||||
const result = await client.deleteAll({
|
||||
filters: { user_id: `nonexistent-user-${randomUUID()}` },
|
||||
userId: `nonexistent-user-${randomUUID()}`,
|
||||
});
|
||||
|
||||
expect(result).toBeDefined();
|
||||
@@ -239,7 +239,7 @@ describeIntegration("MemoryClient Integration — CRUD", () => {
|
||||
describe("cleanup operations", () => {
|
||||
test("deletes all memories for test user", async () => {
|
||||
const result = await client.deleteAll({
|
||||
filters: { user_id: TEST_USER_ID },
|
||||
userId: TEST_USER_ID,
|
||||
});
|
||||
expect(result).toBeDefined();
|
||||
expect(typeof result.message).toBe("string");
|
||||
|
||||
@@ -18,12 +18,10 @@ export default async function globalSetup() {
|
||||
// Full project wipe — all four filters set explicitly
|
||||
try {
|
||||
await client.deleteAll({
|
||||
filters: {
|
||||
user_id: "*",
|
||||
agent_id: "*",
|
||||
app_id: "*",
|
||||
run_id: "*",
|
||||
},
|
||||
userId: "*",
|
||||
agentId: "*",
|
||||
appId: "*",
|
||||
runId: "*",
|
||||
});
|
||||
} catch {
|
||||
// ignore — may 404 if no data exists
|
||||
|
||||
@@ -17,12 +17,10 @@ export default async function globalTeardown() {
|
||||
|
||||
try {
|
||||
await client.deleteAll({
|
||||
filters: {
|
||||
user_id: "*",
|
||||
agent_id: "*",
|
||||
app_id: "*",
|
||||
run_id: "*",
|
||||
},
|
||||
userId: "*",
|
||||
agentId: "*",
|
||||
appId: "*",
|
||||
runId: "*",
|
||||
});
|
||||
} catch {
|
||||
// ignore
|
||||
|
||||
@@ -187,7 +187,7 @@ export async function cleanupTestUser(
|
||||
userId: string,
|
||||
): Promise<void> {
|
||||
try {
|
||||
await client.deleteAll({ filters: { user_id: userId } });
|
||||
await client.deleteAll({ userId });
|
||||
} catch {
|
||||
// ignore
|
||||
}
|
||||
@@ -210,12 +210,10 @@ export async function fullProjectCleanup(client: MemoryClient): Promise<void> {
|
||||
// Delete all memories — all four filters set explicitly
|
||||
try {
|
||||
await client.deleteAll({
|
||||
filters: {
|
||||
user_id: "*",
|
||||
agent_id: "*",
|
||||
app_id: "*",
|
||||
run_id: "*",
|
||||
},
|
||||
userId: "*",
|
||||
agentId: "*",
|
||||
appId: "*",
|
||||
runId: "*",
|
||||
});
|
||||
} catch {
|
||||
// ignore — may 404 if no data exists
|
||||
|
||||
@@ -27,9 +27,7 @@ describe("MemoryClient - add()", () => {
|
||||
const mock = setupMockFetch(extra);
|
||||
|
||||
const client = new MemoryClient({ apiKey: TEST_API_KEY });
|
||||
await client.add([{ role: "user", content: "Hello" }], {
|
||||
filters: { user_id: "u1" },
|
||||
});
|
||||
await client.add([{ role: "user", content: "Hello" }], { userId: "u1" });
|
||||
|
||||
expect(findFetchCall(mock, "/v3/memories/", "POST")).toBeDefined();
|
||||
});
|
||||
@@ -41,26 +39,24 @@ describe("MemoryClient - add()", () => {
|
||||
const mock = setupMockFetch(extra);
|
||||
|
||||
const client = new MemoryClient({ apiKey: TEST_API_KEY });
|
||||
await client.add(messages, { filters: { user_id: "u1" } });
|
||||
await client.add(messages, { userId: "u1" });
|
||||
|
||||
const call = findFetchCall(mock, "/v3/memories/", "POST");
|
||||
expect(getFetchBody(call!).messages).toEqual(messages);
|
||||
});
|
||||
|
||||
test("spreads filters into top-level request body (API expects flat entity IDs)", async () => {
|
||||
test("includes user_id in request body", async () => {
|
||||
const extra = new Map<string, { status: number; body: unknown }>();
|
||||
extra.set("/v3/memories/", { status: 200, body: [createMockMemory()] });
|
||||
const mock = setupMockFetch(extra);
|
||||
|
||||
const client = new MemoryClient({ apiKey: TEST_API_KEY });
|
||||
await client.add([{ role: "user", content: "test" }], {
|
||||
filters: { user_id: "user_1" },
|
||||
user_id: "user_1",
|
||||
});
|
||||
|
||||
const call = findFetchCall(mock, "/v3/memories/", "POST");
|
||||
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");
|
||||
expect(getFetchBody(call!).user_id).toBe("user_1");
|
||||
});
|
||||
|
||||
test("sends empty messages array without crashing", async () => {
|
||||
@@ -69,7 +65,7 @@ describe("MemoryClient - add()", () => {
|
||||
const mock = setupMockFetch(extra);
|
||||
|
||||
const client = new MemoryClient({ apiKey: TEST_API_KEY });
|
||||
await client.add([], { filters: { user_id: "u1" } });
|
||||
await client.add([], { userId: "u1" });
|
||||
|
||||
const call = findFetchCall(mock, "/v3/memories/", "POST");
|
||||
expect(getFetchBody(call!).messages).toEqual([]);
|
||||
@@ -219,7 +215,7 @@ describe("MemoryClient - deleteAll()", () => {
|
||||
const mock = setupMockFetch(extra);
|
||||
|
||||
const client = new MemoryClient({ apiKey: TEST_API_KEY });
|
||||
await client.deleteAll({ filters: { user_id: "u1" } });
|
||||
await client.deleteAll({ userId: "u1" });
|
||||
|
||||
const call = mock.mock.calls.find(
|
||||
(c: [string, RequestInit]) =>
|
||||
@@ -235,7 +231,7 @@ describe("MemoryClient - deleteAll()", () => {
|
||||
const mock = setupMockFetch(extra);
|
||||
|
||||
const client = new MemoryClient({ apiKey: TEST_API_KEY });
|
||||
await client.deleteAll({ filters: { user_id: "user@email.com" } });
|
||||
await client.deleteAll({ userId: "user@email.com" });
|
||||
|
||||
const call = mock.mock.calls.find(
|
||||
(c: [string, RequestInit]) =>
|
||||
|
||||
@@ -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.",
|
||||
{
|
||||
filters: { user_id: "john" },
|
||||
userId: "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." },
|
||||
],
|
||||
{
|
||||
filters: { user_id: "john" },
|
||||
userId: "john",
|
||||
},
|
||||
);
|
||||
console.log("Added messages:", result2);
|
||||
@@ -58,7 +58,7 @@ async function runTests(memory: Memory) {
|
||||
},
|
||||
],
|
||||
{
|
||||
filters: { user_id: "john" },
|
||||
userId: "john",
|
||||
},
|
||||
);
|
||||
console.log("Updated messages:", result3);
|
||||
|
||||
@@ -52,7 +52,7 @@ ${memoriesStr}`;
|
||||
const assistantResponse = response.message.content || "";
|
||||
|
||||
messages.push({ role: "assistant", content: assistantResponse });
|
||||
await memory.add(messages, { filters: { user_id: userId } });
|
||||
await memory.add(messages, { userId: 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.",
|
||||
{
|
||||
filters: { user_id: "john" },
|
||||
userId: "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." },
|
||||
],
|
||||
{
|
||||
filters: { user_id: "john" },
|
||||
userId: "john",
|
||||
},
|
||||
);
|
||||
console.log("Added messages:", result2);
|
||||
@@ -40,7 +40,7 @@ export async function runTests(memory: Memory) {
|
||||
},
|
||||
],
|
||||
{
|
||||
filters: { user_id: "john" },
|
||||
userId: "john",
|
||||
},
|
||||
);
|
||||
console.log("Updated messages:", result3);
|
||||
|
||||
@@ -53,8 +53,15 @@ 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"];
|
||||
// Entity params that must be passed via filters - check both snake_case and camelCase
|
||||
const ENTITY_PARAMS = [
|
||||
"user_id",
|
||||
"agent_id",
|
||||
"run_id",
|
||||
"userId",
|
||||
"agentId",
|
||||
"runId",
|
||||
];
|
||||
|
||||
/**
|
||||
* Validates that no top-level entity parameters are passed in config.
|
||||
@@ -70,7 +77,7 @@ function rejectTopLevelEntityParams(
|
||||
if (invalidKeys.length > 0) {
|
||||
throw new Error(
|
||||
`Top-level entity parameters [${invalidKeys.join(", ")}] are not supported in ${methodName}(). ` +
|
||||
`Use filters: { user_id: "..." } instead.`,
|
||||
`Use filters: { userId: "..." } instead.`,
|
||||
);
|
||||
}
|
||||
}
|
||||
@@ -276,21 +283,26 @@ export class Memory {
|
||||
has_filters: !!config.filters,
|
||||
infer: config.infer,
|
||||
});
|
||||
const { metadata = {}, filters = {}, infer = true } = config;
|
||||
const {
|
||||
userId,
|
||||
agentId,
|
||||
runId,
|
||||
metadata = {},
|
||||
filters = {},
|
||||
infer = true,
|
||||
} = config;
|
||||
|
||||
// Convert camelCase entity params to snake_case for storage (matches API and search/getAll filters)
|
||||
if (userId) filters.user_id = metadata.user_id = userId;
|
||||
if (agentId) filters.agent_id = metadata.agent_id = agentId;
|
||||
if (runId) filters.run_id = metadata.run_id = runId;
|
||||
|
||||
// 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' } }",
|
||||
"One of the filters: userId, agentId or runId is required!",
|
||||
);
|
||||
}
|
||||
|
||||
// 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 }];
|
||||
@@ -1017,18 +1029,22 @@ export class Memory {
|
||||
config: DeleteAllMemoryOptions,
|
||||
): Promise<{ message: string }> {
|
||||
await this._ensureInitialized();
|
||||
const { filters = {} } = config;
|
||||
|
||||
await this._captureEvent("delete_all", {
|
||||
has_user_id: !!filters.user_id,
|
||||
has_agent_id: !!filters.agent_id,
|
||||
has_run_id: !!filters.run_id,
|
||||
has_user_id: !!config.userId,
|
||||
has_agent_id: !!config.agentId,
|
||||
has_run_id: !!config.runId,
|
||||
});
|
||||
const { userId, agentId, runId } = config;
|
||||
|
||||
if (!filters.user_id && !filters.agent_id && !filters.run_id) {
|
||||
// Convert camelCase entity params to snake_case for filters (matches storage and search/getAll)
|
||||
const filters: SearchFilters = {};
|
||||
if (userId) filters.user_id = userId;
|
||||
if (agentId) filters.agent_id = agentId;
|
||||
if (runId) filters.run_id = runId;
|
||||
|
||||
if (!Object.keys(filters).length) {
|
||||
throw new Error(
|
||||
"filters must contain at least one of: user_id, agent_id, run_id. " +
|
||||
"If you want to delete all memories, use the `reset()` method.",
|
||||
"At least one filter is required to delete all memories. If you want to delete all memories, use the `reset()` method.",
|
||||
);
|
||||
}
|
||||
|
||||
|
||||
@@ -1,7 +1,13 @@
|
||||
import { Message } from "../types";
|
||||
import { SearchFilters } from "../types";
|
||||
|
||||
export interface AddMemoryOptions {
|
||||
export interface Entity {
|
||||
userId?: string;
|
||||
agentId?: string;
|
||||
runId?: string;
|
||||
}
|
||||
|
||||
export interface AddMemoryOptions extends Entity {
|
||||
metadata?: Record<string, any>;
|
||||
filters?: SearchFilters;
|
||||
infer?: boolean;
|
||||
@@ -18,6 +24,4 @@ export interface GetAllMemoryOptions {
|
||||
filters?: SearchFilters;
|
||||
}
|
||||
|
||||
export interface DeleteAllMemoryOptions {
|
||||
filters?: SearchFilters;
|
||||
}
|
||||
export interface DeleteAllMemoryOptions extends Entity {}
|
||||
|
||||
@@ -589,7 +589,7 @@ describe("Memory – LM Studio end-to-end flow", () => {
|
||||
disableHistory: true,
|
||||
});
|
||||
|
||||
await mem.add("I love sushi", { filters: { user_id: "u1" } });
|
||||
await mem.add("I love sushi", { userId: "u1" });
|
||||
|
||||
expect(mockLlm.generateResponse).toHaveBeenCalled();
|
||||
expect(mockEmbedder.embed).toHaveBeenCalled();
|
||||
|
||||
@@ -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", {
|
||||
filters: { user_id: userId },
|
||||
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",
|
||||
{ filters: { user_id: userId } },
|
||||
{ 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", {
|
||||
filters: { user_id: userId },
|
||||
userId,
|
||||
});
|
||||
expect(typeof result.results[0].memory).toBe("string");
|
||||
});
|
||||
@@ -122,35 +122,31 @@ describe("Memory - add()", () => {
|
||||
{ role: "user", content: "What is your favorite city?" },
|
||||
{ role: "assistant", content: "I love Paris." },
|
||||
];
|
||||
const result: SearchResult = await memory.add(messages, {
|
||||
filters: { user_id: userId },
|
||||
});
|
||||
const result: SearchResult = await memory.add(messages, { userId });
|
||||
expect(result.results.length).toBeGreaterThan(0);
|
||||
});
|
||||
|
||||
test("works with agent_id filter instead of user_id", async () => {
|
||||
test("works with agentId instead of userId", async () => {
|
||||
const result: SearchResult = await memory.add("test", {
|
||||
filters: { agent_id: "agent_1" },
|
||||
agentId: "agent_1",
|
||||
});
|
||||
expect(result.results.length).toBeGreaterThan(0);
|
||||
});
|
||||
|
||||
test("works with run_id filter instead of user_id", async () => {
|
||||
const result: SearchResult = await memory.add("test", {
|
||||
filters: { run_id: "run_1" },
|
||||
});
|
||||
test("works with runId instead of userId", async () => {
|
||||
const result: SearchResult = await memory.add("test", { runId: "run_1" });
|
||||
expect(result.results.length).toBeGreaterThan(0);
|
||||
});
|
||||
|
||||
test("throws when no user_id/agent_id/run_id provided in filters", async () => {
|
||||
test("throws when no userId/agentId/runId provided", async () => {
|
||||
await expect(memory.add("test", {} as any)).rejects.toThrow(
|
||||
"filters must contain at least one of: user_id, agent_id, run_id",
|
||||
"One of the filters: userId, agentId or runId is required!",
|
||||
);
|
||||
});
|
||||
|
||||
test("passes metadata through to stored memory", async () => {
|
||||
const result: SearchResult = await memory.add("I love TypeScript", {
|
||||
filters: { user_id: userId },
|
||||
userId,
|
||||
metadata: { source: "chat", tag: "programming" },
|
||||
});
|
||||
const stored: MemoryItem | null = await memory.get(result.results[0].id);
|
||||
@@ -162,7 +158,7 @@ describe("Memory - add()", () => {
|
||||
|
||||
test("with infer=false skips LLM and stores messages directly", async () => {
|
||||
const result: SearchResult = await memory.add("Direct storage content", {
|
||||
filters: { user_id: userId },
|
||||
userId,
|
||||
infer: false,
|
||||
});
|
||||
expect(result.results.length).toBeGreaterThan(0);
|
||||
@@ -172,7 +168,7 @@ describe("Memory - add()", () => {
|
||||
|
||||
test("with infer=false marks event as ADD in metadata", async () => {
|
||||
const result: SearchResult = await memory.add("Direct fact", {
|
||||
filters: { user_id: userId },
|
||||
userId,
|
||||
infer: false,
|
||||
});
|
||||
expect(result.results[0].metadata).toEqual(
|
||||
|
||||
@@ -97,7 +97,7 @@ describe("Memory - get()", () => {
|
||||
|
||||
test("returns the memory matching the ID from add()", async () => {
|
||||
const addResult: SearchResult = await memory.add("I love AI", {
|
||||
filters: { user_id: userId },
|
||||
userId,
|
||||
});
|
||||
const id = addResult.results[0].id;
|
||||
const item: MemoryItem | null = await memory.get(id);
|
||||
@@ -107,7 +107,7 @@ describe("Memory - get()", () => {
|
||||
|
||||
test("returns a string for the memory field", async () => {
|
||||
const addResult: SearchResult = await memory.add("Testing get", {
|
||||
filters: { user_id: userId },
|
||||
userId,
|
||||
});
|
||||
const item: MemoryItem | null = await memory.get(addResult.results[0].id);
|
||||
expect(typeof item!.memory).toBe("string");
|
||||
@@ -120,7 +120,7 @@ describe("Memory - get()", () => {
|
||||
|
||||
test("returns hash and createdAt on stored memory", async () => {
|
||||
const addResult: SearchResult = await memory.add("Hash test", {
|
||||
filters: { user_id: userId },
|
||||
userId,
|
||||
});
|
||||
const item: MemoryItem | null = await memory.get(addResult.results[0].id);
|
||||
expect(typeof item!.hash).toBe("string");
|
||||
@@ -146,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", {
|
||||
filters: { user_id: userId },
|
||||
userId,
|
||||
infer: false,
|
||||
});
|
||||
const id = addResult.results[0].id;
|
||||
@@ -156,7 +156,7 @@ describe("Memory - update()", () => {
|
||||
|
||||
test("persists the updated text", async () => {
|
||||
const addResult: SearchResult = await memory.add("Before update", {
|
||||
filters: { user_id: userId },
|
||||
userId,
|
||||
infer: false,
|
||||
});
|
||||
const id = addResult.results[0].id;
|
||||
@@ -167,7 +167,7 @@ describe("Memory - update()", () => {
|
||||
|
||||
test("preserves createdAt and sets updatedAt", async () => {
|
||||
const addResult: SearchResult = await memory.add("Timestamp test", {
|
||||
filters: { user_id: userId },
|
||||
userId,
|
||||
infer: false,
|
||||
});
|
||||
const id = addResult.results[0].id;
|
||||
@@ -182,7 +182,7 @@ describe("Memory - update()", () => {
|
||||
|
||||
test("updates the hash", async () => {
|
||||
const addResult: SearchResult = await memory.add("Hash change", {
|
||||
filters: { user_id: userId },
|
||||
userId,
|
||||
infer: false,
|
||||
});
|
||||
const id = addResult.results[0].id;
|
||||
@@ -209,7 +209,7 @@ describe("Memory - delete()", () => {
|
||||
|
||||
test("returns success message", async () => {
|
||||
const addResult: SearchResult = await memory.add("Delete me", {
|
||||
filters: { user_id: userId },
|
||||
userId,
|
||||
infer: false,
|
||||
});
|
||||
const result = await memory.delete(addResult.results[0].id);
|
||||
@@ -218,7 +218,7 @@ describe("Memory - delete()", () => {
|
||||
|
||||
test("get() returns null after deletion", async () => {
|
||||
const addResult: SearchResult = await memory.add("Temporary", {
|
||||
filters: { user_id: userId },
|
||||
userId,
|
||||
infer: false,
|
||||
});
|
||||
const id = addResult.results[0].id;
|
||||
@@ -242,9 +242,9 @@ describe("Memory - deleteAll()", () => {
|
||||
});
|
||||
|
||||
test("removes all memories for the user and returns success", async () => {
|
||||
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 } });
|
||||
await memory.add("Fact A", { userId });
|
||||
await memory.add("Fact B", { userId });
|
||||
const result = await memory.deleteAll({ userId });
|
||||
expect(result.message).toBe("Memories deleted successfully!");
|
||||
const remaining: SearchResult = await memory.getAll({
|
||||
filters: { user_id: userId },
|
||||
@@ -254,7 +254,7 @@ describe("Memory - deleteAll()", () => {
|
||||
|
||||
test("throws when no filter is provided", async () => {
|
||||
await expect(memory.deleteAll({} as any)).rejects.toThrow(
|
||||
"filters must contain at least one of: user_id, agent_id, run_id",
|
||||
"At least one filter is required to delete all memories",
|
||||
);
|
||||
});
|
||||
});
|
||||
@@ -274,8 +274,8 @@ describe("Memory - getAll()", () => {
|
||||
});
|
||||
|
||||
test("returns all stored memories for the user", async () => {
|
||||
await memory.add("First", { filters: { user_id: userId } });
|
||||
await memory.add("Second", { filters: { user_id: userId } });
|
||||
await memory.add("First", { userId });
|
||||
await memory.add("Second", { userId });
|
||||
const result: SearchResult = await memory.getAll({
|
||||
filters: { user_id: userId },
|
||||
});
|
||||
@@ -309,7 +309,7 @@ describe("Memory - search()", () => {
|
||||
|
||||
beforeAll(async () => {
|
||||
memory = createMemory();
|
||||
await memory.add("I love TypeScript", { filters: { user_id: userId } });
|
||||
await memory.add("I love TypeScript", { userId });
|
||||
});
|
||||
|
||||
afterAll(async () => {
|
||||
@@ -362,7 +362,7 @@ describe("Memory - history()", () => {
|
||||
|
||||
test("records ADD event after add()", async () => {
|
||||
const addResult: SearchResult = await memory.add("New fact", {
|
||||
filters: { user_id: userId },
|
||||
userId,
|
||||
});
|
||||
const history = await memory.history(addResult.results[0].id);
|
||||
expect(Array.isArray(history)).toBe(true);
|
||||
@@ -371,7 +371,7 @@ describe("Memory - history()", () => {
|
||||
|
||||
test("records additional entry after update()", async () => {
|
||||
const addResult: SearchResult = await memory.add("Before", {
|
||||
filters: { user_id: userId },
|
||||
userId,
|
||||
});
|
||||
const id = addResult.results[0].id;
|
||||
await memory.update(id, "After");
|
||||
|
||||
@@ -118,7 +118,7 @@ describe("Memory - reset()", () => {
|
||||
const mem = createMemory();
|
||||
const userId = `reset_test_${Date.now()}`;
|
||||
|
||||
await mem.add("Remember this fact", { filters: { user_id: userId } });
|
||||
await mem.add("Remember this fact", { userId });
|
||||
const before: SearchResult = await mem.getAll({
|
||||
filters: { user_id: userId },
|
||||
});
|
||||
|
||||
@@ -1070,9 +1070,7 @@ describe("Memory class – backward compat with all providers", () => {
|
||||
expect(deleteResult.message).toBe("Memory deleted successfully!");
|
||||
|
||||
// deleteAll
|
||||
const deleteAllResult = await mem.deleteAll({
|
||||
filters: { user_id: "u1" },
|
||||
});
|
||||
const deleteAllResult = await mem.deleteAll({ userId: "u1" });
|
||||
expect(deleteAllResult.message).toBe("Memories deleted successfully!");
|
||||
|
||||
// history
|
||||
|
||||
+9
-32
@@ -142,7 +142,8 @@ 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 metadata, filters.
|
||||
**kwargs: Additional parameters such as user_id, agent_id, app_id,
|
||||
metadata, filters.
|
||||
|
||||
Returns:
|
||||
A dictionary containing the API response in v1.1 format.
|
||||
@@ -165,11 +166,7 @@ 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:
|
||||
@@ -370,7 +367,8 @@ class MemoryClient:
|
||||
|
||||
Args:
|
||||
options: Typed options for the delete_all operation (DeleteAllMemoryOptions).
|
||||
**kwargs: Optional parameters for filtering (filters).
|
||||
**kwargs: Optional parameters for filtering (user_id, agent_id,
|
||||
app_id).
|
||||
|
||||
Returns:
|
||||
A dictionary containing the API response.
|
||||
@@ -385,16 +383,7 @@ class MemoryClient:
|
||||
"""
|
||||
kwargs = {**(options.model_dump(exclude_unset=True) if options else {}), **kwargs}
|
||||
params = self._prepare_params(kwargs)
|
||||
|
||||
# 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 = self.client.delete("/v1/memories/", params=params)
|
||||
response.raise_for_status()
|
||||
capture_client_event(
|
||||
"client.delete_all",
|
||||
@@ -1081,7 +1070,8 @@ 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 metadata, filters.
|
||||
**kwargs: Additional parameters such as user_id, agent_id, app_id,
|
||||
metadata, filters.
|
||||
|
||||
Returns:
|
||||
A dictionary containing the API response in v1.1 format.
|
||||
@@ -1104,11 +1094,7 @@ 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:
|
||||
@@ -1293,7 +1279,7 @@ class AsyncMemoryClient:
|
||||
|
||||
Args:
|
||||
options: Typed options for the delete_all operation (DeleteAllMemoryOptions).
|
||||
**kwargs: Optional parameters for filtering (filters).
|
||||
**kwargs: Optional parameters for filtering (user_id, agent_id, app_id).
|
||||
|
||||
Returns:
|
||||
A dictionary containing the API response.
|
||||
@@ -1308,16 +1294,7 @@ class AsyncMemoryClient:
|
||||
"""
|
||||
kwargs = {**(options.model_dump(exclude_unset=True) if options else {}), **kwargs}
|
||||
params = self._prepare_params(kwargs)
|
||||
|
||||
# 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 = await self.async_client.delete("/v1/memories/", params=params)
|
||||
response.raise_for_status()
|
||||
capture_client_event("client.delete_all", self, {"keys": list(kwargs.keys()), "sync_type": "async"})
|
||||
return response.json()
|
||||
|
||||
+45
-61
@@ -425,27 +425,26 @@ class Memory(MemoryBase):
|
||||
self,
|
||||
messages,
|
||||
*,
|
||||
filters: Optional[Dict[str, Any]] = None,
|
||||
user_id: Optional[str] = None,
|
||||
agent_id: Optional[str] = None,
|
||||
run_id: Optional[str] = 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 in filters.
|
||||
Adds new memories scoped to a single session id (e.g. `user_id`, `agent_id`, or `run_id`). One of those ids is required.
|
||||
|
||||
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.
|
||||
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"}`
|
||||
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.
|
||||
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.
|
||||
@@ -469,19 +468,11 @@ 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=filters.get("user_id"),
|
||||
agent_id=filters.get("agent_id"),
|
||||
run_id=filters.get("run_id"),
|
||||
user_id=user_id,
|
||||
agent_id=agent_id,
|
||||
run_id=run_id,
|
||||
input_metadata=metadata,
|
||||
)
|
||||
|
||||
@@ -507,7 +498,7 @@ class Memory(MemoryBase):
|
||||
suggestion="Convert your input to a string, dictionary, or list of dictionaries."
|
||||
)
|
||||
|
||||
if filters.get("agent_id") is not None and memory_type == MemoryType.PROCEDURAL.value:
|
||||
if agent_id is not None and memory_type == MemoryType.PROCEDURAL.value:
|
||||
results = self._create_procedural_memory(messages, metadata=processed_metadata, prompt=prompt)
|
||||
return results
|
||||
|
||||
@@ -1360,23 +1351,26 @@ class Memory(MemoryBase):
|
||||
self._delete_memory(memory_id, existing_memory)
|
||||
return {"message": "Memory deleted successfully!"}
|
||||
|
||||
def delete_all(self, *, filters: Optional[Dict[str, Any]] = None, **kwargs):
|
||||
def delete_all(self, user_id: Optional[str] = None, agent_id: Optional[str] = None, run_id: Optional[str] = None):
|
||||
"""
|
||||
Delete all memories.
|
||||
|
||||
Args:
|
||||
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"}`
|
||||
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 = filters or {}
|
||||
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
|
||||
|
||||
if not any(key in filters for key in ("user_id", "agent_id", "run_id")):
|
||||
if not filters:
|
||||
raise ValueError(
|
||||
"filters must contain at least one of: user_id, agent_id, run_id. "
|
||||
"If you want to delete all memories, use the `reset()` method."
|
||||
"At least one filter is required to delete all memories. If you want to delete all memories, use the `reset()` method."
|
||||
)
|
||||
|
||||
keys, encoded_ids = process_telemetry_filters(filters)
|
||||
@@ -1694,24 +1688,23 @@ class AsyncMemory(MemoryBase):
|
||||
self,
|
||||
messages,
|
||||
*,
|
||||
filters: Optional[Dict[str, Any]] = None,
|
||||
user_id: Optional[str] = None,
|
||||
agent_id: Optional[str] = None,
|
||||
run_id: Optional[str] = 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.
|
||||
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"}`
|
||||
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.
|
||||
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.
|
||||
@@ -1721,20 +1714,8 @@ 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=filters.get("user_id"),
|
||||
agent_id=filters.get("agent_id"),
|
||||
run_id=filters.get("run_id"),
|
||||
input_metadata=metadata,
|
||||
user_id=user_id, agent_id=agent_id, run_id=run_id, input_metadata=metadata
|
||||
)
|
||||
|
||||
if memory_type is not None and memory_type != MemoryType.PROCEDURAL.value:
|
||||
@@ -1756,7 +1737,7 @@ class AsyncMemory(MemoryBase):
|
||||
suggestion="Convert your input to a string, dictionary, or list of dictionaries."
|
||||
)
|
||||
|
||||
if filters.get("agent_id") is not None and memory_type == MemoryType.PROCEDURAL.value:
|
||||
if 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
|
||||
)
|
||||
@@ -2606,23 +2587,26 @@ class AsyncMemory(MemoryBase):
|
||||
await self._delete_memory(memory_id, existing_memory)
|
||||
return {"message": "Memory deleted successfully!"}
|
||||
|
||||
async def delete_all(self, *, filters: Optional[Dict[str, Any]] = None, **kwargs):
|
||||
async def delete_all(self, user_id=None, agent_id=None, run_id=None):
|
||||
"""
|
||||
Delete all memories asynchronously.
|
||||
|
||||
Args:
|
||||
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"}`
|
||||
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 = filters or {}
|
||||
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
|
||||
|
||||
if not any(key in filters for key in ("user_id", "agent_id", "run_id")):
|
||||
if not filters:
|
||||
raise ValueError(
|
||||
"filters must contain at least one of: user_id, agent_id, run_id. "
|
||||
"If you want to delete all memories, use the `reset()` method."
|
||||
"At least one filter is required to delete all memories. If you want to delete all memories, use the `reset()` method."
|
||||
)
|
||||
|
||||
keys, encoded_ids = process_telemetry_filters(filters)
|
||||
|
||||
+2
-5
@@ -57,10 +57,7 @@ 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"}],
|
||||
filters={"user_id": "test_user"},
|
||||
)
|
||||
result = memory_instance.add(messages=[{"role": "user", "content": "Test message"}], user_id="test_user")
|
||||
|
||||
assert "results" in result
|
||||
assert result["results"] == [{"memory": "Test memory", "event": "ADD"}]
|
||||
@@ -188,7 +185,7 @@ def test_delete_all(memory_instance, version):
|
||||
memory_instance.vector_store.reset = Mock()
|
||||
memory_instance._delete_memory = Mock()
|
||||
|
||||
result = memory_instance.delete_all(filters={"user_id": "test_user"})
|
||||
result = memory_instance.delete_all(user_id="test_user")
|
||||
|
||||
assert memory_instance._delete_memory.call_count == 2
|
||||
# Ensure the collection is NOT dropped — only matched memories should be removed
|
||||
|
||||
+10
-10
@@ -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}], filters={"user_id": "test_user"})
|
||||
result = memory_client.add([{"role": "user", "content": data}], 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}], filters={"user_id": "test_user"})
|
||||
memory_client.add([{"role": "user", "content": data}], 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}], filters={"user_id": "test_user"})
|
||||
memory_client.add([{"role": "user", "content": data}], 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}], filters={"user_id": "test_user"})
|
||||
memory_client.add([{"role": "user", "content": data}], 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}], filters={"user_id": "test_user"})
|
||||
memory_client.add([{"role": "user", "content": data}], 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,8 +71,8 @@ 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}], filters={"user_id": "test_user"})
|
||||
memory_client.add([{"role": "user", "content": data2}], filters={"user_id": "test_user"})
|
||||
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(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", filters={"user_id": "test_user"}, infer=False)
|
||||
memory.add("foo", 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", filters={"user_id": "test_user"}, infer=True)
|
||||
memory.add("I like Python", 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", filters={"user_id": "test_user"}, infer=True)
|
||||
memory.add("I love Python now", 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)
|
||||
|
||||
Reference in New Issue
Block a user