Compare commits

..

2 Commits

Author SHA1 Message Date
kartik-mem0 189aea04c9 chore: add config.baseURL fallback for Ollama host 2026-03-14 23:49:00 +05:30
kartik-mem0 e830351691 fix(oss): OllamaLLM now respects configured url instead of always falling back to localhost 2026-03-13 18:27:49 +05:30
36 changed files with 173 additions and 5318 deletions
-13
View File
@@ -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:**
+2 -2
View File
@@ -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",
+1 -5
View File
@@ -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 });
+4 -1
View File
@@ -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" },
]);
});
});
+4 -13
View File
@@ -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"');
+13 -23
View File
@@ -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();
});
});
+2 -2
View File
@@ -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}")
+5 -5
View File
@@ -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(
+1 -1
View File
@@ -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
View File
@@ -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}],
+1 -1
View File
@@ -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:
-25
View File
@@ -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:
+1 -4
View File
@@ -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]
+1 -4
View File
@@ -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())
-2
View File
@@ -1,2 +0,0 @@
package-manager-strict-version=false
approve-builds=esbuild
+6 -34
View File
@@ -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 {
-30
View File
@@ -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
View File
@@ -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"
}
}
-4105
View File
File diff suppressed because it is too large Load Diff
-6
View File
@@ -1,6 +0,0 @@
approveBuilds: esbuild
onlyBuiltDependencies:
- better-sqlite3
- esbuild
- protobufjs
-288
View File
@@ -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);
});
});
-22
View File
@@ -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"]
}
-9
View File
@@ -1,9 +0,0 @@
import { defineConfig } from "tsup";
export default defineConfig({
entry: ["index.ts"],
format: ["esm"],
dts: true,
sourcemap: true,
clean: true,
});
+1 -1
View File
@@ -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]
-6
View File
@@ -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
-214
View File
@@ -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 -45
View File
@@ -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(
"""
-11
View File
@@ -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
View File
@@ -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
-73
View File
@@ -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
-46
View File
@@ -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