Compare commits
2 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 189aea04c9 | |||
| e830351691 |
@@ -729,19 +729,6 @@ mode: "wide"
|
||||
|
||||
<Tab title="TypeScript">
|
||||
|
||||
<Update label="2026-03-14" description="v2.4.0">
|
||||
|
||||
**Bug Fixes:**
|
||||
- **OSS Storage:** Fixed `SQLITE_CANTOPEN` errors when running as a LaunchAgent, systemd service, or in containers where `process.cwd()` is read-only (e.g. `/`). Default `vector_store.db` location changed from `process.cwd()/vector_store.db` to `~/.mem0/vector_store.db`.
|
||||
- **OSS Storage:** Fixed `historyDbPath` config being silently ignored — config merging always overwrote it with defaults. Top-level `historyDbPath` is now correctly propagated into `historyStore.config` with proper precedence.
|
||||
- **OSS Storage:** Added `ensureSQLiteDirectory()` — parent directories for SQLite database files are now auto-created before opening, preventing `SQLITE_CANTOPEN` when using nested paths.
|
||||
|
||||
**Improvements:**
|
||||
- **Migration:** Added deprecation warning when an existing `vector_store.db` is found at the old `process.cwd()` location, guiding users to move it or set `vectorStore.config.dbPath` explicitly.
|
||||
- **Config:** Limited default SQLite config spreading to only SQLite history providers, preventing config leaking into Supabase or other providers.
|
||||
|
||||
</Update>
|
||||
|
||||
<Update label="2026-03-09" description="v2.3.0">
|
||||
|
||||
**Breaking Changes:**
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
{
|
||||
"name": "mem0ai",
|
||||
"version": "2.4.0",
|
||||
"version": "2.3.0",
|
||||
"description": "The Memory Layer For Your AI Apps",
|
||||
"main": "./dist/index.js",
|
||||
"module": "./dist/index.mjs",
|
||||
@@ -103,7 +103,7 @@
|
||||
"@azure/search-documents": "^12.0.0",
|
||||
"@cloudflare/workers-types": "^4.20250504.0",
|
||||
"@google/genai": "^1.2.0",
|
||||
"@langchain/core": "^1.0.0",
|
||||
"@langchain/core": "^0.3.44",
|
||||
"@mistralai/mistralai": "^1.5.2",
|
||||
"@qdrant/js-client-rest": "1.13.0",
|
||||
"@supabase/supabase-js": "^2.49.1",
|
||||
|
||||
@@ -278,9 +278,5 @@ export function parseMessages(messages: string[]): string {
|
||||
}
|
||||
|
||||
export function removeCodeBlocks(text: string): string {
|
||||
// Extract content inside code fences instead of deleting it.
|
||||
// The old regex /```[^`]*```/g replaced the entire block (including
|
||||
// its content) with an empty string, so when an LLM returned JSON
|
||||
// wrapped in ```json ... ``` the actual payload was discarded.
|
||||
return text.replace(/```(?:\w+)?\n?([\s\S]*?)```/g, "$1").trim();
|
||||
return text.replace(/```[^`]*```/g, "");
|
||||
}
|
||||
|
||||
@@ -210,7 +210,11 @@ describe("backward compat: MemoryVectorStore", () => {
|
||||
|
||||
try {
|
||||
const store = new MemoryVectorStore({ dimension: 3, dbPath });
|
||||
await store.insert([normalize([1, 0, 0])], ["id1"], [{ text: "hello" }]);
|
||||
await store.insert(
|
||||
[normalize([1, 0, 0])],
|
||||
["id1"],
|
||||
[{ text: "hello" }],
|
||||
);
|
||||
|
||||
expect(fs.existsSync(dbPath)).toBe(true);
|
||||
|
||||
|
||||
@@ -200,9 +200,9 @@ describe("MemoryVectorStore – path handling", () => {
|
||||
expect(
|
||||
fs.existsSync(path.join(fakeHome, ".mem0", "vector_store.db")),
|
||||
).toBe(true);
|
||||
expect(fs.existsSync(path.join(readOnlyCwd, "vector_store.db"))).toBe(
|
||||
false,
|
||||
);
|
||||
expect(
|
||||
fs.existsSync(path.join(readOnlyCwd, "vector_store.db")),
|
||||
).toBe(false);
|
||||
} finally {
|
||||
fs.chmodSync(readOnlyCwd, 0o755);
|
||||
fs.rmSync(fakeHome, { recursive: true, force: true });
|
||||
|
||||
@@ -273,7 +273,10 @@ export class Qdrant implements VectorStore {
|
||||
}
|
||||
}
|
||||
|
||||
private async ensureCollection(name: string, size: number): Promise<void> {
|
||||
private async ensureCollection(
|
||||
name: string,
|
||||
size: number,
|
||||
): Promise<void> {
|
||||
try {
|
||||
await this.client.createCollection(name, {
|
||||
vectors: {
|
||||
|
||||
@@ -13,10 +13,7 @@
|
||||
*/
|
||||
|
||||
import { MemoryGraph } from "../src/memory/graph_memory";
|
||||
import {
|
||||
EXTRACT_RELATIONS_PROMPT,
|
||||
getDeleteMessages,
|
||||
} from "../src/graphs/utils";
|
||||
import { EXTRACT_RELATIONS_PROMPT, getDeleteMessages } from "../src/graphs/utils";
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Mocks – we replace heavy dependencies so tests run without Neo4j / OpenAI
|
||||
@@ -103,10 +100,7 @@ describe("_retrieveNodesFromData", () => {
|
||||
});
|
||||
|
||||
const mg = graph();
|
||||
const result = await mg._retrieveNodesFromData(
|
||||
"Alice likes pizza",
|
||||
FILTERS,
|
||||
);
|
||||
const result = await mg._retrieveNodesFromData("Alice likes pizza", FILTERS);
|
||||
|
||||
expect(result).toEqual({ alice: "person", pizza: "food" });
|
||||
});
|
||||
@@ -172,9 +166,7 @@ describe("_retrieveNodesFromData", () => {
|
||||
toolCalls: [
|
||||
{
|
||||
name: "some_other_tool",
|
||||
arguments: JSON.stringify({
|
||||
entities: [{ entity: "X", entity_type: "Y" }],
|
||||
}),
|
||||
arguments: JSON.stringify({ entities: [{ entity: "X", entity_type: "Y" }] }),
|
||||
},
|
||||
],
|
||||
});
|
||||
@@ -281,7 +273,9 @@ describe("_establishNodesRelationsFromData", () => {
|
||||
|
||||
it("throws on malformed JSON in tool call arguments (no try/catch in source)", async () => {
|
||||
mockGenerateResponse.mockResolvedValueOnce({
|
||||
toolCalls: [{ name: "establish_relationships", arguments: "<<BROKEN>>" }],
|
||||
toolCalls: [
|
||||
{ name: "establish_relationships", arguments: "<<BROKEN>>" },
|
||||
],
|
||||
});
|
||||
|
||||
const mg = graph();
|
||||
@@ -353,11 +347,7 @@ describe("_getDeleteEntitiesFromSearchOutput", () => {
|
||||
});
|
||||
|
||||
const mg = graph();
|
||||
const result = await mg._getDeleteEntitiesFromSearchOutput(
|
||||
SEARCH_OUTPUT,
|
||||
"Alice hates pizza",
|
||||
FILTERS,
|
||||
);
|
||||
const result = await mg._getDeleteEntitiesFromSearchOutput(SEARCH_OUTPUT, "Alice hates pizza", FILTERS);
|
||||
|
||||
expect(result).toEqual([
|
||||
{ source: "alice", relationship: "likes", destination: "pizza" },
|
||||
@@ -368,11 +358,7 @@ describe("_getDeleteEntitiesFromSearchOutput", () => {
|
||||
mockGenerateResponse.mockResolvedValueOnce("string response");
|
||||
|
||||
const mg = graph();
|
||||
const result = await mg._getDeleteEntitiesFromSearchOutput(
|
||||
SEARCH_OUTPUT,
|
||||
"x",
|
||||
FILTERS,
|
||||
);
|
||||
const result = await mg._getDeleteEntitiesFromSearchOutput(SEARCH_OUTPUT, "x", FILTERS);
|
||||
expect(result).toEqual([]);
|
||||
});
|
||||
|
||||
@@ -380,11 +366,7 @@ describe("_getDeleteEntitiesFromSearchOutput", () => {
|
||||
mockGenerateResponse.mockResolvedValueOnce({ toolCalls: [] });
|
||||
|
||||
const mg = graph();
|
||||
const result = await mg._getDeleteEntitiesFromSearchOutput(
|
||||
SEARCH_OUTPUT,
|
||||
"x",
|
||||
FILTERS,
|
||||
);
|
||||
const result = await mg._getDeleteEntitiesFromSearchOutput(SEARCH_OUTPUT, "x", FILTERS);
|
||||
expect(result).toEqual([]);
|
||||
});
|
||||
|
||||
@@ -399,11 +381,7 @@ describe("_getDeleteEntitiesFromSearchOutput", () => {
|
||||
});
|
||||
|
||||
const mg = graph();
|
||||
const result = await mg._getDeleteEntitiesFromSearchOutput(
|
||||
SEARCH_OUTPUT,
|
||||
"x",
|
||||
FILTERS,
|
||||
);
|
||||
const result = await mg._getDeleteEntitiesFromSearchOutput(SEARCH_OUTPUT, "x", FILTERS);
|
||||
expect(result).toEqual([]);
|
||||
});
|
||||
|
||||
@@ -412,29 +390,17 @@ describe("_getDeleteEntitiesFromSearchOutput", () => {
|
||||
toolCalls: [
|
||||
{
|
||||
name: "delete_graph_memory",
|
||||
arguments: JSON.stringify({
|
||||
source: "A",
|
||||
relationship: "r1",
|
||||
destination: "B",
|
||||
}),
|
||||
arguments: JSON.stringify({ source: "A", relationship: "r1", destination: "B" }),
|
||||
},
|
||||
{
|
||||
name: "delete_graph_memory",
|
||||
arguments: JSON.stringify({
|
||||
source: "C",
|
||||
relationship: "r2",
|
||||
destination: "D",
|
||||
}),
|
||||
arguments: JSON.stringify({ source: "C", relationship: "r2", destination: "D" }),
|
||||
},
|
||||
],
|
||||
});
|
||||
|
||||
const mg = graph();
|
||||
const result = await mg._getDeleteEntitiesFromSearchOutput(
|
||||
SEARCH_OUTPUT,
|
||||
"x",
|
||||
FILTERS,
|
||||
);
|
||||
const result = await mg._getDeleteEntitiesFromSearchOutput(SEARCH_OUTPUT, "x", FILTERS);
|
||||
expect(result).toHaveLength(2);
|
||||
expect(result[0].source).toBe("a");
|
||||
expect(result[1].source).toBe("c");
|
||||
@@ -459,11 +425,7 @@ describe("_getDeleteEntitiesFromSearchOutput", () => {
|
||||
mockGenerateResponse.mockResolvedValueOnce({ toolCalls: [] });
|
||||
|
||||
const mg = graph();
|
||||
const result = await mg._getDeleteEntitiesFromSearchOutput(
|
||||
[],
|
||||
"data",
|
||||
FILTERS,
|
||||
);
|
||||
const result = await mg._getDeleteEntitiesFromSearchOutput([], "data", FILTERS);
|
||||
expect(result).toEqual([]);
|
||||
});
|
||||
});
|
||||
@@ -519,11 +481,7 @@ describe("_removeSpacesFromEntities (via _establishNodesRelationsFromData)", ()
|
||||
name: "establish_relationships",
|
||||
arguments: JSON.stringify({
|
||||
entities: [
|
||||
{
|
||||
source: "New York",
|
||||
relationship: "Capital Of",
|
||||
destination: "United States",
|
||||
},
|
||||
{ source: "New York", relationship: "Capital Of", destination: "United States" },
|
||||
],
|
||||
}),
|
||||
},
|
||||
@@ -531,18 +489,10 @@ describe("_removeSpacesFromEntities (via _establishNodesRelationsFromData)", ()
|
||||
});
|
||||
|
||||
const mg = graph();
|
||||
const result = await mg._establishNodesRelationsFromData(
|
||||
"test",
|
||||
FILTERS,
|
||||
{},
|
||||
);
|
||||
const result = await mg._establishNodesRelationsFromData("test", FILTERS, {});
|
||||
|
||||
expect(result).toEqual([
|
||||
{
|
||||
source: "new_york",
|
||||
relationship: "capital_of",
|
||||
destination: "united_states",
|
||||
},
|
||||
{ source: "new_york", relationship: "capital_of", destination: "united_states" },
|
||||
]);
|
||||
});
|
||||
});
|
||||
|
||||
@@ -31,8 +31,7 @@ describe("Graph prompts — JSON keyword requirement", () => {
|
||||
it("EXTRACT_RELATIONS_PROMPT produces a message containing 'json' once the suffix is appended", () => {
|
||||
// graph_memory.ts appends "\nPlease provide your response in JSON format."
|
||||
const withSuffix =
|
||||
EXTRACT_RELATIONS_PROMPT +
|
||||
"\nPlease provide your response in JSON format.";
|
||||
EXTRACT_RELATIONS_PROMPT + "\nPlease provide your response in JSON format.";
|
||||
expect(withSuffix.toLowerCase()).toContain("json");
|
||||
});
|
||||
|
||||
@@ -87,21 +86,13 @@ describe("getDeleteMessages", () => {
|
||||
});
|
||||
|
||||
it("handles special characters in userId (e.g. angle brackets, quotes)", () => {
|
||||
const [system] = getDeleteMessages(
|
||||
"mem",
|
||||
"data",
|
||||
'<script>alert("xss")</script>',
|
||||
);
|
||||
const [system] = getDeleteMessages("mem", "data", '<script>alert("xss")</script>');
|
||||
expect(system).toContain('<script>alert("xss")</script>');
|
||||
expect(system).not.toContain("USER_ID");
|
||||
});
|
||||
|
||||
it("handles unicode input", () => {
|
||||
const [system, user] = getDeleteMessages(
|
||||
"日本語メモリ",
|
||||
"新しい情報",
|
||||
"ユーザー1",
|
||||
);
|
||||
const [system, user] = getDeleteMessages("日本語メモリ", "新しい情報", "ユーザー1");
|
||||
expect(system).toContain("ユーザー1");
|
||||
expect(user).toContain("日本語メモリ");
|
||||
expect(user).toContain("新しい情報");
|
||||
@@ -139,7 +130,7 @@ describe("formatEntities", () => {
|
||||
|
||||
it("preserves special characters in entity fields", () => {
|
||||
const result = formatEntities([
|
||||
{ source: "O'Brien", relationship: 'said "hello"', destination: "café" },
|
||||
{ source: "O'Brien", relationship: "said \"hello\"", destination: "café" },
|
||||
]);
|
||||
expect(result).toContain("O'Brien");
|
||||
expect(result).toContain('said "hello"');
|
||||
|
||||
@@ -151,14 +151,10 @@ describeIfQdrant("Issue #4212/#4173: Qdrant dimension mismatch", () => {
|
||||
const id1 = uuidv4();
|
||||
const id2 = uuidv4();
|
||||
|
||||
await store.insert(
|
||||
[vec1, vec2],
|
||||
[id1, id2],
|
||||
[
|
||||
{ data: "hello", userId: "u1" },
|
||||
{ data: "world", userId: "u1" },
|
||||
],
|
||||
);
|
||||
await store.insert([vec1, vec2], [id1, id2], [
|
||||
{ data: "hello", userId: "u1" },
|
||||
{ data: "world", userId: "u1" },
|
||||
]);
|
||||
|
||||
// Search with 768-dim query — this USED TO fail with Bad Request
|
||||
const results = await store.search(vec1, 2, { userId: "u1" });
|
||||
@@ -207,11 +203,9 @@ describeIfQdrant("Issue #4212/#4173: Qdrant dimension mismatch", () => {
|
||||
return {
|
||||
EmbedderFactory: { create: jest.fn().mockReturnValue(fakeEmbedder) },
|
||||
VectorStoreFactory: {
|
||||
create: jest
|
||||
.fn()
|
||||
.mockImplementation((_provider: string, config: any) => {
|
||||
return new QdrantStore(config);
|
||||
}),
|
||||
create: jest.fn().mockImplementation((_provider: string, config: any) => {
|
||||
return new QdrantStore(config);
|
||||
}),
|
||||
},
|
||||
LLMFactory: {
|
||||
create: jest.fn().mockReturnValue({
|
||||
@@ -286,11 +280,9 @@ describeIfQdrant("Issue #4212/#4173: Qdrant dimension mismatch", () => {
|
||||
return {
|
||||
EmbedderFactory: { create: jest.fn().mockReturnValue(fakeEmbedder) },
|
||||
VectorStoreFactory: {
|
||||
create: jest
|
||||
.fn()
|
||||
.mockImplementation((_provider: string, config: any) => {
|
||||
return new QdrantStore(config);
|
||||
}),
|
||||
create: jest.fn().mockImplementation((_provider: string, config: any) => {
|
||||
return new QdrantStore(config);
|
||||
}),
|
||||
},
|
||||
LLMFactory: {
|
||||
create: jest.fn().mockReturnValue({
|
||||
@@ -358,11 +350,9 @@ describeIfQdrant("Issue #4212/#4173: Qdrant dimension mismatch", () => {
|
||||
return {
|
||||
EmbedderFactory: { create: jest.fn().mockReturnValue(fakeEmbedder) },
|
||||
VectorStoreFactory: {
|
||||
create: jest
|
||||
.fn()
|
||||
.mockImplementation((_provider: string, config: any) => {
|
||||
return new QdrantStore(config);
|
||||
}),
|
||||
create: jest.fn().mockImplementation((_provider: string, config: any) => {
|
||||
return new QdrantStore(config);
|
||||
}),
|
||||
},
|
||||
LLMFactory: {
|
||||
create: jest.fn().mockReturnValue({
|
||||
|
||||
@@ -1,30 +0,0 @@
|
||||
import { removeCodeBlocks } from "../src/prompts";
|
||||
|
||||
describe("removeCodeBlocks", () => {
|
||||
it("extracts JSON from ```json code fence", () => {
|
||||
const input = '```json\n{"facts": ["hello"]}\n```';
|
||||
expect(removeCodeBlocks(input)).toBe('{"facts": ["hello"]}');
|
||||
});
|
||||
|
||||
it("extracts content from bare ``` code fence", () => {
|
||||
const input = '```\n{"key": "value"}\n```';
|
||||
expect(removeCodeBlocks(input)).toBe('{"key": "value"}');
|
||||
});
|
||||
|
||||
it("returns plain text unchanged", () => {
|
||||
const input = '{"facts": ["hello"]}';
|
||||
expect(removeCodeBlocks(input)).toBe('{"facts": ["hello"]}');
|
||||
});
|
||||
|
||||
it("handles multiple code blocks", () => {
|
||||
const input = '```json\n{"a":1}\n```\nsome text\n```json\n{"b":2}\n```';
|
||||
expect(removeCodeBlocks(input)).toBe('{"a":1}\n\nsome text\n{"b":2}');
|
||||
});
|
||||
|
||||
it("handles Claude-style response with surrounding text", () => {
|
||||
const input =
|
||||
'Here is the JSON:\n```json\n{"facts": ["user likes TypeScript"]}\n```';
|
||||
expect(removeCodeBlocks(input)).toContain('"facts"');
|
||||
expect(removeCodeBlocks(input)).not.toContain("```");
|
||||
});
|
||||
});
|
||||
@@ -76,14 +76,10 @@ describe("MemoryVectorStore – full backward compat", () => {
|
||||
vec2[1] = 1.0;
|
||||
|
||||
// Insert
|
||||
await store.insert(
|
||||
[vec1, vec2],
|
||||
["id-1", "id-2"],
|
||||
[
|
||||
{ data: "alpha", userId: "u1" },
|
||||
{ data: "beta", userId: "u1" },
|
||||
],
|
||||
);
|
||||
await store.insert([vec1, vec2], ["id-1", "id-2"], [
|
||||
{ data: "alpha", userId: "u1" },
|
||||
{ data: "beta", userId: "u1" },
|
||||
]);
|
||||
|
||||
// Get
|
||||
const item = await store.get("id-1");
|
||||
@@ -208,87 +204,66 @@ describe("MemoryVectorStore – full backward compat", () => {
|
||||
describe("Qdrant – backward compat with mocked client", () => {
|
||||
function createMockQdrantClient() {
|
||||
const collections = new Map<string, number>();
|
||||
const points = new Map<
|
||||
string,
|
||||
{ id: string; vector: number[]; payload: any }
|
||||
>();
|
||||
const points = new Map<string, { id: string; vector: number[]; payload: any }>();
|
||||
|
||||
return {
|
||||
_collections: collections,
|
||||
_points: points,
|
||||
createCollection: jest
|
||||
.fn()
|
||||
.mockImplementation(async (name: string, opts: any) => {
|
||||
if (collections.has(name)) {
|
||||
const err: any = new Error("Collection already exists");
|
||||
err.status = 409;
|
||||
throw err;
|
||||
}
|
||||
collections.set(name, opts.vectors.size);
|
||||
}),
|
||||
createCollection: jest.fn().mockImplementation(async (name: string, opts: any) => {
|
||||
if (collections.has(name)) {
|
||||
const err: any = new Error("Collection already exists");
|
||||
err.status = 409;
|
||||
throw err;
|
||||
}
|
||||
collections.set(name, opts.vectors.size);
|
||||
}),
|
||||
getCollection: jest.fn().mockImplementation(async (name: string) => {
|
||||
if (!collections.has(name)) {
|
||||
const err: any = new Error("Not found");
|
||||
err.status = 404;
|
||||
throw err;
|
||||
}
|
||||
return {
|
||||
config: { params: { vectors: { size: collections.get(name) } } },
|
||||
};
|
||||
return { config: { params: { vectors: { size: collections.get(name) } } } };
|
||||
}),
|
||||
getCollections: jest.fn().mockResolvedValue({
|
||||
collections: [],
|
||||
}),
|
||||
upsert: jest
|
||||
.fn()
|
||||
.mockImplementation(async (collName: string, opts: any) => {
|
||||
for (const pt of opts.points) {
|
||||
points.set(`${collName}:${pt.id}`, {
|
||||
id: pt.id,
|
||||
vector: pt.vector,
|
||||
payload: pt.payload,
|
||||
});
|
||||
upsert: jest.fn().mockImplementation(async (collName: string, opts: any) => {
|
||||
for (const pt of opts.points) {
|
||||
points.set(`${collName}:${pt.id}`, { id: pt.id, vector: pt.vector, payload: pt.payload });
|
||||
}
|
||||
}),
|
||||
retrieve: jest.fn().mockImplementation(async (collName: string, opts: any) => {
|
||||
const results = [];
|
||||
for (const id of opts.ids) {
|
||||
const pt = points.get(`${collName}:${id}`);
|
||||
if (pt) results.push({ id: pt.id, payload: pt.payload });
|
||||
}
|
||||
return results;
|
||||
}),
|
||||
search: jest.fn().mockImplementation(async (collName: string, opts: any) => {
|
||||
const results: any[] = [];
|
||||
points.forEach((pt, key) => {
|
||||
if (key.startsWith(`${collName}:`)) {
|
||||
results.push({ id: pt.id, payload: pt.payload, score: 0.9 });
|
||||
}
|
||||
}),
|
||||
retrieve: jest
|
||||
.fn()
|
||||
.mockImplementation(async (collName: string, opts: any) => {
|
||||
const results = [];
|
||||
for (const id of opts.ids) {
|
||||
const pt = points.get(`${collName}:${id}`);
|
||||
if (pt) results.push({ id: pt.id, payload: pt.payload });
|
||||
});
|
||||
return results.slice(0, opts.limit);
|
||||
}),
|
||||
scroll: jest.fn().mockImplementation(async (collName: string, opts: any) => {
|
||||
const results: any[] = [];
|
||||
points.forEach((pt, key) => {
|
||||
if (key.startsWith(`${collName}:`)) {
|
||||
results.push({ id: pt.id, payload: pt.payload });
|
||||
}
|
||||
return results;
|
||||
}),
|
||||
search: jest
|
||||
.fn()
|
||||
.mockImplementation(async (collName: string, opts: any) => {
|
||||
const results: any[] = [];
|
||||
points.forEach((pt, key) => {
|
||||
if (key.startsWith(`${collName}:`)) {
|
||||
results.push({ id: pt.id, payload: pt.payload, score: 0.9 });
|
||||
}
|
||||
});
|
||||
return results.slice(0, opts.limit);
|
||||
}),
|
||||
scroll: jest
|
||||
.fn()
|
||||
.mockImplementation(async (collName: string, opts: any) => {
|
||||
const results: any[] = [];
|
||||
points.forEach((pt, key) => {
|
||||
if (key.startsWith(`${collName}:`)) {
|
||||
results.push({ id: pt.id, payload: pt.payload });
|
||||
}
|
||||
});
|
||||
return { points: results.slice(0, opts.limit) };
|
||||
}),
|
||||
delete: jest
|
||||
.fn()
|
||||
.mockImplementation(async (collName: string, opts: any) => {
|
||||
for (const id of opts.points) {
|
||||
points.delete(`${collName}:${id}`);
|
||||
}
|
||||
}),
|
||||
});
|
||||
return { points: results.slice(0, opts.limit) };
|
||||
}),
|
||||
delete: jest.fn().mockImplementation(async (collName: string, opts: any) => {
|
||||
for (const id of opts.points) {
|
||||
points.delete(`${collName}:${id}`);
|
||||
}
|
||||
}),
|
||||
deleteCollection: jest.fn().mockImplementation(async (name: string) => {
|
||||
collections.delete(name);
|
||||
}),
|
||||
@@ -349,10 +324,7 @@ describe("Qdrant – backward compat with mocked client", () => {
|
||||
|
||||
// Insert
|
||||
await store.insert(
|
||||
[
|
||||
[1, 2, 3],
|
||||
[4, 5, 6],
|
||||
],
|
||||
[[1, 2, 3], [4, 5, 6]],
|
||||
["id-1", "id-2"],
|
||||
[{ data: "alpha" }, { data: "beta" }],
|
||||
);
|
||||
@@ -419,9 +391,9 @@ describe("Redis – backward compat with mocked client", () => {
|
||||
connect: jest.fn().mockResolvedValue(undefined),
|
||||
on: jest.fn(),
|
||||
isOpen: false,
|
||||
moduleList: jest
|
||||
.fn()
|
||||
.mockResolvedValue([["name", "search", "ver", 20000]]),
|
||||
moduleList: jest.fn().mockResolvedValue([
|
||||
["name", "search", "ver", 20000],
|
||||
]),
|
||||
ft: {
|
||||
dropIndex: jest.fn().mockResolvedValue(undefined),
|
||||
create: jest.fn().mockResolvedValue(undefined),
|
||||
@@ -627,17 +599,14 @@ describe("AzureAISearch – backward compat with mocked client", () => {
|
||||
createOrUpdateIndex: jest.fn().mockResolvedValue({}),
|
||||
deleteIndex: jest.fn().mockResolvedValue({}),
|
||||
})),
|
||||
AzureKeyCredential: jest
|
||||
.fn()
|
||||
.mockImplementation((key: string) => ({ key })),
|
||||
AzureKeyCredential: jest.fn().mockImplementation((key: string) => ({ key })),
|
||||
}));
|
||||
|
||||
jest.doMock("@azure/identity", () => ({
|
||||
DefaultAzureCredential: jest.fn(),
|
||||
}));
|
||||
|
||||
AzureAISearch =
|
||||
require("../src/vector_stores/azure_ai_search").AzureAISearch;
|
||||
AzureAISearch = require("../src/vector_stores/azure_ai_search").AzureAISearch;
|
||||
});
|
||||
|
||||
afterEach(() => {
|
||||
@@ -692,9 +661,7 @@ describe("Vectorize – backward compat with mocked client", () => {
|
||||
jest.doMock("cloudflare", () => {
|
||||
const mockIndexes = {
|
||||
list: jest.fn().mockReturnValue({
|
||||
[Symbol.asyncIterator]: () => ({
|
||||
next: async () => ({ done: true }),
|
||||
}),
|
||||
[Symbol.asyncIterator]: () => ({ next: async () => ({ done: true }) }),
|
||||
}),
|
||||
create: jest.fn().mockResolvedValue({}),
|
||||
delete: jest.fn().mockResolvedValue({}),
|
||||
@@ -805,14 +772,9 @@ describe("LangchainVectorStore – backward compat", () => {
|
||||
const { LangchainVectorStore } = require("../src/vector_stores/langchain");
|
||||
const mockLcStore = {
|
||||
addVectors: jest.fn().mockResolvedValue(undefined),
|
||||
similaritySearchVectorWithScore: jest
|
||||
.fn()
|
||||
.mockResolvedValue([
|
||||
[
|
||||
{ metadata: { _mem0_id: "id-1", data: "test" }, pageContent: "" },
|
||||
0.95,
|
||||
],
|
||||
]),
|
||||
similaritySearchVectorWithScore: jest.fn().mockResolvedValue([
|
||||
[{ metadata: { _mem0_id: "id-1", data: "test" }, pageContent: "" }, 0.95],
|
||||
]),
|
||||
};
|
||||
const store = new LangchainVectorStore({
|
||||
client: mockLcStore,
|
||||
@@ -859,9 +821,9 @@ describe("LangchainVectorStore – backward compat", () => {
|
||||
dimension: 4,
|
||||
});
|
||||
|
||||
await expect(store.insert([[1, 2, 3]], ["id-1"], [{}])).rejects.toThrow(
|
||||
"Vector dimension mismatch",
|
||||
);
|
||||
await expect(
|
||||
store.insert([[1, 2, 3]], ["id-1"], [{}]),
|
||||
).rejects.toThrow("Vector dimension mismatch");
|
||||
});
|
||||
});
|
||||
|
||||
@@ -1014,27 +976,11 @@ describe("Memory class – backward compat with all providers", () => {
|
||||
]);
|
||||
mockVStore.get.mockResolvedValue({
|
||||
id: "id-1",
|
||||
payload: {
|
||||
memory: "test",
|
||||
hash: "h",
|
||||
created_at: new Date().toISOString(),
|
||||
updated_at: new Date().toISOString(),
|
||||
},
|
||||
payload: { memory: "test", hash: "h", created_at: new Date().toISOString(), updated_at: new Date().toISOString() },
|
||||
});
|
||||
mockVStore.list.mockResolvedValue([
|
||||
[
|
||||
{
|
||||
id: "id-1",
|
||||
payload: {
|
||||
memory: "test",
|
||||
hash: "h",
|
||||
created_at: new Date().toISOString(),
|
||||
updated_at: new Date().toISOString(),
|
||||
},
|
||||
},
|
||||
],
|
||||
1,
|
||||
]);
|
||||
mockVStore.list.mockResolvedValue([[
|
||||
{ id: "id-1", payload: { memory: "test", hash: "h", created_at: new Date().toISOString(), updated_at: new Date().toISOString() } },
|
||||
], 1]);
|
||||
mockVectorStoreFactory.create.mockReturnValue(mockVStore);
|
||||
|
||||
const mem = new MemoryClass({
|
||||
@@ -1109,9 +1055,7 @@ describe("Memory class – backward compat with all providers", () => {
|
||||
};
|
||||
mockEmbedderFactory.create.mockReturnValue(failingEmbedder);
|
||||
|
||||
const consoleSpy = jest
|
||||
.spyOn(console, "error")
|
||||
.mockImplementation(() => {});
|
||||
const consoleSpy = jest.spyOn(console, "error").mockImplementation(() => {});
|
||||
|
||||
const mem = new MemoryClass({
|
||||
embedder: { provider: "ollama", config: { model: "test" } },
|
||||
@@ -1133,3 +1077,4 @@ describe("Memory class – backward compat with all providers", () => {
|
||||
consoleSpy.mockRestore();
|
||||
});
|
||||
});
|
||||
|
||||
|
||||
@@ -97,7 +97,7 @@ class NeptuneBase(ABC):
|
||||
for tool_call in search_results["tool_calls"]:
|
||||
if tool_call["name"] != "extract_entities":
|
||||
continue
|
||||
for item in tool_call.get("arguments", {}).get("entities", []):
|
||||
for item in tool_call["arguments"]["entities"]:
|
||||
entity_type_map[item["entity"]] = item["entity_type"]
|
||||
except Exception as e:
|
||||
logger.exception(
|
||||
@@ -144,7 +144,7 @@ class NeptuneBase(ABC):
|
||||
|
||||
entities = []
|
||||
if extracted_entities["tool_calls"]:
|
||||
entities = extracted_entities["tool_calls"][0].get("arguments", {}).get("entities", [])
|
||||
entities = extracted_entities["tool_calls"][0]["arguments"]["entities"]
|
||||
|
||||
entities = self._remove_spaces_from_entities(entities)
|
||||
logger.debug(f"Extracted entities: {entities}")
|
||||
|
||||
@@ -30,10 +30,10 @@ class MemoryGraph:
|
||||
def __init__(self, config):
|
||||
self.config = config
|
||||
self.graph = Neo4jGraph(
|
||||
url=self.config.graph_store.config.url,
|
||||
username=self.config.graph_store.config.username,
|
||||
password=self.config.graph_store.config.password,
|
||||
database=self.config.graph_store.config.database,
|
||||
self.config.graph_store.config.url,
|
||||
self.config.graph_store.config.username,
|
||||
self.config.graph_store.config.password,
|
||||
self.config.graph_store.config.database,
|
||||
refresh_schema=False,
|
||||
driver_config={"notifications_min_severity": "OFF"},
|
||||
)
|
||||
@@ -215,7 +215,7 @@ class MemoryGraph:
|
||||
for tool_call in search_results["tool_calls"]:
|
||||
if tool_call["name"] != "extract_entities":
|
||||
continue
|
||||
for item in tool_call.get("arguments", {}).get("entities", []):
|
||||
for item in tool_call["arguments"]["entities"]:
|
||||
entity_type_map[item["entity"]] = item["entity_type"]
|
||||
except Exception as e:
|
||||
logger.exception(
|
||||
|
||||
@@ -241,7 +241,7 @@ class MemoryGraph:
|
||||
for tool_call in search_results["tool_calls"]:
|
||||
if tool_call["name"] != "extract_entities":
|
||||
continue
|
||||
for item in tool_call.get("arguments", {}).get("entities", []):
|
||||
for item in tool_call["arguments"]["entities"]:
|
||||
entity_type_map[item["entity"]] = item["entity_type"]
|
||||
except Exception as e:
|
||||
logger.exception(
|
||||
|
||||
+35
-42
@@ -24,9 +24,8 @@ from mem0.exceptions import ValidationError as Mem0ValidationError
|
||||
from mem0.memory.base import MemoryBase
|
||||
from mem0.memory.setup import mem0_dir, setup_config
|
||||
from mem0.memory.storage import SQLiteManager
|
||||
from mem0.memory.telemetry import MEM0_TELEMETRY, capture_event
|
||||
from mem0.memory.telemetry import capture_event
|
||||
from mem0.memory.utils import (
|
||||
ensure_json_instruction,
|
||||
extract_json,
|
||||
get_fact_retrieval_messages,
|
||||
parse_messages,
|
||||
@@ -38,8 +37,8 @@ from mem0.utils.factory import (
|
||||
EmbedderFactory,
|
||||
GraphStoreFactory,
|
||||
LlmFactory,
|
||||
RerankerFactory,
|
||||
VectorStoreFactory,
|
||||
RerankerFactory,
|
||||
)
|
||||
|
||||
# Suppress SWIG deprecation warnings globally
|
||||
@@ -205,33 +204,32 @@ class Memory(MemoryBase):
|
||||
self.enable_graph = True
|
||||
else:
|
||||
self.graph = None
|
||||
if MEM0_TELEMETRY:
|
||||
# Create telemetry config manually to avoid deepcopy issues with thread locks
|
||||
telemetry_config_dict = {}
|
||||
if hasattr(self.config.vector_store.config, 'model_dump'):
|
||||
# For pydantic models
|
||||
telemetry_config_dict = self.config.vector_store.config.model_dump()
|
||||
else:
|
||||
# For other objects, manually copy common attributes
|
||||
for attr in ['host', 'port', 'path', 'api_key', 'index_name', 'dimension', 'metric']:
|
||||
if hasattr(self.config.vector_store.config, attr):
|
||||
telemetry_config_dict[attr] = getattr(self.config.vector_store.config, attr)
|
||||
# Create telemetry config manually to avoid deepcopy issues with thread locks
|
||||
telemetry_config_dict = {}
|
||||
if hasattr(self.config.vector_store.config, 'model_dump'):
|
||||
# For pydantic models
|
||||
telemetry_config_dict = self.config.vector_store.config.model_dump()
|
||||
else:
|
||||
# For other objects, manually copy common attributes
|
||||
for attr in ['host', 'port', 'path', 'api_key', 'index_name', 'dimension', 'metric']:
|
||||
if hasattr(self.config.vector_store.config, attr):
|
||||
telemetry_config_dict[attr] = getattr(self.config.vector_store.config, attr)
|
||||
|
||||
# Override collection name for telemetry
|
||||
telemetry_config_dict['collection_name'] = "mem0migrations"
|
||||
# Override collection name for telemetry
|
||||
telemetry_config_dict['collection_name'] = "mem0migrations"
|
||||
|
||||
# Set path for file-based vector stores
|
||||
telemetry_config = _safe_deepcopy_config(self.config.vector_store.config)
|
||||
if self.config.vector_store.provider in ["faiss", "qdrant"]:
|
||||
provider_path = f"migrations_{self.config.vector_store.provider}"
|
||||
telemetry_config_dict['path'] = os.path.join(mem0_dir, provider_path)
|
||||
os.makedirs(telemetry_config_dict['path'], exist_ok=True)
|
||||
# Set path for file-based vector stores
|
||||
telemetry_config = _safe_deepcopy_config(self.config.vector_store.config)
|
||||
if self.config.vector_store.provider in ["faiss", "qdrant"]:
|
||||
provider_path = f"migrations_{self.config.vector_store.provider}"
|
||||
telemetry_config_dict['path'] = os.path.join(mem0_dir, provider_path)
|
||||
os.makedirs(telemetry_config_dict['path'], exist_ok=True)
|
||||
|
||||
# Create the config object using the same class as the original
|
||||
telemetry_config = self.config.vector_store.config.__class__(**telemetry_config_dict)
|
||||
self._telemetry_vector_store = VectorStoreFactory.create(
|
||||
self.config.vector_store.provider, telemetry_config
|
||||
)
|
||||
# Create the config object using the same class as the original
|
||||
telemetry_config = self.config.vector_store.config.__class__(**telemetry_config_dict)
|
||||
self._telemetry_vector_store = VectorStoreFactory.create(
|
||||
self.config.vector_store.provider, telemetry_config
|
||||
)
|
||||
capture_event("mem0.init", self, {"sync_type": "sync"})
|
||||
|
||||
@classmethod
|
||||
@@ -433,9 +431,6 @@ class Memory(MemoryBase):
|
||||
is_agent_memory = self._should_use_agent_memory_extraction(messages, metadata)
|
||||
system_prompt, user_prompt = get_fact_retrieval_messages(parsed_messages, is_agent_memory)
|
||||
|
||||
# Ensure 'json' appears in prompts for json_object response format compatibility
|
||||
system_prompt, user_prompt = ensure_json_instruction(system_prompt, user_prompt)
|
||||
|
||||
response = self.llm.generate_response(
|
||||
messages=[
|
||||
{"role": "system", "content": system_prompt},
|
||||
@@ -1051,10 +1046,11 @@ class Memory(MemoryBase):
|
||||
|
||||
keys, encoded_ids = process_telemetry_filters(filters)
|
||||
capture_event("mem0.delete_all", self, {"keys": keys, "encoded_ids": encoded_ids, "sync_type": "sync"})
|
||||
# delete matching vector memories individually (do NOT reset the collection)
|
||||
# delete all vector memories and reset the collections
|
||||
memories = self.vector_store.list(filters=filters)[0]
|
||||
for memory in memories:
|
||||
self._delete_memory(memory.id)
|
||||
self.vector_store.reset()
|
||||
|
||||
logger.info(f"Deleted {len(memories)} memories")
|
||||
|
||||
@@ -1276,14 +1272,14 @@ class AsyncMemory(MemoryBase):
|
||||
else:
|
||||
self.graph = None
|
||||
|
||||
if MEM0_TELEMETRY:
|
||||
telemetry_config = _safe_deepcopy_config(self.config.vector_store.config)
|
||||
telemetry_config.collection_name = "mem0migrations"
|
||||
if self.config.vector_store.provider in ["faiss", "qdrant"]:
|
||||
provider_path = f"migrations_{self.config.vector_store.provider}"
|
||||
telemetry_config.path = os.path.join(mem0_dir, provider_path)
|
||||
os.makedirs(telemetry_config.path, exist_ok=True)
|
||||
self._telemetry_vector_store = VectorStoreFactory.create(self.config.vector_store.provider, telemetry_config)
|
||||
telemetry_config = _safe_deepcopy_config(self.config.vector_store.config)
|
||||
telemetry_config.collection_name = "mem0migrations"
|
||||
if self.config.vector_store.provider in ["faiss", "qdrant"]:
|
||||
provider_path = f"migrations_{self.config.vector_store.provider}"
|
||||
telemetry_config.path = os.path.join(mem0_dir, provider_path)
|
||||
os.makedirs(telemetry_config.path, exist_ok=True)
|
||||
self._telemetry_vector_store = VectorStoreFactory.create(self.config.vector_store.provider, telemetry_config)
|
||||
|
||||
capture_event("mem0.init", self, {"sync_type": "async"})
|
||||
|
||||
@classmethod
|
||||
@@ -1464,9 +1460,6 @@ class AsyncMemory(MemoryBase):
|
||||
is_agent_memory = self._should_use_agent_memory_extraction(messages, metadata)
|
||||
system_prompt, user_prompt = get_fact_retrieval_messages(parsed_messages, is_agent_memory)
|
||||
|
||||
# Ensure 'json' appears in prompts for json_object response format compatibility
|
||||
system_prompt, user_prompt = ensure_json_instruction(system_prompt, user_prompt)
|
||||
|
||||
response = await asyncio.to_thread(
|
||||
self.llm.generate_response,
|
||||
messages=[{"role": "system", "content": system_prompt}, {"role": "user", "content": user_prompt}],
|
||||
|
||||
@@ -218,7 +218,7 @@ class MemoryGraph:
|
||||
for tool_call in search_results["tool_calls"]:
|
||||
if tool_call["name"] != "extract_entities":
|
||||
continue
|
||||
for item in tool_call.get("arguments", {}).get("entities", []):
|
||||
for item in tool_call["arguments"]["entities"]:
|
||||
if "entity" in item and "entity_type" in item:
|
||||
entity_type_map[item["entity"]] = item["entity_type"]
|
||||
except Exception as e:
|
||||
|
||||
@@ -29,31 +29,6 @@ def get_fact_retrieval_messages_legacy(message):
|
||||
return FACT_RETRIEVAL_PROMPT, f"Input:\n{message}"
|
||||
|
||||
|
||||
def ensure_json_instruction(system_prompt, user_prompt):
|
||||
"""Ensure the word 'json' appears in the prompts when using json_object response format.
|
||||
|
||||
OpenAI's API requires the word 'json' to appear in the messages when
|
||||
response_format is set to {"type": "json_object"}. When users provide a
|
||||
custom_fact_extraction_prompt that doesn't include 'json', this causes a
|
||||
400 error. This function appends a JSON format instruction to the system
|
||||
prompt if 'json' is not already present in either prompt.
|
||||
|
||||
Args:
|
||||
system_prompt: The system prompt string
|
||||
user_prompt: The user prompt string
|
||||
|
||||
Returns:
|
||||
tuple: (system_prompt, user_prompt) with JSON instruction added if needed
|
||||
"""
|
||||
combined = (system_prompt + user_prompt).lower()
|
||||
if "json" not in combined:
|
||||
system_prompt += (
|
||||
"\n\nYou must return your response in valid JSON format "
|
||||
"with a 'facts' key containing an array of strings."
|
||||
)
|
||||
return system_prompt, user_prompt
|
||||
|
||||
|
||||
def parse_messages(messages):
|
||||
response = ""
|
||||
for msg in messages:
|
||||
|
||||
@@ -192,12 +192,9 @@ class RedisDB(VectorStoreBase):
|
||||
"memory": payload["data"],
|
||||
"created_at": int(datetime.fromisoformat(payload["created_at"]).timestamp()),
|
||||
"updated_at": int(datetime.fromisoformat(payload["updated_at"]).timestamp()),
|
||||
"embedding": np.array(vector, dtype=np.float32).tobytes(),
|
||||
}
|
||||
|
||||
# Only update embedding if vector is provided
|
||||
if vector is not None:
|
||||
data["embedding"] = np.array(vector, dtype=np.float32).tobytes()
|
||||
|
||||
for field in ["agent_id", "run_id", "user_id"]:
|
||||
if field in payload:
|
||||
data[field] = payload[field]
|
||||
|
||||
@@ -491,12 +491,9 @@ class ValkeyDB(VectorStoreBase):
|
||||
"hash": payload.get("hash", f"hash_{vector_id}"), # Use a default hash if not provided
|
||||
"memory": payload.get("data", f"data_{vector_id}"), # Use a default data if not provided
|
||||
"created_at": int(datetime.fromisoformat(payload["created_at"]).timestamp()),
|
||||
"embedding": np.array(vector, dtype=np.float32).tobytes(),
|
||||
}
|
||||
|
||||
# Only update embedding if vector is provided
|
||||
if vector is not None:
|
||||
hash_data["embedding"] = np.array(vector, dtype=np.float32).tobytes()
|
||||
|
||||
# Add updated_at if available
|
||||
if "updated_at" in payload:
|
||||
hash_data["updated_at"] = int(datetime.fromisoformat(payload["updated_at"]).timestamp())
|
||||
|
||||
@@ -1,2 +0,0 @@
|
||||
package-manager-strict-version=false
|
||||
approve-builds=esbuild
|
||||
+6
-34
@@ -41,7 +41,6 @@ type Mem0Config = {
|
||||
vectorStore?: { provider: string; config: Record<string, unknown> };
|
||||
llm?: { provider: string; config: Record<string, unknown> };
|
||||
historyDbPath?: string;
|
||||
disableHistory?: boolean;
|
||||
};
|
||||
// Shared
|
||||
userId: string;
|
||||
@@ -133,16 +132,13 @@ class PlatformProvider implements Mem0Provider {
|
||||
private async ensureClient(): Promise<void> {
|
||||
if (this.client) return;
|
||||
if (this.initPromise) return this.initPromise;
|
||||
this.initPromise = this._init().catch((err) => {
|
||||
this.initPromise = null;
|
||||
throw err;
|
||||
});
|
||||
this.initPromise = this._init();
|
||||
return this.initPromise;
|
||||
}
|
||||
|
||||
private async _init(): Promise<void> {
|
||||
const { default: MemoryClient } = await import("mem0ai");
|
||||
const opts: { apiKey: string; org_id?: string; project_id?: string } = { apiKey: this.apiKey };
|
||||
const opts: Record<string, string> = { apiKey: this.apiKey };
|
||||
if (this.orgId) opts.org_id = this.orgId;
|
||||
if (this.projectId) opts.project_id = this.projectId;
|
||||
this.client = new MemoryClient(opts);
|
||||
@@ -229,10 +225,7 @@ class OSSProvider implements Mem0Provider {
|
||||
private async ensureMemory(): Promise<void> {
|
||||
if (this.memory) return;
|
||||
if (this.initPromise) return this.initPromise;
|
||||
this.initPromise = this._init().catch((err) => {
|
||||
this.initPromise = null;
|
||||
throw err;
|
||||
});
|
||||
this.initPromise = this._init();
|
||||
return this.initPromise;
|
||||
}
|
||||
|
||||
@@ -253,30 +246,9 @@ class OSSProvider implements Mem0Provider {
|
||||
config.historyDbPath = dbPath;
|
||||
}
|
||||
|
||||
if (this.ossConfig?.disableHistory) {
|
||||
config.disableHistory = true;
|
||||
}
|
||||
|
||||
if (this.customPrompt) config.customPrompt = this.customPrompt;
|
||||
|
||||
try {
|
||||
this.memory = new Memory(config);
|
||||
} catch (err) {
|
||||
// If initialization fails (e.g. native SQLite binding resolution under
|
||||
// jiti), retry with history disabled — the history DB is the most common
|
||||
// source of native-binding failures and is not required for core
|
||||
// memory operations.
|
||||
if (!config.disableHistory) {
|
||||
console.warn(
|
||||
"[mem0] Memory initialization failed, retrying with history disabled:",
|
||||
err instanceof Error ? err.message : err,
|
||||
);
|
||||
config.disableHistory = true;
|
||||
this.memory = new Memory(config);
|
||||
} else {
|
||||
throw err;
|
||||
}
|
||||
}
|
||||
this.memory = new Memory(config);
|
||||
}
|
||||
|
||||
async add(
|
||||
@@ -549,7 +521,7 @@ function assertAllowedKeys(
|
||||
throw new Error(`${label} has unknown keys: ${unknown.join(", ")}`);
|
||||
}
|
||||
|
||||
export const mem0ConfigSchema = {
|
||||
const mem0ConfigSchema = {
|
||||
parse(value: unknown): Mem0Config {
|
||||
if (!value || typeof value !== "object" || Array.isArray(value)) {
|
||||
throw new Error("openclaw-mem0 config required");
|
||||
@@ -615,7 +587,7 @@ export const mem0ConfigSchema = {
|
||||
// Provider Factory
|
||||
// ============================================================================
|
||||
|
||||
export function createProvider(
|
||||
function createProvider(
|
||||
cfg: Mem0Config,
|
||||
api: OpenClawPluginApi,
|
||||
): Mem0Provider {
|
||||
|
||||
Vendored
-30
@@ -1,30 +0,0 @@
|
||||
declare module "openclaw/plugin-sdk" {
|
||||
export interface OpenClawPluginApi {
|
||||
pluginConfig: Record<string, unknown>;
|
||||
logger: {
|
||||
info(msg: string): void;
|
||||
warn(msg: string): void;
|
||||
error(msg: string): void;
|
||||
debug(msg: string): void;
|
||||
};
|
||||
resolvePath(p: string): string;
|
||||
registerTool(
|
||||
definition: Record<string, unknown>,
|
||||
metadata?: Record<string, unknown>,
|
||||
): void;
|
||||
on(
|
||||
event: string,
|
||||
handler: (event: any, ctx: any) => any,
|
||||
): void;
|
||||
registerCli(
|
||||
handler: (context: { program: any }) => void,
|
||||
options?: Record<string, unknown>,
|
||||
): void;
|
||||
registerService(service: {
|
||||
id: string;
|
||||
start: () => void;
|
||||
stop: () => void;
|
||||
}): void;
|
||||
[key: string]: unknown;
|
||||
}
|
||||
}
|
||||
+2
-18
@@ -1,6 +1,6 @@
|
||||
{
|
||||
"name": "@mem0/openclaw-mem0",
|
||||
"version": "0.3.3",
|
||||
"version": "0.3.1",
|
||||
"type": "module",
|
||||
"description": "Mem0 memory backend for OpenClaw — platform or self-hosted open-source",
|
||||
"license": "Apache-2.0",
|
||||
@@ -11,20 +11,7 @@
|
||||
"mem0",
|
||||
"long-term-memory"
|
||||
],
|
||||
"main": "./dist/index.js",
|
||||
"types": "./dist/index.d.ts",
|
||||
"exports": {
|
||||
".": {
|
||||
"types": "./dist/index.d.ts",
|
||||
"import": "./dist/index.js"
|
||||
}
|
||||
},
|
||||
"files": [
|
||||
"dist",
|
||||
"openclaw.plugin.json"
|
||||
],
|
||||
"scripts": {
|
||||
"build": "tsup",
|
||||
"test": "vitest run"
|
||||
},
|
||||
"dependencies": {
|
||||
@@ -33,13 +20,10 @@
|
||||
},
|
||||
"openclaw": {
|
||||
"extensions": [
|
||||
"./dist/index.js"
|
||||
"./index.ts"
|
||||
]
|
||||
},
|
||||
"devDependencies": {
|
||||
"@types/node": "^22.15.0",
|
||||
"tsup": "^8.5.0",
|
||||
"typescript": "^5.8.3",
|
||||
"vitest": "^4.0.18"
|
||||
}
|
||||
}
|
||||
|
||||
Generated
-4105
File diff suppressed because it is too large
Load Diff
@@ -1,6 +0,0 @@
|
||||
approveBuilds: esbuild
|
||||
|
||||
onlyBuiltDependencies:
|
||||
- better-sqlite3
|
||||
- esbuild
|
||||
- protobufjs
|
||||
@@ -1,288 +0,0 @@
|
||||
/**
|
||||
* Tests for SQLite resilience fixes:
|
||||
* 1. disableHistory config passthrough
|
||||
* 2. initPromise poisoning fix (retry after failure)
|
||||
* 3. Graceful SQLite fallback in OSSProvider
|
||||
*/
|
||||
import { describe, it, expect, vi, beforeEach } from "vitest";
|
||||
import { mem0ConfigSchema, createProvider } from "./index.ts";
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// 1. Config: disableHistory passthrough
|
||||
// ---------------------------------------------------------------------------
|
||||
describe("mem0ConfigSchema — disableHistory", () => {
|
||||
const baseConfig = {
|
||||
mode: "open-source",
|
||||
oss: {
|
||||
embedder: { provider: "openai", config: { apiKey: "sk-test" } },
|
||||
},
|
||||
};
|
||||
|
||||
it("preserves oss.disableHistory: true through config parsing", () => {
|
||||
const cfg = mem0ConfigSchema.parse({
|
||||
...baseConfig,
|
||||
oss: { ...baseConfig.oss, disableHistory: true },
|
||||
});
|
||||
expect(cfg.oss?.disableHistory).toBe(true);
|
||||
});
|
||||
|
||||
it("preserves oss.disableHistory: false through config parsing", () => {
|
||||
const cfg = mem0ConfigSchema.parse({
|
||||
...baseConfig,
|
||||
oss: { ...baseConfig.oss, disableHistory: false },
|
||||
});
|
||||
expect(cfg.oss?.disableHistory).toBe(false);
|
||||
});
|
||||
|
||||
it("omits disableHistory when not provided", () => {
|
||||
const cfg = mem0ConfigSchema.parse(baseConfig);
|
||||
expect(cfg.oss?.disableHistory).toBeUndefined();
|
||||
});
|
||||
|
||||
it("does not reject unknown keys inside oss object", () => {
|
||||
// oss sub-object is passed through resolveEnvVarsDeep, not key-checked
|
||||
expect(() =>
|
||||
mem0ConfigSchema.parse({
|
||||
...baseConfig,
|
||||
oss: { ...baseConfig.oss, disableHistory: true },
|
||||
}),
|
||||
).not.toThrow();
|
||||
});
|
||||
});
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// 2. OSSProvider: disableHistory flows to Memory constructor
|
||||
// ---------------------------------------------------------------------------
|
||||
describe("OSSProvider — disableHistory passthrough to Memory", () => {
|
||||
let capturedConfig: Record<string, unknown> | undefined;
|
||||
let memoryCallCount: number;
|
||||
|
||||
beforeEach(() => {
|
||||
capturedConfig = undefined;
|
||||
memoryCallCount = 0;
|
||||
|
||||
vi.doMock("mem0ai/oss", () => ({
|
||||
Memory: class MockMemory {
|
||||
constructor(config: Record<string, unknown>) {
|
||||
memoryCallCount++;
|
||||
capturedConfig = { ...config };
|
||||
}
|
||||
async add() { return { results: [] }; }
|
||||
async search() { return { results: [] }; }
|
||||
async get() { return {}; }
|
||||
async getAll() { return []; }
|
||||
async delete() { }
|
||||
},
|
||||
}));
|
||||
});
|
||||
|
||||
it("passes disableHistory: true to Memory when configured", async () => {
|
||||
const { createProvider } = await import("./index.ts");
|
||||
const cfg = mem0ConfigSchema.parse({
|
||||
mode: "open-source",
|
||||
oss: { disableHistory: true },
|
||||
});
|
||||
const api = { resolvePath: (p: string) => p } as any;
|
||||
const provider = createProvider(cfg, api);
|
||||
|
||||
// Trigger lazy init by calling search
|
||||
try {
|
||||
await provider.search("test", { user_id: "u1" });
|
||||
} catch { /* provider may fail on mock, that's ok */ }
|
||||
|
||||
expect(capturedConfig).toBeDefined();
|
||||
expect(capturedConfig!.disableHistory).toBe(true);
|
||||
});
|
||||
|
||||
it("does not set disableHistory when not configured", async () => {
|
||||
const { createProvider } = await import("./index.ts");
|
||||
const cfg = mem0ConfigSchema.parse({
|
||||
mode: "open-source",
|
||||
oss: {},
|
||||
});
|
||||
const api = { resolvePath: (p: string) => p } as any;
|
||||
const provider = createProvider(cfg, api);
|
||||
|
||||
try {
|
||||
await provider.search("test", { user_id: "u1" });
|
||||
} catch { }
|
||||
|
||||
expect(capturedConfig).toBeDefined();
|
||||
expect(capturedConfig!.disableHistory).toBeUndefined();
|
||||
});
|
||||
});
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// 3. OSSProvider: initPromise is cleared on failure (allows retry)
|
||||
// ---------------------------------------------------------------------------
|
||||
describe("OSSProvider — initPromise retry after failure", () => {
|
||||
let callCount: number;
|
||||
|
||||
beforeEach(() => {
|
||||
callCount = 0;
|
||||
|
||||
vi.doMock("mem0ai/oss", () => ({
|
||||
Memory: class MockMemory {
|
||||
constructor() {
|
||||
callCount++;
|
||||
if (callCount === 1) {
|
||||
throw new Error("SQLITE_CANTOPEN: simulated binding failure");
|
||||
}
|
||||
// Second+ call succeeds
|
||||
}
|
||||
async search() { return { results: [] }; }
|
||||
async get() { return {}; }
|
||||
async getAll() { return []; }
|
||||
async add() { return { results: [] }; }
|
||||
async delete() { }
|
||||
},
|
||||
}));
|
||||
});
|
||||
|
||||
it("retries initialization after a transient failure", async () => {
|
||||
const { createProvider } = await import("./index.ts");
|
||||
const cfg = mem0ConfigSchema.parse({
|
||||
mode: "open-source",
|
||||
oss: { disableHistory: true },
|
||||
});
|
||||
const api = { resolvePath: (p: string) => p } as any;
|
||||
const provider = createProvider(cfg, api);
|
||||
|
||||
// First call: _init throws, but initPromise is cleared so retry is possible
|
||||
await expect(
|
||||
provider.search("test", { user_id: "u1" }),
|
||||
).rejects.toThrow("SQLITE_CANTOPEN");
|
||||
|
||||
// Second call: should retry _init (not return cached rejection)
|
||||
// callCount === 1 threw, so callCount === 2 should succeed
|
||||
const results = await provider.search("test", { user_id: "u1" });
|
||||
expect(results).toBeDefined();
|
||||
expect(callCount).toBe(2);
|
||||
});
|
||||
});
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// 4. OSSProvider: graceful fallback disables history on init failure
|
||||
// ---------------------------------------------------------------------------
|
||||
describe("OSSProvider — graceful SQLite fallback", () => {
|
||||
let capturedConfigs: Record<string, unknown>[];
|
||||
|
||||
beforeEach(() => {
|
||||
capturedConfigs = [];
|
||||
|
||||
vi.doMock("mem0ai/oss", () => ({
|
||||
Memory: class MockMemory {
|
||||
constructor(config: Record<string, unknown>) {
|
||||
capturedConfigs.push({ ...config });
|
||||
if (!config.disableHistory) {
|
||||
throw new Error("Could not locate the bindings file");
|
||||
}
|
||||
// Succeeds when disableHistory is true
|
||||
}
|
||||
async search() { return { results: [] }; }
|
||||
async get() { return {}; }
|
||||
async getAll() { return []; }
|
||||
async add() { return { results: [] }; }
|
||||
async delete() { }
|
||||
},
|
||||
}));
|
||||
});
|
||||
|
||||
it("retries with disableHistory: true when initial construction fails", async () => {
|
||||
const warnSpy = vi.spyOn(console, "warn").mockImplementation(() => {});
|
||||
const { createProvider } = await import("./index.ts");
|
||||
const cfg = mem0ConfigSchema.parse({
|
||||
mode: "open-source",
|
||||
oss: {},
|
||||
});
|
||||
const api = { resolvePath: (p: string) => p } as any;
|
||||
const provider = createProvider(cfg, api);
|
||||
|
||||
// Should succeed — first attempt fails, fallback with disableHistory succeeds
|
||||
const results = await provider.search("test", { user_id: "u1" });
|
||||
expect(results).toBeDefined();
|
||||
|
||||
// Memory constructor was called twice
|
||||
expect(capturedConfigs).toHaveLength(2);
|
||||
expect(capturedConfigs[0].disableHistory).toBeFalsy();
|
||||
expect(capturedConfigs[1].disableHistory).toBe(true);
|
||||
|
||||
// Warning was logged
|
||||
expect(warnSpy).toHaveBeenCalledWith(
|
||||
expect.stringContaining("[mem0] Memory initialization failed"),
|
||||
expect.stringContaining("bindings file"),
|
||||
);
|
||||
warnSpy.mockRestore();
|
||||
});
|
||||
|
||||
it("does not retry when disableHistory is already true", async () => {
|
||||
vi.doMock("mem0ai/oss", () => ({
|
||||
Memory: class MockMemory {
|
||||
constructor(config: Record<string, unknown>) {
|
||||
// Fail even with disableHistory (e.g. vector store issue)
|
||||
throw new Error("vector store connection refused");
|
||||
}
|
||||
},
|
||||
}));
|
||||
|
||||
const { createProvider } = await import("./index.ts");
|
||||
const cfg = mem0ConfigSchema.parse({
|
||||
mode: "open-source",
|
||||
oss: { disableHistory: true },
|
||||
});
|
||||
const api = { resolvePath: (p: string) => p } as any;
|
||||
const provider = createProvider(cfg, api);
|
||||
|
||||
// Should throw — no fallback possible when disableHistory was already set
|
||||
await expect(
|
||||
provider.search("test", { user_id: "u1" }),
|
||||
).rejects.toThrow("vector store connection refused");
|
||||
});
|
||||
});
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// 5. PlatformProvider — initPromise retry after failure
|
||||
// ---------------------------------------------------------------------------
|
||||
describe("PlatformProvider — initPromise retry after failure", () => {
|
||||
let callCount: number;
|
||||
|
||||
beforeEach(() => {
|
||||
callCount = 0;
|
||||
|
||||
vi.doMock("mem0ai", () => ({
|
||||
default: class MockMemoryClient {
|
||||
constructor() {
|
||||
callCount++;
|
||||
if (callCount === 1) {
|
||||
throw new Error("Network timeout");
|
||||
}
|
||||
}
|
||||
async search() { return []; }
|
||||
async get() { return {}; }
|
||||
async getAll() { return []; }
|
||||
async add() { return { results: [] }; }
|
||||
async delete() { }
|
||||
},
|
||||
}));
|
||||
});
|
||||
|
||||
it("retries initialization after a transient failure", async () => {
|
||||
const { createProvider } = await import("./index.ts");
|
||||
const cfg = mem0ConfigSchema.parse({
|
||||
mode: "platform",
|
||||
apiKey: "test-api-key",
|
||||
});
|
||||
const api = { resolvePath: (p: string) => p } as any;
|
||||
const provider = createProvider(cfg, api);
|
||||
|
||||
// First call fails
|
||||
await expect(
|
||||
provider.search("test", { user_id: "u1" }),
|
||||
).rejects.toThrow("Network timeout");
|
||||
|
||||
// Second call should retry (not return cached rejection)
|
||||
const results = await provider.search("test", { user_id: "u1" });
|
||||
expect(results).toBeDefined();
|
||||
expect(callCount).toBe(2);
|
||||
});
|
||||
});
|
||||
@@ -1,22 +0,0 @@
|
||||
{
|
||||
"compilerOptions": {
|
||||
"target": "ES2022",
|
||||
"module": "ES2022",
|
||||
"moduleResolution": "bundler",
|
||||
"declaration": true,
|
||||
"declarationMap": true,
|
||||
"sourceMap": true,
|
||||
"outDir": "dist",
|
||||
"rootDir": ".",
|
||||
"strict": false,
|
||||
"noImplicitAny": false,
|
||||
"types": ["node"],
|
||||
"esModuleInterop": true,
|
||||
"skipLibCheck": true,
|
||||
"forceConsistentCasingInFileNames": true,
|
||||
"isolatedModules": true,
|
||||
"verbatimModuleSyntax": true
|
||||
},
|
||||
"include": ["index.ts", "openclaw-plugin-sdk.d.ts"],
|
||||
"exclude": ["node_modules", "dist", "**/*.test.ts"]
|
||||
}
|
||||
@@ -1,9 +0,0 @@
|
||||
import { defineConfig } from "tsup";
|
||||
|
||||
export default defineConfig({
|
||||
entry: ["index.ts"],
|
||||
format: ["esm"],
|
||||
dts: true,
|
||||
sourcemap: true,
|
||||
clean: true,
|
||||
});
|
||||
+1
-1
@@ -20,7 +20,7 @@ dependencies = [
|
||||
"posthog>=3.5.0",
|
||||
"pytz>=2024.1",
|
||||
"sqlalchemy>=2.0.31",
|
||||
"protobuf>=5.29.6,<7.0.0",
|
||||
"protobuf>=5.29.0,<6.0.0",
|
||||
]
|
||||
|
||||
[project.optional-dependencies]
|
||||
|
||||
@@ -13,12 +13,6 @@ When installed, Claude can:
|
||||
|
||||
## Installation
|
||||
|
||||
### CLI (Claude Code, OpenCode, OpenClaw, or any tool that supports skills)
|
||||
|
||||
```bash
|
||||
npx skills add https://github.com/mem0ai/mem0 --skill mem0
|
||||
```
|
||||
|
||||
### Claude.ai
|
||||
|
||||
1. Download this `skills/mem0` folder as a ZIP
|
||||
|
||||
@@ -1,214 +0,0 @@
|
||||
"""
|
||||
Tests for issue #3559: Custom prompts crash with response_format json_object
|
||||
when the word 'json' is not present in the prompt.
|
||||
|
||||
OpenAI API requires the word 'json' to appear in messages when using
|
||||
response_format: {"type": "json_object"}. Custom fact extraction prompts
|
||||
may not include this word, causing BadRequestError.
|
||||
|
||||
This tests the ensure_json_instruction utility function and verifies
|
||||
the fix is applied in both sync and async code paths.
|
||||
"""
|
||||
|
||||
import pytest
|
||||
|
||||
from mem0.memory.utils import ensure_json_instruction
|
||||
|
||||
|
||||
class TestEnsureJsonInstruction:
|
||||
"""Tests for the ensure_json_instruction utility function."""
|
||||
|
||||
# -------------------------------------------------------------------
|
||||
# Core behavior: append when missing, skip when present
|
||||
# -------------------------------------------------------------------
|
||||
|
||||
def test_appends_when_json_missing_from_both_prompts(self):
|
||||
"""When neither prompt contains 'json', instruction is appended to system prompt."""
|
||||
system, user = ensure_json_instruction(
|
||||
"Extract facts from the conversation and return them as a list.",
|
||||
"Input:\nuser: Hi my name is John",
|
||||
)
|
||||
assert "json" in system.lower()
|
||||
assert "facts" in system.lower()
|
||||
|
||||
def test_no_change_when_json_in_system_prompt(self):
|
||||
"""When system prompt already contains 'json', no modification."""
|
||||
original = "Extract facts and return in json format."
|
||||
system, user = ensure_json_instruction(original, "Input:\nuser: Hi")
|
||||
assert system == original
|
||||
|
||||
def test_no_change_when_json_in_user_prompt(self):
|
||||
"""When user prompt contains 'json', no modification to system prompt."""
|
||||
original_system = "Extract facts from the conversation."
|
||||
original_user = "Input (respond in json):\nuser: Hi"
|
||||
system, user = ensure_json_instruction(original_system, original_user)
|
||||
assert system == original_system
|
||||
|
||||
def test_user_prompt_never_modified(self):
|
||||
"""The user prompt should never be modified regardless of content."""
|
||||
original_user = "Input:\nuser: I like pizza"
|
||||
_, user = ensure_json_instruction("Extract facts.", original_user)
|
||||
assert user == original_user
|
||||
|
||||
# -------------------------------------------------------------------
|
||||
# Case insensitivity
|
||||
# -------------------------------------------------------------------
|
||||
|
||||
def test_case_insensitive_lowercase(self):
|
||||
original = "Return results in json format."
|
||||
system, _ = ensure_json_instruction(original, "Input:\nuser: Hi")
|
||||
assert system == original
|
||||
|
||||
def test_case_insensitive_uppercase(self):
|
||||
original = "Return results in JSON format."
|
||||
system, _ = ensure_json_instruction(original, "Input:\nuser: Hi")
|
||||
assert system == original
|
||||
|
||||
def test_case_insensitive_mixed(self):
|
||||
original = "Return results in Json format."
|
||||
system, _ = ensure_json_instruction(original, "Input:\nuser: Hi")
|
||||
assert system == original
|
||||
|
||||
def test_case_insensitive_in_user_prompt(self):
|
||||
original_system = "Extract facts."
|
||||
system, _ = ensure_json_instruction(original_system, "Return JSON.\nuser: Hi")
|
||||
assert system == original_system
|
||||
|
||||
# -------------------------------------------------------------------
|
||||
# Parametrized: various custom prompts
|
||||
# -------------------------------------------------------------------
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"prompt,should_append",
|
||||
[
|
||||
# Prompts WITHOUT json — should append
|
||||
("Extract all facts from the conversation.", True),
|
||||
("You are a memory extractor. Return facts as a list.", True),
|
||||
("Analyze the input and find key information.", True),
|
||||
("Return data in structured format.", True),
|
||||
("List the user preferences.", True),
|
||||
# Prompts WITH json — should NOT append
|
||||
("Extract facts and return in json format.", False),
|
||||
("Return a json object with facts.", False),
|
||||
("Output must be valid JSON.", False),
|
||||
("Respond with a JSON array of facts.", False),
|
||||
("Format: json output expected.", False),
|
||||
],
|
||||
)
|
||||
def test_various_custom_prompts(self, prompt, should_append):
|
||||
user_prompt = "Input:\nuser: Hi my name is John"
|
||||
system, _ = ensure_json_instruction(prompt, user_prompt)
|
||||
|
||||
if should_append:
|
||||
assert system != prompt, f"Expected JSON instruction to be appended for: {prompt}"
|
||||
assert "json" in system.lower()
|
||||
else:
|
||||
assert system == prompt, f"Did not expect modification for: {prompt}"
|
||||
|
||||
# -------------------------------------------------------------------
|
||||
# Edge cases
|
||||
# -------------------------------------------------------------------
|
||||
|
||||
def test_empty_system_prompt(self):
|
||||
"""Empty system prompt should get JSON instruction."""
|
||||
system, _ = ensure_json_instruction("", "Input:\nuser: test")
|
||||
assert "json" in system.lower()
|
||||
|
||||
def test_whitespace_only_system_prompt(self):
|
||||
"""Whitespace-only prompt should get JSON instruction."""
|
||||
system, _ = ensure_json_instruction(" \n ", "Input:\nuser: test")
|
||||
assert "json" in system.lower()
|
||||
|
||||
def test_preserves_original_prompt_content(self):
|
||||
"""The fix should only append, never modify the original prompt content."""
|
||||
original = "Extract all user preferences and habits from the conversation."
|
||||
system, _ = ensure_json_instruction(original, "Input:\nuser: I like pizza")
|
||||
assert system.startswith(original)
|
||||
assert len(system) > len(original)
|
||||
|
||||
def test_appended_instruction_mentions_facts_key(self):
|
||||
"""The appended instruction should guide the model to use the 'facts' key."""
|
||||
system, _ = ensure_json_instruction(
|
||||
"Extract information.", "Input:\nuser: test"
|
||||
)
|
||||
assert "facts" in system.lower()
|
||||
|
||||
def test_idempotent_when_already_has_json(self):
|
||||
"""Calling ensure_json_instruction twice doesn't double-append."""
|
||||
system1, user1 = ensure_json_instruction(
|
||||
"Extract facts.", "Input:\nuser: test"
|
||||
)
|
||||
system2, user2 = ensure_json_instruction(system1, user1)
|
||||
assert system1 == system2
|
||||
assert user1 == user2
|
||||
|
||||
def test_json_in_curly_braces_not_detected(self):
|
||||
"""A prompt with JSON-like structure but no 'json' word should get instruction.
|
||||
e.g. '{"facts": [...]}' contains the characters j,s,o,n but not the word 'json'."""
|
||||
prompt = 'Return format: {"facts": [...]}'
|
||||
# This contains the substring "json" inside the key name — let's check
|
||||
if "json" in prompt.lower():
|
||||
# If it does contain json, it won't be modified
|
||||
system, _ = ensure_json_instruction(prompt, "Input:\nuser: test")
|
||||
assert system == prompt
|
||||
else:
|
||||
system, _ = ensure_json_instruction(prompt, "Input:\nuser: test")
|
||||
assert system != prompt
|
||||
|
||||
# -------------------------------------------------------------------
|
||||
# Default prompts verification
|
||||
# -------------------------------------------------------------------
|
||||
|
||||
def test_default_prompts_already_contain_json(self):
|
||||
"""Built-in prompts already contain 'json', so ensure_json_instruction is a no-op."""
|
||||
from mem0.configs.prompts import (
|
||||
FACT_RETRIEVAL_PROMPT,
|
||||
USER_MEMORY_EXTRACTION_PROMPT,
|
||||
AGENT_MEMORY_EXTRACTION_PROMPT,
|
||||
)
|
||||
|
||||
for name, prompt in [
|
||||
("FACT_RETRIEVAL_PROMPT", FACT_RETRIEVAL_PROMPT),
|
||||
("USER_MEMORY_EXTRACTION_PROMPT", USER_MEMORY_EXTRACTION_PROMPT),
|
||||
("AGENT_MEMORY_EXTRACTION_PROMPT", AGENT_MEMORY_EXTRACTION_PROMPT),
|
||||
]:
|
||||
assert "json" in prompt.lower(), (
|
||||
f"{name} should contain 'json' — "
|
||||
"if this fails, the default prompts have changed"
|
||||
)
|
||||
# ensure_json_instruction should be a no-op for defaults
|
||||
system, _ = ensure_json_instruction(prompt, "Input:\nuser: test")
|
||||
assert system == prompt, f"ensure_json_instruction modified {name} unexpectedly"
|
||||
|
||||
# -------------------------------------------------------------------
|
||||
# Integration: verify fix is wired into both sync and async paths
|
||||
# -------------------------------------------------------------------
|
||||
|
||||
def test_fix_applied_in_sync_memory_class(self):
|
||||
"""Verify the ensure_json_instruction call exists in Memory._add_to_vector_store."""
|
||||
import inspect
|
||||
from mem0.memory.main import Memory
|
||||
|
||||
source = inspect.getsource(Memory._add_to_vector_store)
|
||||
assert "ensure_json_instruction" in source, (
|
||||
"ensure_json_instruction not found in Memory._add_to_vector_store (sync)"
|
||||
)
|
||||
|
||||
def test_fix_applied_in_async_memory_class(self):
|
||||
"""Verify the ensure_json_instruction call exists in AsyncMemory._add_to_vector_store."""
|
||||
import inspect
|
||||
from mem0.memory.main import AsyncMemory
|
||||
|
||||
source = inspect.getsource(AsyncMemory._add_to_vector_store)
|
||||
assert "ensure_json_instruction" in source, (
|
||||
"ensure_json_instruction not found in AsyncMemory._add_to_vector_store (async)"
|
||||
)
|
||||
|
||||
def test_import_exists_in_main(self):
|
||||
"""Verify ensure_json_instruction is imported in main.py."""
|
||||
import inspect
|
||||
import mem0.memory.main as main_module
|
||||
|
||||
source = inspect.getsource(main_module)
|
||||
assert "from mem0.memory.utils import" in source
|
||||
assert "ensure_json_instruction" in source
|
||||
@@ -1,8 +1,6 @@
|
||||
from unittest.mock import MagicMock, Mock, patch
|
||||
|
||||
import numpy as np
|
||||
import pytest
|
||||
|
||||
from unittest.mock import Mock, patch
|
||||
from mem0.memory.kuzu_memory import MemoryGraph
|
||||
|
||||
|
||||
@@ -192,48 +190,6 @@ class TestKuzu:
|
||||
assert get_node_count(kuzu_memory) == 0
|
||||
assert get_edge_count(kuzu_memory) == 0
|
||||
|
||||
def _make_kuzu_instance():
|
||||
with patch.object(MemoryGraph, "__init__", return_value=None):
|
||||
instance = MemoryGraph.__new__(MemoryGraph)
|
||||
instance.llm_provider = "openai"
|
||||
instance.llm = MagicMock()
|
||||
instance.embedding_model = MagicMock()
|
||||
instance.config = MagicMock()
|
||||
instance.config.graph_store.custom_prompt = None
|
||||
return instance
|
||||
|
||||
|
||||
class TestRetrieveNodesFromData:
|
||||
"""Tests for _retrieve_nodes_from_data in KuzuMemoryGraph."""
|
||||
|
||||
def test_missing_entities_key_returns_empty(self):
|
||||
"""LLM returns extract_entities tool call without 'entities' key — should not crash.
|
||||
Reproduces the exact scenario from issue #4238."""
|
||||
instance = _make_kuzu_instance()
|
||||
instance.llm.generate_response.return_value = {
|
||||
"tool_calls": [{"name": "extract_entities", "arguments": {"text": "Hello."}}]
|
||||
}
|
||||
result = instance._retrieve_nodes_from_data("Hello.", {"user_id": "u1"})
|
||||
assert result == {}
|
||||
|
||||
def test_normal_entities_extracted(self):
|
||||
instance = _make_kuzu_instance()
|
||||
instance.llm.generate_response.return_value = {
|
||||
"tool_calls": [{"name": "extract_entities", "arguments": {"entities": [
|
||||
{"entity": "Alice", "entity_type": "person"},
|
||||
{"entity": "hiking", "entity_type": "activity"},
|
||||
]}}]
|
||||
}
|
||||
result = instance._retrieve_nodes_from_data("Alice loves hiking", {"user_id": "u1"})
|
||||
assert result == {"alice": "person", "hiking": "activity"}
|
||||
|
||||
def test_none_tool_calls_returns_empty(self):
|
||||
instance = _make_kuzu_instance()
|
||||
instance.llm.generate_response.return_value = {"tool_calls": None}
|
||||
result = instance._retrieve_nodes_from_data("hello world", {"user_id": "u1"})
|
||||
assert result == {}
|
||||
|
||||
|
||||
def get_node_count(kuzu_memory):
|
||||
results = kuzu_memory.kuzu_execute(
|
||||
"""
|
||||
|
||||
@@ -10,7 +10,6 @@ patch.dict("sys.modules", {
|
||||
}).start()
|
||||
|
||||
from mem0.memory.memgraph_memory import MemoryGraph as MemgraphMemoryGraph # noqa: E402
|
||||
|
||||
MemoryGraph = MemgraphMemoryGraph
|
||||
|
||||
|
||||
@@ -55,16 +54,6 @@ class TestRetrieveNodesFromData:
|
||||
assert "relu" in result
|
||||
assert "task" not in result
|
||||
|
||||
def test_missing_entities_key_returns_empty(self):
|
||||
"""LLM returns extract_entities tool call without 'entities' key — should not crash.
|
||||
Reproduces the exact scenario from issue #4238."""
|
||||
instance = _make_instance()
|
||||
instance.llm.generate_response.return_value = {
|
||||
"tool_calls": [{"name": "extract_entities", "arguments": {"text": "Hello."}}]
|
||||
}
|
||||
result = instance._retrieve_nodes_from_data("Hello.", {"user_id": "u1"})
|
||||
assert result == {}
|
||||
|
||||
def test_none_tool_calls_returns_empty(self):
|
||||
instance = _make_instance()
|
||||
instance.llm.generate_response.return_value = {"tool_calls": None}
|
||||
|
||||
+1
-44
@@ -60,7 +60,7 @@ def memory_custom_instance():
|
||||
|
||||
config = MemoryConfig(
|
||||
version="v1.1",
|
||||
custom_fact_extraction_prompt="custom prompt extracting memory in json format",
|
||||
custom_fact_extraction_prompt="custom prompt extracting memory",
|
||||
custom_update_memory_prompt="custom prompt determining memory update",
|
||||
)
|
||||
config.graph_store.config = {"some_config": "value"}
|
||||
@@ -196,15 +196,12 @@ def test_delete_all(memory_instance, version, enable_graph):
|
||||
memory_instance.enable_graph = enable_graph
|
||||
mock_memories = [Mock(id="1"), Mock(id="2")]
|
||||
memory_instance.vector_store.list = Mock(return_value=(mock_memories, None))
|
||||
memory_instance.vector_store.reset = Mock()
|
||||
memory_instance._delete_memory = Mock()
|
||||
memory_instance.graph.delete_all = Mock()
|
||||
|
||||
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
|
||||
memory_instance.vector_store.reset.assert_not_called()
|
||||
|
||||
if enable_graph:
|
||||
memory_instance.graph.delete_all.assert_called_once_with({"user_id": "test_user"})
|
||||
@@ -299,43 +296,3 @@ def test_custom_prompts(memory_custom_instance):
|
||||
messages=[{"role": "user", "content": mock_get_update_memory_messages.return_value}],
|
||||
response_format={"type": "json_object"},
|
||||
)
|
||||
|
||||
|
||||
def test_no_telemetry_vector_store_when_disabled():
|
||||
"""VectorStoreFactory should only be called once (for user data) when telemetry is disabled."""
|
||||
with (
|
||||
patch("mem0.memory.main.MEM0_TELEMETRY", False),
|
||||
patch("mem0.utils.factory.EmbedderFactory") as mock_embedder,
|
||||
patch("mem0.memory.main.VectorStoreFactory") as mock_vector_store,
|
||||
patch("mem0.utils.factory.LlmFactory") as mock_llm,
|
||||
patch("mem0.memory.telemetry.capture_event"),
|
||||
):
|
||||
mock_embedder.create.return_value = Mock()
|
||||
mock_vector_store.create.return_value = Mock()
|
||||
mock_llm.create.return_value = Mock()
|
||||
|
||||
config = MemoryConfig(version="v1.1")
|
||||
Memory(config)
|
||||
|
||||
# VectorStoreFactory.create should be called exactly once — for user data only, not telemetry
|
||||
assert mock_vector_store.create.call_count == 1
|
||||
|
||||
|
||||
def test_telemetry_vector_store_created_when_enabled():
|
||||
"""VectorStoreFactory should be called twice (user data + telemetry) when telemetry is enabled."""
|
||||
with (
|
||||
patch("mem0.memory.main.MEM0_TELEMETRY", True),
|
||||
patch("mem0.utils.factory.EmbedderFactory") as mock_embedder,
|
||||
patch("mem0.memory.main.VectorStoreFactory") as mock_vector_store,
|
||||
patch("mem0.utils.factory.LlmFactory") as mock_llm,
|
||||
patch("mem0.memory.telemetry.capture_event"),
|
||||
):
|
||||
mock_embedder.create.return_value = Mock()
|
||||
mock_vector_store.create.return_value = Mock()
|
||||
mock_llm.create.return_value = Mock()
|
||||
|
||||
config = MemoryConfig(version="v1.1")
|
||||
Memory(config)
|
||||
|
||||
# VectorStoreFactory.create should be called twice — user data + telemetry
|
||||
assert mock_vector_store.create.call_count == 2
|
||||
|
||||
@@ -1,73 +0,0 @@
|
||||
"""Tests for Redis vector store update() — embedding corruption fix.
|
||||
|
||||
Regression tests for #4336: when update() is called with vector=None
|
||||
(metadata-only update), np.array(None) silently creates a 4-byte scalar,
|
||||
overwriting the real embedding. The fix skips the embedding field entirely
|
||||
when vector is None.
|
||||
"""
|
||||
|
||||
from datetime import datetime
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import numpy as np
|
||||
import pytz
|
||||
|
||||
|
||||
def _make_redis_db():
|
||||
"""Create a RedisDB instance with mocked internals, bypassing __init__
|
||||
to avoid the redis module name collision with mem0.vector_stores.redis."""
|
||||
from mem0.vector_stores.redis import RedisDB
|
||||
|
||||
db = RedisDB.__new__(RedisDB)
|
||||
mock_index = MagicMock()
|
||||
db.index = mock_index
|
||||
db.schema = {"index": {"prefix": "mem0:test"}}
|
||||
return db, mock_index
|
||||
|
||||
|
||||
def test_update_with_none_vector_preserves_embedding():
|
||||
"""update() with vector=None should not include embedding in the data."""
|
||||
db, mock_index = _make_redis_db()
|
||||
|
||||
payload = {
|
||||
"hash": "test_hash",
|
||||
"data": "updated_data",
|
||||
"created_at": datetime.now(pytz.timezone("UTC")).isoformat(),
|
||||
"updated_at": datetime.now(pytz.timezone("UTC")).isoformat(),
|
||||
"user_id": "test_user",
|
||||
}
|
||||
|
||||
db.update(vector_id="test_id", vector=None, payload=payload)
|
||||
|
||||
mock_index.load.assert_called_once()
|
||||
call_kwargs = mock_index.load.call_args
|
||||
data_dict = call_kwargs[1]["data"][0] if "data" in call_kwargs[1] else call_kwargs[0][0][0]
|
||||
assert "embedding" not in data_dict, (
|
||||
"embedding should not be in data when vector is None"
|
||||
)
|
||||
assert data_dict["memory_id"] == "test_id"
|
||||
|
||||
|
||||
def test_update_with_vector_includes_embedding():
|
||||
"""update() with a real vector should include embedding in the data."""
|
||||
db, mock_index = _make_redis_db()
|
||||
|
||||
vector = np.random.rand(1536).tolist()
|
||||
payload = {
|
||||
"hash": "test_hash",
|
||||
"data": "updated_data",
|
||||
"created_at": datetime.now(pytz.timezone("UTC")).isoformat(),
|
||||
"updated_at": datetime.now(pytz.timezone("UTC")).isoformat(),
|
||||
"user_id": "test_user",
|
||||
}
|
||||
|
||||
db.update(vector_id="test_id", vector=vector, payload=payload)
|
||||
|
||||
mock_index.load.assert_called_once()
|
||||
call_kwargs = mock_index.load.call_args
|
||||
data_dict = call_kwargs[1]["data"][0] if "data" in call_kwargs[1] else call_kwargs[0][0][0]
|
||||
assert "embedding" in data_dict, (
|
||||
"embedding should be in data when vector is provided"
|
||||
)
|
||||
expected_bytes = np.array(vector, dtype=np.float32).tobytes()
|
||||
assert data_dict["embedding"] == expected_bytes
|
||||
@@ -204,52 +204,6 @@ def test_update_handles_missing_created_at(valkey_db, mock_valkey_client):
|
||||
assert "created_at" in kwargs["mapping"] # Should be added automatically
|
||||
|
||||
|
||||
def test_update_with_none_vector_preserves_embedding(valkey_db, mock_valkey_client):
|
||||
"""Test that update with vector=None does not corrupt the stored embedding.
|
||||
|
||||
Regression test for #4336: when vector=None is passed (metadata-only update),
|
||||
np.array(None) silently creates a 4-byte scalar, overwriting the real embedding.
|
||||
The fix skips the embedding field entirely so the existing value is preserved.
|
||||
"""
|
||||
payload = {
|
||||
"hash": "test_hash",
|
||||
"data": "updated_data",
|
||||
"created_at": datetime.now(pytz.timezone("UTC")).isoformat(),
|
||||
"user_id": "test_user",
|
||||
}
|
||||
|
||||
valkey_db.update(vector_id="test_id", vector=None, payload=payload)
|
||||
|
||||
mock_valkey_client.hset.assert_called_once()
|
||||
args, kwargs = mock_valkey_client.hset.call_args
|
||||
assert "embedding" not in kwargs["mapping"], (
|
||||
"embedding should not be in hash_data when vector is None"
|
||||
)
|
||||
assert kwargs["mapping"]["memory_id"] == "test_id"
|
||||
assert kwargs["mapping"]["memory"] == "updated_data"
|
||||
|
||||
|
||||
def test_update_with_vector_includes_embedding(valkey_db, mock_valkey_client):
|
||||
"""Test that update with a real vector includes the embedding in hash_data."""
|
||||
vector = np.random.rand(1536).tolist()
|
||||
payload = {
|
||||
"hash": "test_hash",
|
||||
"data": "updated_data",
|
||||
"created_at": datetime.now(pytz.timezone("UTC")).isoformat(),
|
||||
"user_id": "test_user",
|
||||
}
|
||||
|
||||
valkey_db.update(vector_id="test_id", vector=vector, payload=payload)
|
||||
|
||||
mock_valkey_client.hset.assert_called_once()
|
||||
args, kwargs = mock_valkey_client.hset.call_args
|
||||
assert "embedding" in kwargs["mapping"], (
|
||||
"embedding should be in hash_data when vector is provided"
|
||||
)
|
||||
expected_bytes = np.array(vector, dtype=np.float32).tobytes()
|
||||
assert kwargs["mapping"]["embedding"] == expected_bytes
|
||||
|
||||
|
||||
def test_get(valkey_db, mock_valkey_client):
|
||||
"""Test getting a vector."""
|
||||
# Mock hgetall to return a vector
|
||||
|
||||
Reference in New Issue
Block a user