feat(oss): port v3 pipeline with hybrid search, entity extraction, and additive scoring (#4805)
Co-authored-by: Soumil Rathi <soumilrathi@gmail.com> Co-authored-by: Saket Aryan <saketaryan2002@gmail.com> Co-authored-by: chaithanyak42 <chaithanya.kumar42a@gmail.com> Co-authored-by: kartik-mem0 <kartik.labhshetwar@mem0.ai>
This commit is contained in:
@@ -88,6 +88,13 @@ Install the sdk via pip:
|
||||
pip install mem0ai
|
||||
```
|
||||
|
||||
For enhanced hybrid search with BM25 keyword matching and entity extraction, install with NLP support:
|
||||
|
||||
```bash
|
||||
pip install mem0ai[nlp]
|
||||
python -m spacy download en_core_web_sm
|
||||
```
|
||||
|
||||
Install sdk via npm:
|
||||
```bash
|
||||
npm install mem0ai
|
||||
@@ -109,7 +116,9 @@ See the [CLI documentation](https://docs.mem0.ai/platform/cli) for the full comm
|
||||
|
||||
### Basic Usage
|
||||
|
||||
Mem0 requires an LLM to function, with `gpt-4.1-nano-2025-04-14 from OpenAI as the default. However, it supports a variety of LLMs; for details, refer to our [Supported LLMs documentation](https://docs.mem0.ai/components/llms/overview).
|
||||
Mem0 requires an LLM to function, with `gpt-4.1-nano-2025-04-14` from OpenAI as the default. However, it supports a variety of LLMs; for details, refer to our [Supported LLMs documentation](https://docs.mem0.ai/components/llms/overview).
|
||||
|
||||
Mem0 uses `text-embedding-3-small` from OpenAI as the default embedding model. For best results with hybrid search (semantic + keyword + entity boosting), we recommend using at least [Qwen 600M](https://huggingface.co/Alibaba-NLP/gte-Qwen2-1.5B-instruct) or a comparable embedding model. See [Supported Embeddings](https://docs.mem0.ai/components/embedders/overview) for configuration details.
|
||||
|
||||
First step is to instantiate the memory:
|
||||
|
||||
|
||||
@@ -120,10 +120,11 @@
|
||||
"better-sqlite3": "^12.6.2",
|
||||
"cloudflare": "^4.2.0",
|
||||
"groq-sdk": "0.3.0",
|
||||
"neo4j-driver": "^5.28.1",
|
||||
"ollama": "^0.5.14",
|
||||
"pg": "8.11.3",
|
||||
"redis": "^4.6.13"
|
||||
"redis": "^4.6.13",
|
||||
"compromise": "^14.0.0",
|
||||
"natural": "^8.0.1"
|
||||
},
|
||||
"engines": {
|
||||
"node": ">=18"
|
||||
|
||||
Generated
+2355
-4130
File diff suppressed because it is too large
Load Diff
@@ -3,7 +3,6 @@ import type * as MemoryTypes from "./mem0.types";
|
||||
|
||||
// Re-export all types from mem0.types
|
||||
export type {
|
||||
EntityOptions,
|
||||
AddMemoryOptions,
|
||||
SearchMemoryOptions,
|
||||
GetAllMemoryOptions,
|
||||
|
||||
+56
-10
@@ -23,6 +23,37 @@ import { captureClientEvent, generateHash } from "./telemetry";
|
||||
import { camelToSnake, camelToSnakeKeys, snakeToCamelKeys } from "./utils";
|
||||
import { createExceptionFromResponse, MemoryError } from "../common/exceptions";
|
||||
|
||||
// Entity params that must be passed via filters - check both snake_case and camelCase
|
||||
const ENTITY_PARAMS = [
|
||||
"user_id",
|
||||
"agent_id",
|
||||
"app_id",
|
||||
"run_id",
|
||||
"userId",
|
||||
"agentId",
|
||||
"appId",
|
||||
"runId",
|
||||
];
|
||||
|
||||
/**
|
||||
* Validates that no top-level entity parameters are passed.
|
||||
* @throws Error if entity params are found at top level
|
||||
*/
|
||||
function rejectTopLevelEntityParams(
|
||||
options: Record<string, any> | undefined,
|
||||
methodName: string,
|
||||
): void {
|
||||
const invalidKeys = Object.keys(options ?? {}).filter((k) =>
|
||||
ENTITY_PARAMS.includes(k),
|
||||
);
|
||||
if (invalidKeys.length > 0) {
|
||||
throw new Error(
|
||||
`Top-level entity parameters [${invalidKeys.join(", ")}] are not supported in ${methodName}(). ` +
|
||||
`Use filters: { user_id: "..." } instead.`,
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
class APIError extends Error {
|
||||
constructor(message: string) {
|
||||
super(message);
|
||||
@@ -190,7 +221,7 @@ export default class MemoryClient {
|
||||
this._captureEvent("add", [payloadKeys]);
|
||||
|
||||
const response = await this._fetchWithErrorHandling(
|
||||
`${this.host}/v1/memories/`,
|
||||
`${this.host}/v3/memories/`,
|
||||
{
|
||||
method: "POST",
|
||||
headers: this.headers,
|
||||
@@ -254,13 +285,17 @@ export default class MemoryClient {
|
||||
}
|
||||
|
||||
async getAll(options?: GetAllMemoryOptions): Promise<Array<Memory>> {
|
||||
// Reject top-level entity params - must use filters instead
|
||||
rejectTopLevelEntityParams(options as Record<string, any>, "getAll");
|
||||
|
||||
if (this.telemetryId === "") await this.ping();
|
||||
const payloadKeys = Object.keys(options || {});
|
||||
this._captureEvent("get_all", [payloadKeys]);
|
||||
const { page, pageSize, ...rest } = options ?? {};
|
||||
const { page, pageSize, filters, ...rest } = options ?? {};
|
||||
const body: Record<string, any> = {
|
||||
output_format: "v1.1",
|
||||
...camelToSnakeKeys(rest),
|
||||
...(filters && { filters }),
|
||||
};
|
||||
|
||||
let url = `${this.host}/v2/memories/`;
|
||||
@@ -273,33 +308,36 @@ export default class MemoryClient {
|
||||
headers: this.headers,
|
||||
body: JSON.stringify(body),
|
||||
});
|
||||
// Unwrap v1.1 format: { results: [...] } → [...]
|
||||
return Array.isArray(response) ? response : (response?.results ?? response);
|
||||
}
|
||||
|
||||
async search(
|
||||
query: string,
|
||||
options?: SearchMemoryOptions,
|
||||
): Promise<Array<Memory>> {
|
||||
): Promise<{ results: Array<Memory> }> {
|
||||
// Reject top-level entity params - must use filters instead
|
||||
rejectTopLevelEntityParams(options as Record<string, any>, "search");
|
||||
|
||||
if (this.telemetryId === "") await this.ping();
|
||||
const payloadKeys = Object.keys(options || {});
|
||||
this._captureEvent("search", [payloadKeys]);
|
||||
const { filters, ...rest } = options ?? {};
|
||||
const payload: Record<string, any> = {
|
||||
query,
|
||||
output_format: "v1.1",
|
||||
...camelToSnakeKeys(options ?? {}),
|
||||
...camelToSnakeKeys(rest),
|
||||
...(filters && { filters }),
|
||||
};
|
||||
|
||||
const response = await this._fetchWithErrorHandling(
|
||||
`${this.host}/v2/memories/search/`,
|
||||
`${this.host}/v3/memories/search/`,
|
||||
{
|
||||
method: "POST",
|
||||
headers: this.headers,
|
||||
body: JSON.stringify(payload),
|
||||
},
|
||||
);
|
||||
// Unwrap v1.1 format: { results: [...] } → [...]
|
||||
return Array.isArray(response) ? response : (response?.results ?? response);
|
||||
return response;
|
||||
}
|
||||
|
||||
async delete(memoryId: string): Promise<{ message: string }> {
|
||||
@@ -614,12 +652,16 @@ export default class MemoryClient {
|
||||
throw new Error("Missing filters or schema");
|
||||
}
|
||||
|
||||
const { filters, ...rest } = data;
|
||||
const response = await this._fetchWithErrorHandling(
|
||||
`${this.host}/v1/exports/`,
|
||||
{
|
||||
method: "POST",
|
||||
headers: this.headers,
|
||||
body: JSON.stringify(camelToSnakeKeys(data)),
|
||||
body: JSON.stringify({
|
||||
...camelToSnakeKeys(rest),
|
||||
filters,
|
||||
}),
|
||||
},
|
||||
);
|
||||
|
||||
@@ -636,12 +678,16 @@ export default class MemoryClient {
|
||||
throw new Error("Missing memoryExportId or filters");
|
||||
}
|
||||
|
||||
const { filters, ...rest } = data;
|
||||
const response = await this._fetchWithErrorHandling(
|
||||
`${this.host}/v1/exports/get/`,
|
||||
{
|
||||
method: "POST",
|
||||
headers: this.headers,
|
||||
body: JSON.stringify(camelToSnakeKeys(data)),
|
||||
body: JSON.stringify({
|
||||
...camelToSnakeKeys(rest),
|
||||
...(filters && { filters }),
|
||||
}),
|
||||
},
|
||||
);
|
||||
return response;
|
||||
|
||||
@@ -51,17 +51,13 @@ describeIntegration("MemoryClient Integration — CRUD", () => {
|
||||
},
|
||||
];
|
||||
|
||||
const result = await client.add(messages, { userId: TEST_USER_ID });
|
||||
const result = await client.add(messages, {
|
||||
userId: TEST_USER_ID,
|
||||
});
|
||||
|
||||
// API processes memories asynchronously — returns PENDING
|
||||
expect(Array.isArray(result)).toBe(true);
|
||||
expect(result.length).toBeGreaterThan(0);
|
||||
|
||||
// Validate response shape
|
||||
for (const item of result) {
|
||||
expect(item).toHaveProperty("status");
|
||||
expect(item).toHaveProperty("eventId");
|
||||
}
|
||||
// v3 API processes memories asynchronously — returns PENDING
|
||||
expect(result).toHaveProperty("eventId");
|
||||
expect(result).toHaveProperty("status");
|
||||
});
|
||||
|
||||
test("adds a second batch of messages", async () => {
|
||||
@@ -76,8 +72,10 @@ describeIntegration("MemoryClient Integration — CRUD", () => {
|
||||
},
|
||||
];
|
||||
|
||||
const result = await client.add(messages, { userId: TEST_USER_ID });
|
||||
expect(Array.isArray(result)).toBe(true);
|
||||
const result = await client.add(messages, {
|
||||
userId: TEST_USER_ID,
|
||||
});
|
||||
expect(result).toHaveProperty("eventId");
|
||||
});
|
||||
|
||||
test("memories become available after async processing", async () => {
|
||||
@@ -123,7 +121,7 @@ describeIntegration("MemoryClient Integration — CRUD", () => {
|
||||
describe("get all memories", () => {
|
||||
test("returns all memories for test user", async () => {
|
||||
const memories = await client.getAll({
|
||||
filters: { userId: TEST_USER_ID },
|
||||
filters: { user_id: TEST_USER_ID },
|
||||
});
|
||||
|
||||
expect(Array.isArray(memories)).toBe(true);
|
||||
@@ -137,7 +135,7 @@ describeIntegration("MemoryClient Integration — CRUD", () => {
|
||||
|
||||
test("returns paginated results with page and page_size", async () => {
|
||||
const page1 = await client.getAll({
|
||||
filters: { userId: TEST_USER_ID },
|
||||
filters: { user_id: TEST_USER_ID },
|
||||
page: 1,
|
||||
pageSize: 1,
|
||||
});
|
||||
@@ -196,13 +194,12 @@ describeIntegration("MemoryClient Integration — CRUD", () => {
|
||||
},
|
||||
);
|
||||
|
||||
expect(Array.isArray(result)).toBe(true);
|
||||
expect(result.length).toBeGreaterThan(0);
|
||||
expect(result).toHaveProperty("eventId");
|
||||
});
|
||||
|
||||
test("getAll for non-existent user returns empty array", async () => {
|
||||
const memories = await client.getAll({
|
||||
filters: { userId: `nonexistent-user-${randomUUID()}` },
|
||||
filters: { user_id: `nonexistent-user-${randomUUID()}` },
|
||||
});
|
||||
|
||||
expect(Array.isArray(memories)).toBe(true);
|
||||
@@ -241,7 +238,9 @@ describeIntegration("MemoryClient Integration — CRUD", () => {
|
||||
// ─── Delete all + delete user ─────────────────────────────
|
||||
describe("cleanup operations", () => {
|
||||
test("deletes all memories for test user", async () => {
|
||||
const result = await client.deleteAll({ userId: TEST_USER_ID });
|
||||
const result = await client.deleteAll({
|
||||
userId: TEST_USER_ID,
|
||||
});
|
||||
expect(result).toBeDefined();
|
||||
expect(typeof result.message).toBe("string");
|
||||
});
|
||||
|
||||
@@ -64,7 +64,7 @@ export async function waitForMemories(
|
||||
): Promise<Memory[]> {
|
||||
for (let attempt = 1; attempt <= maxRetries; attempt++) {
|
||||
const memories = await withRetry(() =>
|
||||
client.getAll({ filters: { userId } }),
|
||||
client.getAll({ filters: { user_id: userId } }),
|
||||
);
|
||||
if (Array.isArray(memories) && memories.length >= minCount) {
|
||||
return memories;
|
||||
@@ -92,8 +92,9 @@ export async function waitForSearchResults(
|
||||
maxRetries = 4,
|
||||
): Promise<Memory[]> {
|
||||
for (let attempt = 1; attempt <= maxRetries; attempt++) {
|
||||
const results = await withRetry(() => client.search(query, options));
|
||||
if (Array.isArray(results) && results.length > 0) {
|
||||
const response = await withRetry(() => client.search(query, options));
|
||||
const results = response?.results ?? [];
|
||||
if (results.length > 0) {
|
||||
return results;
|
||||
}
|
||||
if (attempt < maxRetries) {
|
||||
|
||||
@@ -43,7 +43,7 @@ describeIntegration("MemoryClient Integration — Search & History", () => {
|
||||
const results = await waitForSearchResults(
|
||||
client,
|
||||
"What is my favorite color?",
|
||||
{ filters: { userId: TEST_USER_ID } },
|
||||
{ filters: { user_id: TEST_USER_ID } },
|
||||
);
|
||||
|
||||
expect(Array.isArray(results)).toBe(true);
|
||||
@@ -64,7 +64,7 @@ describeIntegration("MemoryClient Integration — Search & History", () => {
|
||||
client,
|
||||
"What do you know about me?",
|
||||
{
|
||||
filters: { OR: [{ userId: TEST_USER_ID }] },
|
||||
filters: { OR: [{ user_id: TEST_USER_ID }] },
|
||||
},
|
||||
);
|
||||
|
||||
@@ -108,24 +108,24 @@ describeIntegration("MemoryClient Integration — Search & History", () => {
|
||||
// ─── Edge cases ─────────────────────────────────────────
|
||||
describe("edge cases", () => {
|
||||
test("search for non-existent user returns empty results", async () => {
|
||||
const results = await client.search("anything", {
|
||||
filters: { userId: `nonexistent-user-${randomUUID()}` },
|
||||
const response = await client.search("test search query", {
|
||||
filters: { user_id: `nonexistent-user-${randomUUID()}` },
|
||||
});
|
||||
|
||||
expect(Array.isArray(results)).toBe(true);
|
||||
expect(results.length).toBe(0);
|
||||
expect(response).toHaveProperty("results");
|
||||
expect(response.results).toHaveLength(0);
|
||||
});
|
||||
|
||||
test("search with top_k param does not throw", async () => {
|
||||
const results = await client.search(
|
||||
const response = await client.search(
|
||||
"Tell me about integration test user",
|
||||
{
|
||||
filters: { userId: TEST_USER_ID },
|
||||
filters: { user_id: TEST_USER_ID },
|
||||
topK: 1,
|
||||
},
|
||||
);
|
||||
|
||||
expect(Array.isArray(results)).toBe(true);
|
||||
expect(response).toHaveProperty("results");
|
||||
});
|
||||
});
|
||||
});
|
||||
|
||||
@@ -21,33 +21,33 @@ installConsoleSuppression();
|
||||
// ─── add() ───────────────────────────────────────────────
|
||||
|
||||
describe("MemoryClient - add()", () => {
|
||||
test("sends POST to /v1/memories/", async () => {
|
||||
test("sends POST to /v3/memories/", async () => {
|
||||
const extra = new Map<string, { status: number; body: unknown }>();
|
||||
extra.set("/v1/memories/", { status: 200, body: [createMockMemory()] });
|
||||
extra.set("/v3/memories/", { status: 200, body: [createMockMemory()] });
|
||||
const mock = setupMockFetch(extra);
|
||||
|
||||
const client = new MemoryClient({ apiKey: TEST_API_KEY });
|
||||
await client.add([{ role: "user", content: "Hello" }], { userId: "u1" });
|
||||
|
||||
expect(findFetchCall(mock, "/v1/memories/", "POST")).toBeDefined();
|
||||
expect(findFetchCall(mock, "/v3/memories/", "POST")).toBeDefined();
|
||||
});
|
||||
|
||||
test("includes messages in request body", async () => {
|
||||
const messages = [{ role: "user" as const, content: "Hello, I am Alex" }];
|
||||
const extra = new Map<string, { status: number; body: unknown }>();
|
||||
extra.set("/v1/memories/", { status: 200, body: [createMockMemory()] });
|
||||
extra.set("/v3/memories/", { status: 200, body: [createMockMemory()] });
|
||||
const mock = setupMockFetch(extra);
|
||||
|
||||
const client = new MemoryClient({ apiKey: TEST_API_KEY });
|
||||
await client.add(messages, { userId: "u1" });
|
||||
|
||||
const call = findFetchCall(mock, "/v1/memories/", "POST");
|
||||
const call = findFetchCall(mock, "/v3/memories/", "POST");
|
||||
expect(getFetchBody(call!).messages).toEqual(messages);
|
||||
});
|
||||
|
||||
test("includes user_id in request body", async () => {
|
||||
const extra = new Map<string, { status: number; body: unknown }>();
|
||||
extra.set("/v1/memories/", { status: 200, body: [createMockMemory()] });
|
||||
extra.set("/v3/memories/", { status: 200, body: [createMockMemory()] });
|
||||
const mock = setupMockFetch(extra);
|
||||
|
||||
const client = new MemoryClient({ apiKey: TEST_API_KEY });
|
||||
@@ -55,19 +55,19 @@ describe("MemoryClient - add()", () => {
|
||||
user_id: "user_1",
|
||||
});
|
||||
|
||||
const call = findFetchCall(mock, "/v1/memories/", "POST");
|
||||
const call = findFetchCall(mock, "/v3/memories/", "POST");
|
||||
expect(getFetchBody(call!).user_id).toBe("user_1");
|
||||
});
|
||||
|
||||
test("sends empty messages array without crashing", async () => {
|
||||
const extra = new Map<string, { status: number; body: unknown }>();
|
||||
extra.set("/v1/memories/", { status: 200, body: [] });
|
||||
extra.set("/v3/memories/", { status: 200, body: [] });
|
||||
const mock = setupMockFetch(extra);
|
||||
|
||||
const client = new MemoryClient({ apiKey: TEST_API_KEY });
|
||||
await client.add([], { userId: "u1" });
|
||||
|
||||
const call = findFetchCall(mock, "/v1/memories/", "POST");
|
||||
const call = findFetchCall(mock, "/v3/memories/", "POST");
|
||||
expect(getFetchBody(call!).messages).toEqual([]);
|
||||
});
|
||||
});
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
/**
|
||||
* MemoryClient unit tests — search (v2 default, filters).
|
||||
* MemoryClient unit tests — search (v3 endpoint, filters).
|
||||
* Tests verify request construction, not mock response echo.
|
||||
*/
|
||||
import { MemoryClient } from "../mem0";
|
||||
@@ -15,48 +15,60 @@ import {
|
||||
installConsoleSuppression();
|
||||
|
||||
describe("MemoryClient - search()", () => {
|
||||
test("sends POST to /v2/memories/search/ by default", async () => {
|
||||
test("sends POST to /v3/memories/search/ by default", async () => {
|
||||
const extra = new Map<string, { status: number; body: unknown }>();
|
||||
extra.set("/v2/memories/search/", { status: 200, body: [] });
|
||||
extra.set("/v3/memories/search/", {
|
||||
status: 200,
|
||||
body: { results: [] },
|
||||
});
|
||||
const mock = setupMockFetch(extra);
|
||||
|
||||
const client = new MemoryClient({ apiKey: TEST_API_KEY });
|
||||
await client.search("What is my name?", {
|
||||
filters: { userId: "u1" },
|
||||
filters: { user_id: "u1" },
|
||||
});
|
||||
|
||||
expect(findFetchCall(mock, "/v2/memories/search/", "POST")).toBeDefined();
|
||||
expect(findFetchCall(mock, "/v3/memories/search/", "POST")).toBeDefined();
|
||||
});
|
||||
|
||||
test("includes query in request body", async () => {
|
||||
const extra = new Map<string, { status: number; body: unknown }>();
|
||||
extra.set("/v2/memories/search/", { status: 200, body: [] });
|
||||
extra.set("/v3/memories/search/", {
|
||||
status: 200,
|
||||
body: { results: [] },
|
||||
});
|
||||
const mock = setupMockFetch(extra);
|
||||
|
||||
const client = new MemoryClient({ apiKey: TEST_API_KEY });
|
||||
await client.search("What is my name?", {
|
||||
filters: { userId: "u1" },
|
||||
filters: { user_id: "u1" },
|
||||
});
|
||||
|
||||
const call = findFetchCall(mock, "/v2/memories/search/", "POST");
|
||||
const call = findFetchCall(mock, "/v3/memories/search/", "POST");
|
||||
expect(getFetchBody(call!).query).toBe("What is my name?");
|
||||
});
|
||||
|
||||
test("passes filters through to the API body", async () => {
|
||||
const extra = new Map<string, { status: number; body: unknown }>();
|
||||
extra.set("/v2/memories/search/", { status: 200, body: [] });
|
||||
extra.set("/v3/memories/search/", {
|
||||
status: 200,
|
||||
body: { results: [] },
|
||||
});
|
||||
const mock = setupMockFetch(extra);
|
||||
|
||||
const client = new MemoryClient({ apiKey: TEST_API_KEY });
|
||||
await client.search("test", { filters: { userId: "u1" } });
|
||||
await client.search("test", { filters: { user_id: "u1" } });
|
||||
|
||||
const call = findFetchCall(mock, "/v2/memories/search/", "POST");
|
||||
const call = findFetchCall(mock, "/v3/memories/search/", "POST");
|
||||
expect(getFetchBody(call!).filters).toEqual({ user_id: "u1" });
|
||||
});
|
||||
|
||||
test("passes complex OR filters through to the API body", async () => {
|
||||
const extra = new Map<string, { status: number; body: unknown }>();
|
||||
extra.set("/v2/memories/search/", { status: 200, body: [] });
|
||||
extra.set("/v3/memories/search/", {
|
||||
status: 200,
|
||||
body: { results: [] },
|
||||
});
|
||||
const mock = setupMockFetch(extra);
|
||||
|
||||
const client = new MemoryClient({ apiKey: TEST_API_KEY });
|
||||
@@ -64,32 +76,192 @@ describe("MemoryClient - search()", () => {
|
||||
filters: { OR: [{ user_id: "u1" }, { agent_id: "a1" }] },
|
||||
});
|
||||
|
||||
const call = findFetchCall(mock, "/v2/memories/search/", "POST");
|
||||
const call = findFetchCall(mock, "/v3/memories/search/", "POST");
|
||||
const body = getFetchBody(call!);
|
||||
expect(body.filters).toEqual({
|
||||
OR: [{ user_id: "u1" }, { agent_id: "a1" }],
|
||||
});
|
||||
});
|
||||
|
||||
test("passes complex AND filters through to the API body", async () => {
|
||||
const extra = new Map<string, { status: number; body: unknown }>();
|
||||
extra.set("/v3/memories/search/", {
|
||||
status: 200,
|
||||
body: { results: [] },
|
||||
});
|
||||
const mock = setupMockFetch(extra);
|
||||
|
||||
const client = new MemoryClient({ apiKey: TEST_API_KEY });
|
||||
await client.search("query", {
|
||||
filters: {
|
||||
AND: [
|
||||
{ user_id: "u1" },
|
||||
{ created_at: { gte: "2024-01-01T00:00:00Z" } },
|
||||
],
|
||||
},
|
||||
});
|
||||
|
||||
const call = findFetchCall(mock, "/v3/memories/search/", "POST");
|
||||
const body = getFetchBody(call!);
|
||||
expect(body.filters).toEqual({
|
||||
AND: [{ user_id: "u1" }, { created_at: { gte: "2024-01-01T00:00:00Z" } }],
|
||||
});
|
||||
});
|
||||
|
||||
test("passes NOT filters through to the API body", async () => {
|
||||
const extra = new Map<string, { status: number; body: unknown }>();
|
||||
extra.set("/v3/memories/search/", {
|
||||
status: 200,
|
||||
body: { results: [] },
|
||||
});
|
||||
const mock = setupMockFetch(extra);
|
||||
|
||||
const client = new MemoryClient({ apiKey: TEST_API_KEY });
|
||||
await client.search("query", {
|
||||
filters: {
|
||||
AND: [
|
||||
{ user_id: "u1" },
|
||||
{ NOT: { categories: { in: ["spam", "test"] } } },
|
||||
],
|
||||
},
|
||||
});
|
||||
|
||||
const call = findFetchCall(mock, "/v3/memories/search/", "POST");
|
||||
const body = getFetchBody(call!);
|
||||
expect(body.filters).toEqual({
|
||||
AND: [
|
||||
{ user_id: "u1" },
|
||||
{ NOT: { categories: { in: ["spam", "test"] } } },
|
||||
],
|
||||
});
|
||||
});
|
||||
|
||||
test("passes complex nested AND/OR/NOT filters through to the API body", async () => {
|
||||
const extra = new Map<string, { status: number; body: unknown }>();
|
||||
extra.set("/v3/memories/search/", {
|
||||
status: 200,
|
||||
body: { results: [] },
|
||||
});
|
||||
const mock = setupMockFetch(extra);
|
||||
|
||||
const client = new MemoryClient({ apiKey: TEST_API_KEY });
|
||||
const complexFilter = {
|
||||
AND: [
|
||||
{ user_id: "u1" },
|
||||
{ created_at: { gte: "2024-01-01T00:00:00Z" } },
|
||||
{
|
||||
NOT: {
|
||||
OR: [
|
||||
{ categories: { in: ["spam"] } },
|
||||
{ categories: { in: ["test"] } },
|
||||
],
|
||||
},
|
||||
},
|
||||
],
|
||||
};
|
||||
await client.search("query", { filters: complexFilter });
|
||||
|
||||
const call = findFetchCall(mock, "/v3/memories/search/", "POST");
|
||||
const body = getFetchBody(call!);
|
||||
expect(body.filters).toEqual(complexFilter);
|
||||
});
|
||||
|
||||
test("does not crash when called without options", async () => {
|
||||
const extra = new Map<string, { status: number; body: unknown }>();
|
||||
extra.set("/v2/memories/search/", { status: 200, body: [] });
|
||||
extra.set("/v3/memories/search/", {
|
||||
status: 200,
|
||||
body: { results: [] },
|
||||
});
|
||||
setupMockFetch(extra);
|
||||
|
||||
const client = new MemoryClient({ apiKey: TEST_API_KEY });
|
||||
const result: Memory[] = await client.search("query");
|
||||
expect(Array.isArray(result)).toBe(true);
|
||||
const result = await client.search("query");
|
||||
expect(result).toHaveProperty("results");
|
||||
});
|
||||
|
||||
test("handles empty results array", async () => {
|
||||
const extra = new Map<string, { status: number; body: unknown }>();
|
||||
extra.set("/v2/memories/search/", { status: 200, body: [] });
|
||||
extra.set("/v3/memories/search/", {
|
||||
status: 200,
|
||||
body: { results: [] },
|
||||
});
|
||||
setupMockFetch(extra);
|
||||
|
||||
const client = new MemoryClient({ apiKey: TEST_API_KEY });
|
||||
const result: Memory[] = await client.search("nonexistent query", {
|
||||
filters: { userId: "u1" },
|
||||
const result = await client.search("nonexistent query", {
|
||||
filters: { AND: [{ user_id: "u1" }] },
|
||||
});
|
||||
expect(result).toHaveLength(0);
|
||||
expect(result.results).toHaveLength(0);
|
||||
});
|
||||
});
|
||||
|
||||
describe("MemoryClient - search() entity param rejection", () => {
|
||||
test("rejects user_id at top level", async () => {
|
||||
setupMockFetch();
|
||||
const client = new MemoryClient({ apiKey: TEST_API_KEY });
|
||||
await expect(
|
||||
client.search("query", { user_id: "u1" } as any),
|
||||
).rejects.toThrow(/filters/);
|
||||
});
|
||||
|
||||
test("rejects agent_id at top level", async () => {
|
||||
setupMockFetch();
|
||||
const client = new MemoryClient({ apiKey: TEST_API_KEY });
|
||||
await expect(
|
||||
client.search("query", { agent_id: "a1" } as any),
|
||||
).rejects.toThrow(/filters/);
|
||||
});
|
||||
|
||||
test("rejects app_id at top level", async () => {
|
||||
setupMockFetch();
|
||||
const client = new MemoryClient({ apiKey: TEST_API_KEY });
|
||||
await expect(
|
||||
client.search("query", { app_id: "app1" } as any),
|
||||
).rejects.toThrow(/filters/);
|
||||
});
|
||||
|
||||
test("accepts filters with user_id", async () => {
|
||||
const extra = new Map<string, { status: number; body: unknown }>();
|
||||
extra.set("/v3/memories/search/", {
|
||||
status: 200,
|
||||
body: { results: [] },
|
||||
});
|
||||
const mock = setupMockFetch(extra);
|
||||
|
||||
const client = new MemoryClient({ apiKey: TEST_API_KEY });
|
||||
// Should not throw
|
||||
await client.search("query", { filters: { AND: [{ user_id: "u1" }] } });
|
||||
expect(findFetchCall(mock, "/v3/memories/search/", "POST")).toBeDefined();
|
||||
});
|
||||
});
|
||||
|
||||
describe("MemoryClient - getAll() entity param rejection", () => {
|
||||
test("rejects user_id at top level", async () => {
|
||||
setupMockFetch();
|
||||
const client = new MemoryClient({ apiKey: TEST_API_KEY });
|
||||
await expect(client.getAll({ user_id: "u1" } as any)).rejects.toThrow(
|
||||
/filters/,
|
||||
);
|
||||
});
|
||||
|
||||
test("rejects agent_id at top level", async () => {
|
||||
setupMockFetch();
|
||||
const client = new MemoryClient({ apiKey: TEST_API_KEY });
|
||||
await expect(client.getAll({ agent_id: "a1" } as any)).rejects.toThrow(
|
||||
/filters/,
|
||||
);
|
||||
});
|
||||
|
||||
test("accepts filters with user_id", async () => {
|
||||
const extra = new Map<string, { status: number; body: unknown }>();
|
||||
extra.set("/v2/memories/", {
|
||||
status: 200,
|
||||
body: { results: [] },
|
||||
});
|
||||
const mock = setupMockFetch(extra);
|
||||
|
||||
const client = new MemoryClient({ apiKey: TEST_API_KEY });
|
||||
await client.getAll({ filters: { user_id: "u1" } });
|
||||
expect(findFetchCall(mock, "/v2/memories/", "POST")).toBeDefined();
|
||||
});
|
||||
});
|
||||
|
||||
@@ -82,14 +82,14 @@ async function runTests(memory: Memory) {
|
||||
// Get all memories
|
||||
console.log("\nGetting all memories...");
|
||||
const allMemories = await memory.getAll({
|
||||
userId: "john",
|
||||
filters: { user_id: "john" },
|
||||
});
|
||||
console.log("All memories:", allMemories);
|
||||
|
||||
// Search for memories
|
||||
console.log("\nSearching memories...");
|
||||
const searchResult = await memory.search("What do you know about Paris?", {
|
||||
userId: "john",
|
||||
filters: { user_id: "john" },
|
||||
});
|
||||
console.log("Search results:", searchResult);
|
||||
|
||||
@@ -292,86 +292,6 @@ async function demoRedis() {
|
||||
await runTests(memory);
|
||||
}
|
||||
|
||||
async function demoGraphMemory() {
|
||||
console.log("\n=== Testing Graph Memory Store ===\n");
|
||||
|
||||
const memory = new Memory({
|
||||
version: "v1.1",
|
||||
embedder: {
|
||||
provider: "openai",
|
||||
config: {
|
||||
apiKey: process.env.OPENAI_API_KEY || "",
|
||||
model: "text-embedding-3-small",
|
||||
},
|
||||
},
|
||||
vectorStore: {
|
||||
provider: "memory",
|
||||
config: {
|
||||
collectionName: "memories",
|
||||
dimension: 1536,
|
||||
},
|
||||
},
|
||||
llm: {
|
||||
provider: "openai",
|
||||
config: {
|
||||
apiKey: process.env.OPENAI_API_KEY || "",
|
||||
model: "gpt-4-turbo-preview",
|
||||
},
|
||||
},
|
||||
graphStore: {
|
||||
provider: "neo4j",
|
||||
config: {
|
||||
url: process.env.NEO4J_URL || "neo4j://localhost:7687",
|
||||
username: process.env.NEO4J_USERNAME || "neo4j",
|
||||
password: process.env.NEO4J_PASSWORD || "password",
|
||||
},
|
||||
llm: {
|
||||
provider: "openai",
|
||||
config: {
|
||||
model: "gpt-4-turbo-preview",
|
||||
},
|
||||
},
|
||||
},
|
||||
historyDbPath: "memory.db",
|
||||
});
|
||||
|
||||
try {
|
||||
// Reset all memories
|
||||
await memory.reset();
|
||||
|
||||
// Add memories with relationships
|
||||
const result = await memory.add(
|
||||
[
|
||||
{
|
||||
role: "user",
|
||||
content: "Alice is Bob's sister and works as a doctor.",
|
||||
},
|
||||
{
|
||||
role: "assistant",
|
||||
content:
|
||||
"I understand that Alice and Bob are siblings and Alice is a medical professional.",
|
||||
},
|
||||
{ role: "user", content: "Bob is married to Carol who is a teacher." },
|
||||
],
|
||||
{
|
||||
userId: "john",
|
||||
},
|
||||
);
|
||||
console.log("Added memories with relationships:", result);
|
||||
|
||||
// Search for connected information
|
||||
const searchResult = await memory.search(
|
||||
"Tell me about Bob's family connections",
|
||||
{
|
||||
userId: "john",
|
||||
},
|
||||
);
|
||||
console.log("Search results with graph relationships:", searchResult);
|
||||
} catch (error) {
|
||||
console.error("Error in graph memory demo:", error);
|
||||
}
|
||||
}
|
||||
|
||||
async function main() {
|
||||
// Test in-memory store
|
||||
await demoMemoryStore();
|
||||
@@ -379,19 +299,6 @@ async function main() {
|
||||
// Test in-memory store with Ollama
|
||||
await demoLocalMemory();
|
||||
|
||||
// Test graph memory if Neo4j environment variables are set
|
||||
if (
|
||||
process.env.NEO4J_URL &&
|
||||
process.env.NEO4J_USERNAME &&
|
||||
process.env.NEO4J_PASSWORD
|
||||
) {
|
||||
await demoGraphMemory();
|
||||
} else {
|
||||
console.log(
|
||||
"\nSkipping Graph Memory test - Neo4j environment variables not set",
|
||||
);
|
||||
}
|
||||
|
||||
// Test PGVector store if environment variables are set
|
||||
if (process.env.PGVECTOR_DB) {
|
||||
await demoPGVector();
|
||||
|
||||
@@ -26,7 +26,9 @@ const memory = new Memory({
|
||||
});
|
||||
|
||||
async function chatWithMemories(message: string, userId = "default_user") {
|
||||
const relevantMemories = await memory.search(message, { userId: userId });
|
||||
const relevantMemories = await memory.search(message, {
|
||||
filters: { user_id: userId },
|
||||
});
|
||||
|
||||
const memoriesStr = relevantMemories.results
|
||||
.map((entry) => `- ${entry.memory}`)
|
||||
|
||||
@@ -64,14 +64,14 @@ export async function runTests(memory: Memory) {
|
||||
// Get all memories
|
||||
console.log("\nGetting all memories...");
|
||||
const allMemories = await memory.getAll({
|
||||
userId: "john",
|
||||
filters: { user_id: "john" },
|
||||
});
|
||||
console.log("All memories:", allMemories);
|
||||
|
||||
// Search for memories
|
||||
console.log("\nSearching memories...");
|
||||
const searchResult = await memory.search("What do you know about Paris?", {
|
||||
userId: "john",
|
||||
filters: { user_id: "john" },
|
||||
});
|
||||
console.log("Search results:", searchResult);
|
||||
|
||||
|
||||
@@ -22,7 +22,7 @@ export const DEFAULT_MEMORY_CONFIG: MemoryConfig = {
|
||||
config: {
|
||||
baseURL: "https://api.openai.com/v1",
|
||||
apiKey: process.env.OPENAI_API_KEY || "",
|
||||
model: "gpt-4-turbo-preview",
|
||||
model: "gpt-4.1-nano-2025-04-14",
|
||||
modelProperties: undefined,
|
||||
},
|
||||
},
|
||||
|
||||
@@ -133,9 +133,6 @@ export class ConfigManager {
|
||||
userConfig.historyStore?.config?.historyDbPath ||
|
||||
DEFAULT_MEMORY_CONFIG.historyStore?.config?.historyDbPath,
|
||||
customInstructions: userConfig.customInstructions,
|
||||
graphStore: userConfig.graphStore
|
||||
? { ...userConfig.graphStore }
|
||||
: undefined,
|
||||
historyStore: (() => {
|
||||
const defaultHistoryStore = DEFAULT_MEMORY_CONFIG.historyStore!;
|
||||
const historyProvider =
|
||||
|
||||
@@ -35,13 +35,23 @@ export class AzureOpenAIEmbedder implements Embedder {
|
||||
}
|
||||
|
||||
async embedBatch(texts: string[]): Promise<number[][]> {
|
||||
const response = await this.client.embeddings.create({
|
||||
model: this.model,
|
||||
input: texts,
|
||||
...(this.embeddingDims !== undefined && {
|
||||
dimensions: this.embeddingDims,
|
||||
}),
|
||||
});
|
||||
return response.data.map((item) => item.embedding);
|
||||
const MAX_BATCH = 100;
|
||||
const allEmbeddings: number[][] = [];
|
||||
for (let i = 0; i < texts.length; i += MAX_BATCH) {
|
||||
const chunk = texts.slice(i, i + MAX_BATCH);
|
||||
const response = await this.client.embeddings.create({
|
||||
model: this.model,
|
||||
input: chunk,
|
||||
...(this.embeddingDims !== undefined && {
|
||||
dimensions: this.embeddingDims,
|
||||
}),
|
||||
});
|
||||
allEmbeddings.push(
|
||||
...response.data
|
||||
.sort((a, b) => a.index - b.index)
|
||||
.map((item) => item.embedding),
|
||||
);
|
||||
}
|
||||
return allEmbeddings;
|
||||
}
|
||||
}
|
||||
|
||||
@@ -28,13 +28,23 @@ export class OpenAIEmbedder implements Embedder {
|
||||
}
|
||||
|
||||
async embedBatch(texts: string[]): Promise<number[][]> {
|
||||
const response = await this.openai.embeddings.create({
|
||||
model: this.model,
|
||||
input: texts,
|
||||
...(this.embeddingDims !== undefined && {
|
||||
dimensions: this.embeddingDims,
|
||||
}),
|
||||
});
|
||||
return response.data.map((item) => item.embedding);
|
||||
const MAX_BATCH = 100;
|
||||
const allEmbeddings: number[][] = [];
|
||||
for (let i = 0; i < texts.length; i += MAX_BATCH) {
|
||||
const chunk = texts.slice(i, i + MAX_BATCH);
|
||||
const response = await this.openai.embeddings.create({
|
||||
model: this.model,
|
||||
input: chunk,
|
||||
...(this.embeddingDims !== undefined && {
|
||||
dimensions: this.embeddingDims,
|
||||
}),
|
||||
});
|
||||
allEmbeddings.push(
|
||||
...response.data
|
||||
.sort((a, b) => a.index - b.index)
|
||||
.map((item) => item.embedding),
|
||||
);
|
||||
}
|
||||
return allEmbeddings;
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,30 +0,0 @@
|
||||
import { LLMConfig } from "../types";
|
||||
|
||||
export interface Neo4jConfig {
|
||||
url: string | null;
|
||||
username: string | null;
|
||||
password: string | null;
|
||||
}
|
||||
|
||||
export interface GraphStoreConfig {
|
||||
provider: string;
|
||||
config: Neo4jConfig;
|
||||
llm?: LLMConfig;
|
||||
customInstructions?: string;
|
||||
}
|
||||
|
||||
export function validateNeo4jConfig(config: Neo4jConfig): void {
|
||||
const { url, username, password } = config;
|
||||
if (!url || !username || !password) {
|
||||
throw new Error("Please provide 'url', 'username' and 'password'.");
|
||||
}
|
||||
}
|
||||
|
||||
export function validateGraphStoreConfig(config: GraphStoreConfig): void {
|
||||
const { provider } = config;
|
||||
if (provider === "neo4j") {
|
||||
validateNeo4jConfig(config.config);
|
||||
} else {
|
||||
throw new Error(`Unsupported graph store provider: ${provider}`);
|
||||
}
|
||||
}
|
||||
@@ -1,267 +0,0 @@
|
||||
import { z } from "zod";
|
||||
|
||||
export interface GraphToolParameters {
|
||||
source: string;
|
||||
destination: string;
|
||||
relationship: string;
|
||||
source_type?: string;
|
||||
destination_type?: string;
|
||||
}
|
||||
|
||||
export interface GraphEntitiesParameters {
|
||||
entities: Array<{
|
||||
entity: string;
|
||||
entity_type: string;
|
||||
}>;
|
||||
}
|
||||
|
||||
export interface GraphRelationsParameters {
|
||||
entities: Array<{
|
||||
source: string;
|
||||
relationship: string;
|
||||
destination: string;
|
||||
}>;
|
||||
}
|
||||
|
||||
// --- Zod Schemas for Tool Arguments ---
|
||||
|
||||
// Schema for simple relationship arguments (Update, Delete)
|
||||
export const GraphSimpleRelationshipArgsSchema = z.object({
|
||||
source: z
|
||||
.string()
|
||||
.describe("The identifier of the source node in the relationship."),
|
||||
relationship: z
|
||||
.string()
|
||||
.describe("The relationship between the source and destination nodes."),
|
||||
destination: z
|
||||
.string()
|
||||
.describe("The identifier of the destination node in the relationship."),
|
||||
});
|
||||
|
||||
// Schema for adding a relationship (includes types)
|
||||
export const GraphAddRelationshipArgsSchema =
|
||||
GraphSimpleRelationshipArgsSchema.extend({
|
||||
source_type: z
|
||||
.string()
|
||||
.describe("The type or category of the source node."),
|
||||
destination_type: z
|
||||
.string()
|
||||
.describe("The type or category of the destination node."),
|
||||
});
|
||||
|
||||
// Schema for extracting entities
|
||||
export const GraphExtractEntitiesArgsSchema = z.object({
|
||||
entities: z
|
||||
.array(
|
||||
z.object({
|
||||
entity: z.string().describe("The name or identifier of the entity."),
|
||||
entity_type: z.string().describe("The type or category of the entity."),
|
||||
}),
|
||||
)
|
||||
.describe("An array of entities with their types."),
|
||||
});
|
||||
|
||||
// Schema for establishing relationships
|
||||
export const GraphRelationsArgsSchema = z.object({
|
||||
entities: z
|
||||
.array(GraphSimpleRelationshipArgsSchema)
|
||||
.describe("An array of relationships (source, relationship, destination)."),
|
||||
});
|
||||
|
||||
// --- Tool Definitions (using JSON schema, keep as is) ---
|
||||
|
||||
// Note: The tool definitions themselves still use JSON schema format
|
||||
// as expected by the LLM APIs. The Zod schemas above are for internal
|
||||
// validation and potentially for use with Langchain's .withStructuredOutput
|
||||
// if we adapt it to handle tool calls via schema.
|
||||
|
||||
export const UPDATE_MEMORY_TOOL_GRAPH = {
|
||||
type: "function",
|
||||
function: {
|
||||
name: "update_graph_memory",
|
||||
description:
|
||||
"Update the relationship key of an existing graph memory based on new information.",
|
||||
parameters: {
|
||||
type: "object",
|
||||
properties: {
|
||||
source: {
|
||||
type: "string",
|
||||
description:
|
||||
"The identifier of the source node in the relationship to be updated.",
|
||||
},
|
||||
destination: {
|
||||
type: "string",
|
||||
description:
|
||||
"The identifier of the destination node in the relationship to be updated.",
|
||||
},
|
||||
relationship: {
|
||||
type: "string",
|
||||
description:
|
||||
"The new or updated relationship between the source and destination nodes.",
|
||||
},
|
||||
},
|
||||
required: ["source", "destination", "relationship"],
|
||||
additionalProperties: false,
|
||||
},
|
||||
},
|
||||
};
|
||||
|
||||
export const ADD_MEMORY_TOOL_GRAPH = {
|
||||
type: "function",
|
||||
function: {
|
||||
name: "add_graph_memory",
|
||||
description: "Add a new graph memory to the knowledge graph.",
|
||||
parameters: {
|
||||
type: "object",
|
||||
properties: {
|
||||
source: {
|
||||
type: "string",
|
||||
description:
|
||||
"The identifier of the source node in the new relationship.",
|
||||
},
|
||||
destination: {
|
||||
type: "string",
|
||||
description:
|
||||
"The identifier of the destination node in the new relationship.",
|
||||
},
|
||||
relationship: {
|
||||
type: "string",
|
||||
description:
|
||||
"The type of relationship between the source and destination nodes.",
|
||||
},
|
||||
source_type: {
|
||||
type: "string",
|
||||
description: "The type or category of the source node.",
|
||||
},
|
||||
destination_type: {
|
||||
type: "string",
|
||||
description: "The type or category of the destination node.",
|
||||
},
|
||||
},
|
||||
required: [
|
||||
"source",
|
||||
"destination",
|
||||
"relationship",
|
||||
"source_type",
|
||||
"destination_type",
|
||||
],
|
||||
additionalProperties: false,
|
||||
},
|
||||
},
|
||||
};
|
||||
|
||||
export const NOOP_TOOL = {
|
||||
type: "function",
|
||||
function: {
|
||||
name: "noop",
|
||||
description: "No operation should be performed to the graph entities.",
|
||||
parameters: {
|
||||
type: "object",
|
||||
properties: {},
|
||||
required: [],
|
||||
additionalProperties: false,
|
||||
},
|
||||
},
|
||||
};
|
||||
|
||||
export const RELATIONS_TOOL = {
|
||||
type: "function",
|
||||
function: {
|
||||
name: "establish_relationships",
|
||||
description:
|
||||
"Establish relationships among the entities based on the provided text.",
|
||||
parameters: {
|
||||
type: "object",
|
||||
properties: {
|
||||
entities: {
|
||||
type: "array",
|
||||
items: {
|
||||
type: "object",
|
||||
properties: {
|
||||
source: {
|
||||
type: "string",
|
||||
description: "The source entity of the relationship.",
|
||||
},
|
||||
relationship: {
|
||||
type: "string",
|
||||
description:
|
||||
"The relationship between the source and destination entities.",
|
||||
},
|
||||
destination: {
|
||||
type: "string",
|
||||
description: "The destination entity of the relationship.",
|
||||
},
|
||||
},
|
||||
required: ["source", "relationship", "destination"],
|
||||
additionalProperties: false,
|
||||
},
|
||||
},
|
||||
},
|
||||
required: ["entities"],
|
||||
additionalProperties: false,
|
||||
},
|
||||
},
|
||||
};
|
||||
|
||||
export const EXTRACT_ENTITIES_TOOL = {
|
||||
type: "function",
|
||||
function: {
|
||||
name: "extract_entities",
|
||||
description: "Extract entities and their types from the text.",
|
||||
parameters: {
|
||||
type: "object",
|
||||
properties: {
|
||||
entities: {
|
||||
type: "array",
|
||||
items: {
|
||||
type: "object",
|
||||
properties: {
|
||||
entity: {
|
||||
type: "string",
|
||||
description: "The name or identifier of the entity.",
|
||||
},
|
||||
entity_type: {
|
||||
type: "string",
|
||||
description: "The type or category of the entity.",
|
||||
},
|
||||
},
|
||||
required: ["entity", "entity_type"],
|
||||
additionalProperties: false,
|
||||
},
|
||||
description: "An array of entities with their types.",
|
||||
},
|
||||
},
|
||||
required: ["entities"],
|
||||
additionalProperties: false,
|
||||
},
|
||||
},
|
||||
};
|
||||
|
||||
export const DELETE_MEMORY_TOOL_GRAPH = {
|
||||
type: "function",
|
||||
function: {
|
||||
name: "delete_graph_memory",
|
||||
description: "Delete the relationship between two nodes.",
|
||||
parameters: {
|
||||
type: "object",
|
||||
properties: {
|
||||
source: {
|
||||
type: "string",
|
||||
description: "The identifier of the source node in the relationship.",
|
||||
},
|
||||
relationship: {
|
||||
type: "string",
|
||||
description:
|
||||
"The existing relationship between the source and destination nodes that needs to be deleted.",
|
||||
},
|
||||
destination: {
|
||||
type: "string",
|
||||
description:
|
||||
"The identifier of the destination node in the relationship.",
|
||||
},
|
||||
},
|
||||
required: ["source", "relationship", "destination"],
|
||||
additionalProperties: false,
|
||||
},
|
||||
},
|
||||
};
|
||||
@@ -1,116 +0,0 @@
|
||||
export const UPDATE_GRAPH_PROMPT = `
|
||||
You are an AI expert specializing in graph memory management and optimization. Your task is to analyze existing graph memories alongside new information, and update the relationships in the memory list to ensure the most accurate, current, and coherent representation of knowledge.
|
||||
|
||||
Input:
|
||||
1. Existing Graph Memories: A list of current graph memories, each containing source, target, and relationship information.
|
||||
2. New Graph Memory: Fresh information to be integrated into the existing graph structure.
|
||||
|
||||
Guidelines:
|
||||
1. Identification: Use the source and target as primary identifiers when matching existing memories with new information.
|
||||
2. Conflict Resolution:
|
||||
- If new information contradicts an existing memory:
|
||||
a) For matching source and target but differing content, update the relationship of the existing memory.
|
||||
b) If the new memory provides more recent or accurate information, update the existing memory accordingly.
|
||||
3. Comprehensive Review: Thoroughly examine each existing graph memory against the new information, updating relationships as necessary. Multiple updates may be required.
|
||||
4. Consistency: Maintain a uniform and clear style across all memories. Each entry should be concise yet comprehensive.
|
||||
5. Semantic Coherence: Ensure that updates maintain or improve the overall semantic structure of the graph.
|
||||
6. Temporal Awareness: If timestamps are available, consider the recency of information when making updates.
|
||||
7. Relationship Refinement: Look for opportunities to refine relationship descriptions for greater precision or clarity.
|
||||
8. Redundancy Elimination: Identify and merge any redundant or highly similar relationships that may result from the update.
|
||||
|
||||
Memory Format:
|
||||
source -- RELATIONSHIP -- destination
|
||||
|
||||
Task Details:
|
||||
======= Existing Graph Memories:=======
|
||||
{existing_memories}
|
||||
|
||||
======= New Graph Memory:=======
|
||||
{new_memories}
|
||||
|
||||
Output:
|
||||
Provide a list of update instructions, each specifying the source, target, and the new relationship to be set. Only include memories that require updates.
|
||||
`;
|
||||
|
||||
export const EXTRACT_RELATIONS_PROMPT = `
|
||||
You are an advanced algorithm designed to extract structured information from text to construct knowledge graphs. Your goal is to capture comprehensive and accurate information. Follow these key principles:
|
||||
|
||||
1. Extract only explicitly stated information from the text.
|
||||
2. Establish relationships among the entities provided.
|
||||
3. Use "USER_ID" as the source entity for any self-references (e.g., "I," "me," "my," etc.) in user messages.
|
||||
CUSTOM_PROMPT
|
||||
|
||||
Relationships:
|
||||
- Use consistent, general, and timeless relationship types.
|
||||
- Example: Prefer "professor" over "became_professor."
|
||||
- Relationships should only be established among the entities explicitly mentioned in the user message.
|
||||
|
||||
Entity Consistency:
|
||||
- Ensure that relationships are coherent and logically align with the context of the message.
|
||||
- Maintain consistent naming for entities across the extracted data.
|
||||
|
||||
Strive to construct a coherent and easily understandable knowledge graph by eshtablishing all the relationships among the entities and adherence to the user's context.
|
||||
|
||||
Adhere strictly to these guidelines to ensure high-quality knowledge graph extraction.
|
||||
`;
|
||||
|
||||
export const DELETE_RELATIONS_SYSTEM_PROMPT = `
|
||||
You are a graph memory manager specializing in identifying, managing, and optimizing relationships within graph-based memories. Your primary task is to analyze a list of existing relationships and determine which ones should be deleted based on the new information provided.
|
||||
Input:
|
||||
1. Existing Graph Memories: A list of current graph memories, each containing source, relationship, and destination information.
|
||||
2. New Text: The new information to be integrated into the existing graph structure.
|
||||
3. Use "USER_ID" as node for any self-references (e.g., "I," "me," "my," etc.) in user messages.
|
||||
|
||||
Guidelines:
|
||||
1. Identification: Use the new information to evaluate existing relationships in the memory graph.
|
||||
2. Deletion Criteria: Delete a relationship only if it meets at least one of these conditions:
|
||||
- Outdated or Inaccurate: The new information is more recent or accurate.
|
||||
- Contradictory: The new information conflicts with or negates the existing information.
|
||||
3. DO NOT DELETE if their is a possibility of same type of relationship but different destination nodes.
|
||||
4. Comprehensive Analysis:
|
||||
- Thoroughly examine each existing relationship against the new information and delete as necessary.
|
||||
- Multiple deletions may be required based on the new information.
|
||||
5. Semantic Integrity:
|
||||
- Ensure that deletions maintain or improve the overall semantic structure of the graph.
|
||||
- Avoid deleting relationships that are NOT contradictory/outdated to the new information.
|
||||
6. Temporal Awareness: Prioritize recency when timestamps are available.
|
||||
7. Necessity Principle: Only DELETE relationships that must be deleted and are contradictory/outdated to the new information to maintain an accurate and coherent memory graph.
|
||||
|
||||
Note: DO NOT DELETE if their is a possibility of same type of relationship but different destination nodes.
|
||||
|
||||
For example:
|
||||
Existing Memory: alice -- loves_to_eat -- pizza
|
||||
New Information: Alice also loves to eat burger.
|
||||
|
||||
Do not delete in the above example because there is a possibility that Alice loves to eat both pizza and burger.
|
||||
|
||||
Memory Format:
|
||||
source -- relationship -- destination
|
||||
|
||||
Provide a list of deletion instructions, each specifying the relationship to be deleted.
|
||||
|
||||
Respond in JSON format.
|
||||
`;
|
||||
|
||||
export function getDeleteMessages(
|
||||
existingMemoriesString: string,
|
||||
data: string,
|
||||
userId: string,
|
||||
): [string, string] {
|
||||
return [
|
||||
DELETE_RELATIONS_SYSTEM_PROMPT.replace("USER_ID", userId),
|
||||
`Here are the existing memories: ${existingMemoriesString} \n\n New Information: ${data}`,
|
||||
];
|
||||
}
|
||||
|
||||
export function formatEntities(
|
||||
entities: Array<{
|
||||
source: string;
|
||||
relationship: string;
|
||||
destination: string;
|
||||
}>,
|
||||
): string {
|
||||
return entities
|
||||
.map((e) => `${e.source} -- ${e.relationship} -- ${e.destination}`)
|
||||
.join("\n");
|
||||
}
|
||||
@@ -10,12 +10,6 @@ import { LLM, LLMResponse } from "./base";
|
||||
import { LLMConfig, Message } from "../types/index";
|
||||
// Import the schemas directly into LangchainLLM
|
||||
import { FactRetrievalSchema, MemoryUpdateSchema } from "../prompts";
|
||||
// Import graph tool argument schemas
|
||||
import {
|
||||
GraphExtractEntitiesArgsSchema,
|
||||
GraphRelationsArgsSchema,
|
||||
GraphSimpleRelationshipArgsSchema, // Used for delete tool
|
||||
} from "../graphs/tools";
|
||||
|
||||
const convertToLangchainMessages = (messages: Message[]): BaseMessage[] => {
|
||||
return messages.map((msg) => {
|
||||
@@ -73,28 +67,14 @@ export class LangchainLLM implements LLM {
|
||||
const invokeOptions: Record<string, any> = {};
|
||||
let isStructuredOutput = false;
|
||||
let selectedSchema: z.ZodSchema<any> | null = null;
|
||||
let isToolCallResponse = false;
|
||||
|
||||
// --- Internal Schema Selection Logic (runs regardless of response_format) ---
|
||||
const systemPromptContent =
|
||||
(messages.find((m) => m.role === "system")?.content as string) || "";
|
||||
const userPromptContent =
|
||||
(messages.find((m) => m.role === "user")?.content as string) || "";
|
||||
const toolNames = tools?.map((t) => t.function.name) || [];
|
||||
|
||||
// Prioritize tool call argument schemas
|
||||
if (toolNames.includes("extract_entities")) {
|
||||
selectedSchema = GraphExtractEntitiesArgsSchema;
|
||||
isToolCallResponse = true;
|
||||
} else if (toolNames.includes("establish_relationships")) {
|
||||
selectedSchema = GraphRelationsArgsSchema;
|
||||
isToolCallResponse = true;
|
||||
} else if (toolNames.includes("delete_graph_memory")) {
|
||||
selectedSchema = GraphSimpleRelationshipArgsSchema;
|
||||
isToolCallResponse = true;
|
||||
}
|
||||
// Check for memory prompts if no tool schema matched
|
||||
else if (
|
||||
// Check for memory prompts
|
||||
if (
|
||||
systemPromptContent.includes("Personal Information Organizer") &&
|
||||
systemPromptContent.includes("extract relevant pieces of information")
|
||||
) {
|
||||
@@ -111,26 +91,18 @@ export class LangchainLLM implements LLM {
|
||||
selectedSchema &&
|
||||
typeof (this.llmInstance as any).withStructuredOutput === "function"
|
||||
) {
|
||||
// Apply if a schema was selected (for memory or single tool calls)
|
||||
if (
|
||||
!isToolCallResponse ||
|
||||
(isToolCallResponse && tools && tools.length === 1)
|
||||
) {
|
||||
try {
|
||||
runnable = (this.llmInstance as any).withStructuredOutput(
|
||||
selectedSchema,
|
||||
{ name: tools?.[0]?.function.name },
|
||||
);
|
||||
isStructuredOutput = true;
|
||||
} catch (e) {
|
||||
isStructuredOutput = false; // Ensure flag is false on error
|
||||
// No fallback to response_format here unless explicitly passed
|
||||
if (response_format?.type === "json_object") {
|
||||
invokeOptions.response_format = { type: "json_object" };
|
||||
}
|
||||
try {
|
||||
runnable = (this.llmInstance as any).withStructuredOutput(
|
||||
selectedSchema,
|
||||
{ name: tools?.[0]?.function.name },
|
||||
);
|
||||
isStructuredOutput = true;
|
||||
} catch (e) {
|
||||
isStructuredOutput = false; // Ensure flag is false on error
|
||||
// No fallback to response_format here unless explicitly passed
|
||||
if (response_format?.type === "json_object") {
|
||||
invokeOptions.response_format = { type: "json_object" };
|
||||
}
|
||||
} else if (isToolCallResponse) {
|
||||
// If multiple tools, don't apply structured output, handle via tool binding below
|
||||
}
|
||||
} else if (selectedSchema && response_format?.type === "json_object") {
|
||||
// Schema selected, but no .withStructuredOutput. Try basic response_format only if explicitly requested.
|
||||
@@ -164,37 +136,9 @@ export class LangchainLLM implements LLM {
|
||||
try {
|
||||
const response = await runnable.invoke(langchainMessages, invokeOptions);
|
||||
|
||||
if (isStructuredOutput && !isToolCallResponse) {
|
||||
if (isStructuredOutput) {
|
||||
// Memory prompt with structured output
|
||||
return JSON.stringify(response);
|
||||
} else if (isStructuredOutput && isToolCallResponse) {
|
||||
// Tool call with structured arguments
|
||||
if (response?.tool_calls && Array.isArray(response.tool_calls)) {
|
||||
const mappedToolCalls = response.tool_calls.map((call: any) => ({
|
||||
name: call.name || tools?.[0]?.function.name || "unknown_tool",
|
||||
arguments:
|
||||
typeof call.args === "string"
|
||||
? call.args
|
||||
: JSON.stringify(call.args),
|
||||
}));
|
||||
return {
|
||||
content: response.content || "",
|
||||
role: "assistant",
|
||||
toolCalls: mappedToolCalls,
|
||||
};
|
||||
} else {
|
||||
// Direct object response for tool args
|
||||
return {
|
||||
content: "",
|
||||
role: "assistant",
|
||||
toolCalls: [
|
||||
{
|
||||
name: tools?.[0]?.function.name || "unknown_tool",
|
||||
arguments: JSON.stringify(response),
|
||||
},
|
||||
],
|
||||
};
|
||||
}
|
||||
} else if (
|
||||
response &&
|
||||
response.tool_calls &&
|
||||
|
||||
@@ -1,675 +0,0 @@
|
||||
import neo4j, { Driver } from "neo4j-driver";
|
||||
import { BM25 } from "../utils/bm25";
|
||||
import { GraphStoreConfig } from "../graphs/configs";
|
||||
import { MemoryConfig } from "../types";
|
||||
import { EmbedderFactory, LLMFactory } from "../utils/factory";
|
||||
import { Embedder } from "../embeddings/base";
|
||||
import { LLM } from "../llms/base";
|
||||
import {
|
||||
DELETE_MEMORY_TOOL_GRAPH,
|
||||
EXTRACT_ENTITIES_TOOL,
|
||||
RELATIONS_TOOL,
|
||||
} from "../graphs/tools";
|
||||
import { EXTRACT_RELATIONS_PROMPT, getDeleteMessages } from "../graphs/utils";
|
||||
import { logger } from "../utils/logger";
|
||||
|
||||
interface SearchOutput {
|
||||
source: string;
|
||||
source_id: string;
|
||||
relationship: string;
|
||||
relation_id: string;
|
||||
destination: string;
|
||||
destination_id: string;
|
||||
similarity: number;
|
||||
}
|
||||
|
||||
interface ToolCall {
|
||||
name: string;
|
||||
arguments: string;
|
||||
}
|
||||
|
||||
interface LLMResponse {
|
||||
toolCalls?: ToolCall[];
|
||||
}
|
||||
|
||||
interface Tool {
|
||||
type: string;
|
||||
function: {
|
||||
name: string;
|
||||
description: string;
|
||||
parameters: Record<string, any>;
|
||||
};
|
||||
}
|
||||
|
||||
interface GraphMemoryResult {
|
||||
deleted_entities: any[];
|
||||
added_entities: any[];
|
||||
relations?: any[];
|
||||
}
|
||||
|
||||
export class MemoryGraph {
|
||||
private config: MemoryConfig;
|
||||
private graph: Driver;
|
||||
private embeddingModel: Embedder;
|
||||
private llm: LLM;
|
||||
private structuredLlm: LLM;
|
||||
private llmProvider: string;
|
||||
private threshold: number;
|
||||
|
||||
constructor(config: MemoryConfig) {
|
||||
this.config = config;
|
||||
if (
|
||||
!config.graphStore?.config?.url ||
|
||||
!config.graphStore?.config?.username ||
|
||||
!config.graphStore?.config?.password
|
||||
) {
|
||||
throw new Error("Neo4j configuration is incomplete");
|
||||
}
|
||||
|
||||
this.graph = neo4j.driver(
|
||||
config.graphStore.config.url,
|
||||
neo4j.auth.basic(
|
||||
config.graphStore.config.username,
|
||||
config.graphStore.config.password,
|
||||
),
|
||||
);
|
||||
|
||||
this.embeddingModel = EmbedderFactory.create(
|
||||
this.config.embedder.provider,
|
||||
this.config.embedder.config,
|
||||
);
|
||||
|
||||
this.llmProvider = "openai";
|
||||
let llmConfig = this.config.llm.config;
|
||||
|
||||
if (this.config.llm?.provider) {
|
||||
this.llmProvider = this.config.llm.provider;
|
||||
}
|
||||
if (this.config.graphStore?.llm?.provider) {
|
||||
this.llmProvider = this.config.graphStore.llm.provider;
|
||||
llmConfig = this.config.graphStore.llm.config ?? llmConfig;
|
||||
}
|
||||
|
||||
this.llm = LLMFactory.create(this.llmProvider, llmConfig);
|
||||
this.structuredLlm = LLMFactory.create(this.llmProvider, llmConfig);
|
||||
this.threshold = 0.7;
|
||||
}
|
||||
|
||||
async add(
|
||||
data: string,
|
||||
filters: Record<string, any>,
|
||||
): Promise<GraphMemoryResult> {
|
||||
const entityTypeMap = await this._retrieveNodesFromData(data, filters);
|
||||
|
||||
const toBeAdded = await this._establishNodesRelationsFromData(
|
||||
data,
|
||||
filters,
|
||||
entityTypeMap,
|
||||
);
|
||||
|
||||
const searchOutput = await this._searchGraphDb(
|
||||
Object.keys(entityTypeMap),
|
||||
filters,
|
||||
);
|
||||
|
||||
const toBeDeleted = await this._getDeleteEntitiesFromSearchOutput(
|
||||
searchOutput,
|
||||
data,
|
||||
filters,
|
||||
);
|
||||
|
||||
const deletedEntities = await this._deleteEntities(
|
||||
toBeDeleted,
|
||||
filters["userId"],
|
||||
);
|
||||
|
||||
const addedEntities = await this._addEntities(
|
||||
toBeAdded,
|
||||
filters["userId"],
|
||||
entityTypeMap,
|
||||
);
|
||||
|
||||
return {
|
||||
deleted_entities: deletedEntities,
|
||||
added_entities: addedEntities,
|
||||
relations: toBeAdded,
|
||||
};
|
||||
}
|
||||
|
||||
async search(query: string, filters: Record<string, any>, topK = 100) {
|
||||
const entityTypeMap = await this._retrieveNodesFromData(query, filters);
|
||||
const searchOutput = await this._searchGraphDb(
|
||||
Object.keys(entityTypeMap),
|
||||
filters,
|
||||
);
|
||||
|
||||
if (!searchOutput.length) {
|
||||
return [];
|
||||
}
|
||||
|
||||
const searchOutputsSequence = searchOutput.map((item) => [
|
||||
item.source,
|
||||
item.relationship,
|
||||
item.destination,
|
||||
]);
|
||||
|
||||
const bm25 = new BM25(searchOutputsSequence);
|
||||
const tokenizedQuery = query.split(" ");
|
||||
const rerankedResults = bm25.search(tokenizedQuery).slice(0, 5);
|
||||
|
||||
const searchResults = rerankedResults.map((item) => ({
|
||||
source: item[0],
|
||||
relationship: item[1],
|
||||
destination: item[2],
|
||||
}));
|
||||
|
||||
logger.info(`Returned ${searchResults.length} search results`);
|
||||
return searchResults;
|
||||
}
|
||||
|
||||
async deleteAll(filters: Record<string, any>) {
|
||||
const session = this.graph.session();
|
||||
try {
|
||||
await session.run("MATCH (n {user_id: $user_id}) DETACH DELETE n", {
|
||||
user_id: filters["userId"],
|
||||
});
|
||||
} finally {
|
||||
await session.close();
|
||||
}
|
||||
}
|
||||
|
||||
async getAll(filters: Record<string, any>, topK = 100) {
|
||||
const session = this.graph.session();
|
||||
try {
|
||||
const result = await session.run(
|
||||
`
|
||||
MATCH (n {user_id: $user_id})-[r]->(m {user_id: $user_id})
|
||||
RETURN n.name AS source, type(r) AS relationship, m.name AS target
|
||||
LIMIT toInteger($limit)
|
||||
`,
|
||||
{ user_id: filters["userId"], limit: Math.floor(Number(topK)) },
|
||||
);
|
||||
|
||||
const finalResults = result.records.map((record) => ({
|
||||
source: record.get("source"),
|
||||
relationship: record.get("relationship"),
|
||||
target: record.get("target"),
|
||||
}));
|
||||
|
||||
logger.info(`Retrieved ${finalResults.length} relationships`);
|
||||
return finalResults;
|
||||
} finally {
|
||||
await session.close();
|
||||
}
|
||||
}
|
||||
|
||||
private async _retrieveNodesFromData(
|
||||
data: string,
|
||||
filters: Record<string, any>,
|
||||
) {
|
||||
const tools = [EXTRACT_ENTITIES_TOOL] as Tool[];
|
||||
const searchResults = await this.structuredLlm.generateResponse(
|
||||
[
|
||||
{
|
||||
role: "system",
|
||||
content: `You are a smart assistant who understands entities and their types in a given text. If user message contains self reference such as 'I', 'me', 'my' etc. then use ${filters["userId"]} as the source entity. Extract all the entities from the text. ***DO NOT*** answer the question itself if the given text is a question. Respond in JSON format.`,
|
||||
},
|
||||
{ role: "user", content: data },
|
||||
],
|
||||
{ type: "json_object" },
|
||||
tools,
|
||||
);
|
||||
|
||||
let entityTypeMap: Record<string, string> = {};
|
||||
try {
|
||||
if (typeof searchResults !== "string" && searchResults.toolCalls) {
|
||||
for (const call of searchResults.toolCalls) {
|
||||
if (call.name === "extract_entities") {
|
||||
const args = JSON.parse(call.arguments);
|
||||
for (const item of args.entities) {
|
||||
entityTypeMap[item.entity] = item.entity_type;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
} catch (e) {
|
||||
logger.error(`Error in search tool: ${e}`);
|
||||
}
|
||||
|
||||
entityTypeMap = Object.fromEntries(
|
||||
Object.entries(entityTypeMap).map(([k, v]) => [
|
||||
k.toLowerCase().replace(/ /g, "_"),
|
||||
v.toLowerCase().replace(/ /g, "_"),
|
||||
]),
|
||||
);
|
||||
|
||||
logger.debug(`Entity type map: ${JSON.stringify(entityTypeMap)}`);
|
||||
return entityTypeMap;
|
||||
}
|
||||
|
||||
private async _establishNodesRelationsFromData(
|
||||
data: string,
|
||||
filters: Record<string, any>,
|
||||
entityTypeMap: Record<string, string>,
|
||||
) {
|
||||
let messages;
|
||||
if (this.config.graphStore?.customInstructions) {
|
||||
messages = [
|
||||
{
|
||||
role: "system",
|
||||
content:
|
||||
EXTRACT_RELATIONS_PROMPT.replace(
|
||||
"USER_ID",
|
||||
filters["userId"],
|
||||
).replace(
|
||||
"CUSTOM_PROMPT",
|
||||
`4. ${this.config.graphStore.customInstructions}`,
|
||||
) + "\nPlease provide your response in JSON format.",
|
||||
},
|
||||
{ role: "user", content: data },
|
||||
];
|
||||
} else {
|
||||
messages = [
|
||||
{
|
||||
role: "system",
|
||||
content:
|
||||
EXTRACT_RELATIONS_PROMPT.replace("USER_ID", filters["userId"]) +
|
||||
"\nPlease provide your response in JSON format.",
|
||||
},
|
||||
{
|
||||
role: "user",
|
||||
content: `List of entities: ${Object.keys(entityTypeMap)}. \n\nText: ${data}`,
|
||||
},
|
||||
];
|
||||
}
|
||||
|
||||
const tools = [RELATIONS_TOOL] as Tool[];
|
||||
const extractedEntities = await this.structuredLlm.generateResponse(
|
||||
messages,
|
||||
{ type: "json_object" },
|
||||
tools,
|
||||
);
|
||||
|
||||
let entities: any[] = [];
|
||||
if (typeof extractedEntities !== "string" && extractedEntities.toolCalls) {
|
||||
const toolCall = extractedEntities.toolCalls[0];
|
||||
if (toolCall && toolCall.arguments) {
|
||||
const args = JSON.parse(toolCall.arguments);
|
||||
entities = args.entities || [];
|
||||
}
|
||||
}
|
||||
|
||||
entities = this._removeSpacesFromEntities(entities);
|
||||
logger.debug(`Extracted entities: ${JSON.stringify(entities)}`);
|
||||
return entities;
|
||||
}
|
||||
|
||||
private async _searchGraphDb(
|
||||
nodeList: string[],
|
||||
filters: Record<string, any>,
|
||||
topK = 100,
|
||||
): Promise<SearchOutput[]> {
|
||||
const resultRelations: SearchOutput[] = [];
|
||||
const session = this.graph.session();
|
||||
|
||||
try {
|
||||
for (const node of nodeList) {
|
||||
const nEmbedding = await this.embeddingModel.embed(node);
|
||||
|
||||
const cypher = `
|
||||
MATCH (n)
|
||||
WHERE n.embedding IS NOT NULL AND n.user_id = $user_id
|
||||
WITH n,
|
||||
round(reduce(dot = 0.0, i IN range(0, size(n.embedding)-1) | dot + n.embedding[i] * $n_embedding[i]) /
|
||||
(sqrt(reduce(l2 = 0.0, i IN range(0, size(n.embedding)-1) | l2 + n.embedding[i] * n.embedding[i])) *
|
||||
sqrt(reduce(l2 = 0.0, i IN range(0, size($n_embedding)-1) | l2 + $n_embedding[i] * $n_embedding[i]))), 4) AS similarity
|
||||
WHERE similarity >= $threshold
|
||||
MATCH (n)-[r]->(m)
|
||||
RETURN n.name AS source, elementId(n) AS source_id, type(r) AS relationship, elementId(r) AS relation_id, m.name AS destination, elementId(m) AS destination_id, similarity
|
||||
UNION
|
||||
MATCH (n)
|
||||
WHERE n.embedding IS NOT NULL AND n.user_id = $user_id
|
||||
WITH n,
|
||||
round(reduce(dot = 0.0, i IN range(0, size(n.embedding)-1) | dot + n.embedding[i] * $n_embedding[i]) /
|
||||
(sqrt(reduce(l2 = 0.0, i IN range(0, size(n.embedding)-1) | l2 + n.embedding[i] * n.embedding[i])) *
|
||||
sqrt(reduce(l2 = 0.0, i IN range(0, size($n_embedding)-1) | l2 + $n_embedding[i] * $n_embedding[i]))), 4) AS similarity
|
||||
WHERE similarity >= $threshold
|
||||
MATCH (m)-[r]->(n)
|
||||
RETURN m.name AS source, elementId(m) AS source_id, type(r) AS relationship, elementId(r) AS relation_id, n.name AS destination, elementId(n) AS destination_id, similarity
|
||||
ORDER BY similarity DESC
|
||||
LIMIT toInteger($limit)
|
||||
`;
|
||||
|
||||
const result = await session.run(cypher, {
|
||||
n_embedding: nEmbedding,
|
||||
threshold: this.threshold,
|
||||
user_id: filters["userId"],
|
||||
limit: Math.floor(Number(topK)),
|
||||
});
|
||||
|
||||
resultRelations.push(
|
||||
...result.records.map((record) => ({
|
||||
source: record.get("source"),
|
||||
source_id: record.get("source_id").toString(),
|
||||
relationship: record.get("relationship"),
|
||||
relation_id: record.get("relation_id").toString(),
|
||||
destination: record.get("destination"),
|
||||
destination_id: record.get("destination_id").toString(),
|
||||
similarity: record.get("similarity"),
|
||||
})),
|
||||
);
|
||||
}
|
||||
} finally {
|
||||
await session.close();
|
||||
}
|
||||
|
||||
return resultRelations;
|
||||
}
|
||||
|
||||
private async _getDeleteEntitiesFromSearchOutput(
|
||||
searchOutput: SearchOutput[],
|
||||
data: string,
|
||||
filters: Record<string, any>,
|
||||
) {
|
||||
const searchOutputString = searchOutput
|
||||
.map(
|
||||
(item) =>
|
||||
`${item.source} -- ${item.relationship} -- ${item.destination}`,
|
||||
)
|
||||
.join("\n");
|
||||
|
||||
const [systemPrompt, userPrompt] = getDeleteMessages(
|
||||
searchOutputString,
|
||||
data,
|
||||
filters["userId"],
|
||||
);
|
||||
|
||||
const tools = [DELETE_MEMORY_TOOL_GRAPH] as Tool[];
|
||||
const memoryUpdates = await this.structuredLlm.generateResponse(
|
||||
[
|
||||
{ role: "system", content: systemPrompt },
|
||||
{ role: "user", content: userPrompt },
|
||||
],
|
||||
{ type: "json_object" },
|
||||
tools,
|
||||
);
|
||||
|
||||
const toBeDeleted: any[] = [];
|
||||
if (typeof memoryUpdates !== "string" && memoryUpdates.toolCalls) {
|
||||
for (const item of memoryUpdates.toolCalls) {
|
||||
if (item.name === "delete_graph_memory") {
|
||||
toBeDeleted.push(JSON.parse(item.arguments));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
const cleanedToBeDeleted = this._removeSpacesFromEntities(toBeDeleted);
|
||||
logger.debug(
|
||||
`Deleted relationships: ${JSON.stringify(cleanedToBeDeleted)}`,
|
||||
);
|
||||
return cleanedToBeDeleted;
|
||||
}
|
||||
|
||||
private async _deleteEntities(toBeDeleted: any[], userId: string) {
|
||||
const results: any[] = [];
|
||||
const session = this.graph.session();
|
||||
|
||||
try {
|
||||
for (const item of toBeDeleted) {
|
||||
const { source, destination, relationship } = item;
|
||||
|
||||
const cypher = `
|
||||
MATCH (n {name: $source_name, user_id: $user_id})
|
||||
-[r:${relationship}]->
|
||||
(m {name: $dest_name, user_id: $user_id})
|
||||
DELETE r
|
||||
RETURN
|
||||
n.name AS source,
|
||||
m.name AS target,
|
||||
type(r) AS relationship
|
||||
`;
|
||||
|
||||
const result = await session.run(cypher, {
|
||||
source_name: source,
|
||||
dest_name: destination,
|
||||
user_id: userId,
|
||||
});
|
||||
|
||||
results.push(result.records);
|
||||
}
|
||||
} finally {
|
||||
await session.close();
|
||||
}
|
||||
|
||||
return results;
|
||||
}
|
||||
|
||||
private async _addEntities(
|
||||
toBeAdded: any[],
|
||||
userId: string,
|
||||
entityTypeMap: Record<string, string>,
|
||||
) {
|
||||
const results: any[] = [];
|
||||
const session = this.graph.session();
|
||||
|
||||
try {
|
||||
for (const item of toBeAdded) {
|
||||
const { source, destination, relationship } = item;
|
||||
const sourceType = entityTypeMap[source] || "unknown";
|
||||
const destinationType = entityTypeMap[destination] || "unknown";
|
||||
|
||||
const sourceEmbedding = await this.embeddingModel.embed(source);
|
||||
const destEmbedding = await this.embeddingModel.embed(destination);
|
||||
|
||||
const sourceNodeSearchResult = await this._searchSourceNode(
|
||||
sourceEmbedding,
|
||||
userId,
|
||||
);
|
||||
const destinationNodeSearchResult = await this._searchDestinationNode(
|
||||
destEmbedding,
|
||||
userId,
|
||||
);
|
||||
|
||||
let cypher: string;
|
||||
let params: Record<string, any>;
|
||||
|
||||
if (
|
||||
destinationNodeSearchResult.length === 0 &&
|
||||
sourceNodeSearchResult.length > 0
|
||||
) {
|
||||
cypher = `
|
||||
MATCH (source)
|
||||
WHERE elementId(source) = $source_id
|
||||
MERGE (destination:${destinationType} {name: $destination_name, user_id: $user_id})
|
||||
ON CREATE SET
|
||||
destination.created = timestamp(),
|
||||
destination.embedding = $destination_embedding
|
||||
MERGE (source)-[r:${relationship}]->(destination)
|
||||
ON CREATE SET
|
||||
r.created = timestamp()
|
||||
RETURN source.name AS source, type(r) AS relationship, destination.name AS target
|
||||
`;
|
||||
|
||||
params = {
|
||||
source_id: sourceNodeSearchResult[0].elementId,
|
||||
destination_name: destination,
|
||||
destination_embedding: destEmbedding,
|
||||
user_id: userId,
|
||||
};
|
||||
} else if (
|
||||
destinationNodeSearchResult.length > 0 &&
|
||||
sourceNodeSearchResult.length === 0
|
||||
) {
|
||||
cypher = `
|
||||
MATCH (destination)
|
||||
WHERE elementId(destination) = $destination_id
|
||||
MERGE (source:${sourceType} {name: $source_name, user_id: $user_id})
|
||||
ON CREATE SET
|
||||
source.created = timestamp(),
|
||||
source.embedding = $source_embedding
|
||||
MERGE (source)-[r:${relationship}]->(destination)
|
||||
ON CREATE SET
|
||||
r.created = timestamp()
|
||||
RETURN source.name AS source, type(r) AS relationship, destination.name AS target
|
||||
`;
|
||||
|
||||
params = {
|
||||
destination_id: destinationNodeSearchResult[0].elementId,
|
||||
source_name: source,
|
||||
source_embedding: sourceEmbedding,
|
||||
user_id: userId,
|
||||
};
|
||||
} else if (
|
||||
sourceNodeSearchResult.length > 0 &&
|
||||
destinationNodeSearchResult.length > 0
|
||||
) {
|
||||
cypher = `
|
||||
MATCH (source)
|
||||
WHERE elementId(source) = $source_id
|
||||
MATCH (destination)
|
||||
WHERE elementId(destination) = $destination_id
|
||||
MERGE (source)-[r:${relationship}]->(destination)
|
||||
ON CREATE SET
|
||||
r.created_at = timestamp(),
|
||||
r.updated_at = timestamp()
|
||||
RETURN source.name AS source, type(r) AS relationship, destination.name AS target
|
||||
`;
|
||||
|
||||
params = {
|
||||
source_id: sourceNodeSearchResult[0]?.elementId,
|
||||
destination_id: destinationNodeSearchResult[0]?.elementId,
|
||||
user_id: userId,
|
||||
};
|
||||
} else {
|
||||
cypher = `
|
||||
MERGE (n:${sourceType} {name: $source_name, user_id: $user_id})
|
||||
ON CREATE SET n.created = timestamp(), n.embedding = $source_embedding
|
||||
ON MATCH SET n.embedding = $source_embedding
|
||||
MERGE (m:${destinationType} {name: $dest_name, user_id: $user_id})
|
||||
ON CREATE SET m.created = timestamp(), m.embedding = $dest_embedding
|
||||
ON MATCH SET m.embedding = $dest_embedding
|
||||
MERGE (n)-[rel:${relationship}]->(m)
|
||||
ON CREATE SET rel.created = timestamp()
|
||||
RETURN n.name AS source, type(rel) AS relationship, m.name AS target
|
||||
`;
|
||||
|
||||
params = {
|
||||
source_name: source,
|
||||
dest_name: destination,
|
||||
source_embedding: sourceEmbedding,
|
||||
dest_embedding: destEmbedding,
|
||||
user_id: userId,
|
||||
};
|
||||
}
|
||||
|
||||
const result = await session.run(cypher, params);
|
||||
results.push(result.records);
|
||||
}
|
||||
} finally {
|
||||
await session.close();
|
||||
}
|
||||
|
||||
return results;
|
||||
}
|
||||
|
||||
private _removeSpacesFromEntities(entityList: any[]) {
|
||||
return entityList.map((item) => ({
|
||||
...item,
|
||||
source: item.source.toLowerCase().replace(/ /g, "_"),
|
||||
relationship: item.relationship.toLowerCase().replace(/ /g, "_"),
|
||||
destination: item.destination.toLowerCase().replace(/ /g, "_"),
|
||||
}));
|
||||
}
|
||||
|
||||
private async _searchSourceNode(
|
||||
sourceEmbedding: number[],
|
||||
userId: string,
|
||||
threshold = 0.9,
|
||||
) {
|
||||
const session = this.graph.session();
|
||||
try {
|
||||
const cypher = `
|
||||
MATCH (source_candidate)
|
||||
WHERE source_candidate.embedding IS NOT NULL
|
||||
AND source_candidate.user_id = $user_id
|
||||
|
||||
WITH source_candidate,
|
||||
round(
|
||||
reduce(dot = 0.0, i IN range(0, size(source_candidate.embedding)-1) |
|
||||
dot + source_candidate.embedding[i] * $source_embedding[i]) /
|
||||
(sqrt(reduce(l2 = 0.0, i IN range(0, size(source_candidate.embedding)-1) |
|
||||
l2 + source_candidate.embedding[i] * source_candidate.embedding[i])) *
|
||||
sqrt(reduce(l2 = 0.0, i IN range(0, size($source_embedding)-1) |
|
||||
l2 + $source_embedding[i] * $source_embedding[i])))
|
||||
, 4) AS source_similarity
|
||||
WHERE source_similarity >= $threshold
|
||||
|
||||
WITH source_candidate, source_similarity
|
||||
ORDER BY source_similarity DESC
|
||||
LIMIT 1
|
||||
|
||||
RETURN elementId(source_candidate) as element_id
|
||||
`;
|
||||
|
||||
const params = {
|
||||
source_embedding: sourceEmbedding,
|
||||
user_id: userId,
|
||||
threshold,
|
||||
};
|
||||
|
||||
const result = await session.run(cypher, params);
|
||||
|
||||
return result.records.map((record) => ({
|
||||
elementId: record.get("element_id").toString(),
|
||||
}));
|
||||
} finally {
|
||||
await session.close();
|
||||
}
|
||||
}
|
||||
|
||||
private async _searchDestinationNode(
|
||||
destinationEmbedding: number[],
|
||||
userId: string,
|
||||
threshold = 0.9,
|
||||
) {
|
||||
const session = this.graph.session();
|
||||
try {
|
||||
const cypher = `
|
||||
MATCH (destination_candidate)
|
||||
WHERE destination_candidate.embedding IS NOT NULL
|
||||
AND destination_candidate.user_id = $user_id
|
||||
|
||||
WITH destination_candidate,
|
||||
round(
|
||||
reduce(dot = 0.0, i IN range(0, size(destination_candidate.embedding)-1) |
|
||||
dot + destination_candidate.embedding[i] * $destination_embedding[i]) /
|
||||
(sqrt(reduce(l2 = 0.0, i IN range(0, size(destination_candidate.embedding)-1) |
|
||||
l2 + destination_candidate.embedding[i] * destination_candidate.embedding[i])) *
|
||||
sqrt(reduce(l2 = 0.0, i IN range(0, size($destination_embedding)-1) |
|
||||
l2 + $destination_embedding[i] * $destination_embedding[i])))
|
||||
, 4) AS destination_similarity
|
||||
WHERE destination_similarity >= $threshold
|
||||
|
||||
WITH destination_candidate, destination_similarity
|
||||
ORDER BY destination_similarity DESC
|
||||
LIMIT 1
|
||||
|
||||
RETURN elementId(destination_candidate) as element_id
|
||||
`;
|
||||
|
||||
const params = {
|
||||
destination_embedding: destinationEmbedding,
|
||||
user_id: userId,
|
||||
threshold,
|
||||
};
|
||||
|
||||
const result = await session.run(cypher, params);
|
||||
|
||||
return result.records.map((record) => ({
|
||||
elementId: record.get("element_id").toString(),
|
||||
}));
|
||||
} finally {
|
||||
await session.close();
|
||||
}
|
||||
}
|
||||
}
|
||||
+849
-212
File diff suppressed because it is too large
Load Diff
@@ -13,13 +13,15 @@ export interface AddMemoryOptions extends Entity {
|
||||
infer?: boolean;
|
||||
}
|
||||
|
||||
export interface SearchMemoryOptions extends Entity {
|
||||
export interface SearchMemoryOptions {
|
||||
topK?: number;
|
||||
filters?: SearchFilters;
|
||||
threshold?: number;
|
||||
}
|
||||
|
||||
export interface GetAllMemoryOptions {
|
||||
topK?: number;
|
||||
filters?: SearchFilters;
|
||||
}
|
||||
|
||||
export interface GetAllMemoryOptions extends Entity {
|
||||
topK?: number;
|
||||
}
|
||||
|
||||
export interface DeleteAllMemoryOptions extends Entity {}
|
||||
|
||||
@@ -274,6 +274,598 @@ export function getUpdateMemoryMessages(
|
||||
Do not return anything except the JSON format.`;
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// V3 Additive Extraction Prompt
|
||||
// Ported from mem0/configs/prompts.py — ADDITIVE_EXTRACTION_PROMPT
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
export const ADDITIVE_EXTRACTION_PROMPT = `
|
||||
# ROLE
|
||||
|
||||
You are a Memory Extractor — a precise, evidence-bound processor responsible for extracting rich, contextual memories from conversations. Your sole operation is ADD: identify every piece of memorable information and produce self-contained, contextually rich factual statements.
|
||||
|
||||
You extract from BOTH user and assistant messages. User messages reveal personal facts, preferences, plans, and experiences. Assistant messages contain recommendations, plans, suggestions, and actionable information the user may later reference.
|
||||
|
||||
Accuracy and completeness are critical. Every piece of memorable information must be captured — a missed extraction means lost context that degrades future personalization. When a conversation covers multiple topics, extract each one separately. Do not let a dominant topic cause you to miss secondary information.
|
||||
|
||||
# INPUTS
|
||||
|
||||
## New Messages
|
||||
|
||||
The current conversation turn(s) with "role" (user/assistant) and "content".
|
||||
|
||||
Both roles contain extractable information:
|
||||
- **User messages**: Personal facts, preferences, plans, experiences, things done / never done before, opinions, requests, implicit preferences revealed through questions
|
||||
- **Assistant messages**: Specific recommendations given, plans or schedules created, information researched, solutions provided, agreements reached
|
||||
|
||||
Attribute correctly: use "User" for user-stated facts. For assistant-generated content, frame in terms of the user's context (e.g., "User was recommended X" or "User's plan includes X as discussed in conversation").
|
||||
|
||||
Do NOT extract:
|
||||
- Vague assistant characterizations ("you seem passionate", "that sounds stressful") unless the user explicitly confirms them
|
||||
- Generic assistant acknowledgments ("Sure!", "Great question!")
|
||||
- Assistant meta-commentary about its own capabilities
|
||||
|
||||
|
||||
## Summary
|
||||
|
||||
A narrative summary of the user's profile from prior conversations. May be empty for new users. Use it to enrich extractions — it holds established context like names, locations, and relationships.
|
||||
|
||||
|
||||
## Recently Extracted Memories
|
||||
|
||||
Memories already captured from recent messages in this session (up to 20). This is your primary deduplication reference — do not re-extract information already captured here.
|
||||
|
||||
|
||||
## Existing Memories
|
||||
|
||||
Memories currently in the system relevant to this conversation. Formatted as:
|
||||
[{"id": "uuid-string", "text": "..."}, ...]
|
||||
|
||||
Use these ONLY for deduplication and linking — do NOT extract new memories from Existing Memories. Your extractions must come exclusively from New Messages. If new information in New Messages is semantically equivalent to an Existing Memory with no meaningful new context, skip it.
|
||||
|
||||
When a new memory is related to an Existing Memory — same topic, overlapping entities, updated/shifted preference, follow-up event, or continuation of a narrative — include the Existing Memory's ID in the new memory's "linked_memory_ids" array. Your ADD output IDs remain sequential ("0", "1", ...) but linked_memory_ids uses the UUIDs from this list.
|
||||
|
||||
|
||||
IMPORTANT: An existing memory about an entity (e.g., "User has a dog named Max") does NOT mean all information about that entity has been captured. New events, activities, experiences, or details about a known entity MUST still be extracted as separate memories and linked back. Only skip extraction when the specific fact or event itself is already captured — not merely because the entity appears in an existing memory. "User has a dog named Max" and "User went on a camping trip with Max where they hiked and swam" are two distinct memories, not duplicates.
|
||||
|
||||
|
||||
## Last k Messages
|
||||
|
||||
Recent messages (up to 20) preceding New Messages. Use to resolve references and pronouns in New Messages.
|
||||
|
||||
|
||||
## Observation Date
|
||||
|
||||
When the conversation actually took place (e.g., "2023-05-24"). This is your ONLY temporal anchor for resolving time references.
|
||||
|
||||
Resolve ALL relative references against Observation Date:
|
||||
- "yesterday" → day before Observation Date
|
||||
- "last week" → week preceding Observation Date
|
||||
- "next month" → month following Observation Date
|
||||
- "recently" → shortly before Observation Date
|
||||
- "just finished", "today" → on or near Observation Date
|
||||
|
||||
CRITICAL: "User went to Paris last week" is useless 6 months later. "User went to Paris the week of May 15, 2023" is meaningful forever. Always ground relative references to specific dates.
|
||||
|
||||
|
||||
## Current Date
|
||||
|
||||
Today's system date. May be years after Observation Date. Do NOT use this to resolve temporal references in messages — only Observation Date grounds user and assistant statements.
|
||||
|
||||
|
||||
## Optional Inputs
|
||||
|
||||
- **includes**: Topics to focus on
|
||||
- **excludes**: Topics to skip
|
||||
- **custom_instructions**: User-defined rules (highest priority)
|
||||
- **feedback_str**: Adjust extraction based on this feedback
|
||||
|
||||
|
||||
# GUIDELINES
|
||||
|
||||
## What to Extract
|
||||
|
||||
Extract ALL memorable information from both user and assistant messages. Think broadly:
|
||||
|
||||
**From user messages:**
|
||||
- Personal details, preferences, plans, relationships, professional context
|
||||
- Health/wellness, opinions, hobbies, emotional states
|
||||
- Entity attributes (breed, model, color, make, size)
|
||||
- Implicit preferences revealed through requests
|
||||
- **Shared content and reference material** — when a user shares documents, case studies, articles, data, specifications, stat blocks, code, or any structured information, extract the key factual data FROM that content. The user shared it because they want it remembered.
|
||||
- Firsts and milestones — 'first call-out', 'just started', 'recently joined', etc.
|
||||
- Specific foods, meals, and who was present (e.g. 'dinner with mom — salads, sandwiches, homemade desserts').
|
||||
- Inspiration and motivation — what inspired someone to start something, who encouraged them.
|
||||
|
||||
**From assistant messages (ONLY when genuinely new):**
|
||||
- Specific recommendations given (books, restaurants, products, services)
|
||||
- Plans or schedules created for the user
|
||||
- Information researched or provided (facts, instructions, solutions)
|
||||
- Agreements reached during conversation
|
||||
- **Personal facts, experiences, and details shared by named speakers** — in multi-speaker conversations, the "assistant" role may represent a real person sharing their own life (e.g., "Maria: I just got a new cat named Bailey"). Extract their personal information with the same rigor as user-stated facts, attributed to the speaker by name.
|
||||
|
||||
Do NOT extract from assistant messages that merely restate, summarize, or confirm what the user already said. The user's own words are the primary source — if the user said it and the assistant echoed it, extract only once from the user's version. Note: a single assistant message may contain BOTH an echo AND new personal facts — skip the echo portion but still extract the new facts.
|
||||
|
||||
Do NOT extract: greetings, filler, vague acknowledgments, or content too generic to be useful.
|
||||
|
||||
**When in doubt, extract.** A slightly redundant memory is far less costly than a missing one. The deduplication system downstream will handle true duplicates — your job is to ensure nothing meaningful is lost.
|
||||
|
||||
### Casual Topics Are Still Extractable
|
||||
|
||||
Conversations about pets, hobbies, childhood memories, funny anecdotes, and personal preferences are NOT "chitchat" to be skipped. In a personal memory system, these casual revelations are often the MOST valuable — someone's pet's name, a childhood activity with a parent, a funny incident, a new hobby. Only skip messages that are PURELY phatic ("Hi!", "Sounds good!", "Thanks!") with zero informational content.
|
||||
|
||||
### Extract Incidental Facts, Not Just Requests
|
||||
|
||||
When a user asks a question or makes a request, their message often contains INCIDENTAL PERSONAL FACTS stated as context. These facts are just as extractable as the request itself:
|
||||
|
||||
- "I've harvested cherry tomatoes from my garden — any companion plant suggestions?" → Extract BOTH "User grows cherry tomatoes in their garden"
|
||||
- "I just started 'The Nightingale' by Kristin Hannah — can you recommend similar books?" → Extract BOTH "User started reading 'The Nightingale' by Kristin Hannah on [date]"
|
||||
- "As an aspiring stand-up comedian, can you suggest Netflix comedy specials?" → Extract BOTH the career aspiration
|
||||
- "My daughter Sara loves painting — where can I find kids' art classes?" → Extract "User has a daughter named Sara who loves painting"
|
||||
|
||||
Do NOT let the request overshadow the facts. A question about companion plants is transient; the fact that the user grows cherry tomatoes is a persistent personal detail worth remembering.
|
||||
|
||||
**IMPORTANT — Extract ALL dimensions of a conversation.** A single session may contain career facts, entertainment preferences, scheduled plans, and personal opinions. Extract each dimension as a separate memory. Do not let one dominant topic cause you to miss secondary information.
|
||||
|
||||
### Shared Photos and Images
|
||||
|
||||
When a message contains a photo description (e.g., "[Shared photo: ...]" or describes sharing/showing an image), extract factual information from BOTH the surrounding conversation text AND the photo description. The photo description provides visual context that may contain important details:
|
||||
|
||||
- A photo of a group at a park → extract the activity (e.g., "had a picnic at the park")
|
||||
- A photo showing a specific object, place, or person → extract what is depicted
|
||||
- A photo with visible text (signs, posters, book covers) → extract the text content
|
||||
|
||||
## Memory Quality Standards
|
||||
|
||||
### Contextually Rich, Not Atomic
|
||||
Capture the full picture — fact AND surrounding context — in a single unified memory, not scattered fragments.
|
||||
|
||||
Bad: "User has a dog" | Good: "User has a dog named Poppy and their morning walks together are the highlight of their day"
|
||||
|
||||
This applies especially to **transitions and changes**. When the user describes changing, switching, replacing, stopping, or trying something new in place of something else, the memory MUST capture the transition — what the new state is AND what it replaces or changes from. The relationship between old and new is critical context. Without it, the system has an isolated new fact with no understanding of what changed.
|
||||
|
||||
Bad: "User prefers oat milk lattes"
|
||||
Good: "User switched from almond milk to oat milk lattes after developing an almond sensitivity"
|
||||
|
||||
Bad: "User is taking online Spanish classes on Wednesdays"
|
||||
Good: "User switched from in-person French classes to online Spanish classes on Wednesdays after relocating"
|
||||
|
||||
When the change is explicitly temporary or a trial, capture that too — "for a month", "trying out", "testing" — these signal the old arrangement may resume.
|
||||
|
||||
### Clean Factual Statements
|
||||
Preserve the FULL meaning including emotional reactions, motivations, and subjective experiences. Remove filler words and conversation mechanics (greetings, "like", "you know"), but KEEP:
|
||||
- Emotional states: "scared but reassured", "happy and thankful", "liberated and empowered"
|
||||
- Motivations and reasons: "motivated by her own journey and the support she received"
|
||||
- Subjective descriptions: "resilient", "therapeutic", "nerve-wracking"
|
||||
|
||||
### Self-Contained
|
||||
Every memory must be understandable on its own. Replace all pronouns with specific names or "User."
|
||||
|
||||
### Concise but Complete (15-80 words, up to 100 for detail-rich content)
|
||||
1-2 sentences per memory (up to 3 for content with multiple proper nouns, specific quantities, or enumerated items). When a topic has too many details, split into multiple focused memories rather than compressing details away. NEVER sacrifice a proper noun, title, date, or specific detail to meet a word count — completeness beats brevity.
|
||||
|
||||
### Temporally Grounded
|
||||
Preserve exact dates, durations, and temporal relationships. Convert relative → absolute using Observation Date (NOT Current Date). NEVER convert absolute → vague. "18 days" stays "18 days", not "some time."
|
||||
|
||||
### Numerically Precise
|
||||
Preserve exact quantities as stated. "416 pages" stays "416 pages", not "about 400 pages."
|
||||
|
||||
### Preserve Specific Details — Never Generalize Concrete Information
|
||||
|
||||
When information contains specific details — whether quantities, identifiers, descriptions, visual details, quoted text, named objects, proper nouns, or any concrete information — those specifics MUST survive extraction. Replacing a specific detail with a vague category is a critical error.
|
||||
|
||||
#### Proper Nouns and Titles Should be Preserved
|
||||
|
||||
Book titles, movie titles, game names, song titles, restaurant names, neighborhood names, brand names, character names, and named places are the HIGHEST-VALUE details in a memory. Users search by name — a memory without the name is unfindable. ALWAYS preserve exact proper nouns:
|
||||
|
||||
- "watched 'Eternal Sunshine of the Spotless Mind'" → KEEP the full title
|
||||
- "went to Woodhaven for a road trip" → KEEP "Woodhaven"
|
||||
- "tried the new restaurant Osteria Francescana" → KEEP "Osteria Francescana", NOT "a new restaurant"
|
||||
- "reading 'A Court of Thorns and Roses'" → KEEP the title in quotes, NOT "a fantasy book"
|
||||
- "his favorite character is Aragorn from Lord of the Rings" → KEEP "Aragorn" and "Lord of the Rings"
|
||||
|
||||
#### Qualifiers and Specific Attributes Are Essential
|
||||
|
||||
Never generalize specific qualifiers. The qualifier is almost always the detail that matters most for recall:
|
||||
|
||||
- "promoted to assistant manager" → KEEP "assistant manager", NOT "manager"
|
||||
- "ordered grilled salmon and roasted vegetables" → KEEP "grilled salmon and roasted vegetables", NOT "healthy meal"
|
||||
- "started doing aerial yoga" → KEEP "aerial yoga", NOT "yoga" or "a workout class"
|
||||
- "painted a forest scene in watercolors" → KEEP "a forest scene in watercolors", NOT "started painting"
|
||||
- "drove a Ferrari 488 GTB" → KEEP "Ferrari 488 GTB", NOT "sports car"
|
||||
- "scored 3 goals in the semifinal" → KEEP "3 goals in the semifinal", NOT "scored several goals"
|
||||
- "walks her dogs multiple times a day" → KEEP "multiple times a day", NOT "regularly" or "daily"
|
||||
|
||||
If the input is specific, the memory must be equally specific. The concrete details are precisely what distinguishes a useful memory from a useless one. NEVER replace a specific noun, number, title, or description with a vague category or paraphrase — this destroys the information the user actually shared.
|
||||
|
||||
### Meaning-Preserving
|
||||
Capture the EXACT meaning of what was said. Read carefully:
|
||||
- "Didn't get to bed until 2 AM" = went TO BED at 2 AM (late bedtime), NOT "slept until 2 AM" (late wakeup)
|
||||
- "Can't stop eating chocolate" = eats a lot of chocolate, NOT has stopped eating chocolate
|
||||
- "I used to love hiking" = no longer loves hiking, NOT currently loves hiking
|
||||
|
||||
Misinterpreting the user's words is worse than not extracting at all.
|
||||
|
||||
|
||||
## Integrity Rules
|
||||
|
||||
- **No Fabrication**: Every detail must trace to the inputs. If you can't point to where it came from, don't include it.
|
||||
- **No Implicit Attribute Inference**: Don't infer gender, age, ethnicity, etc. from names or context. Only record explicitly stated attributes.
|
||||
- **Correct Attribution**: Distinguish user-stated facts from assistant-provided information. Frame assistant content appropriately.
|
||||
- **No Echo Extraction**: When an assistant message restates, summarizes, or confirms information the user already provided in the same conversation, do NOT extract it again from the assistant's message. Only extract from assistant messages when they contribute genuinely NEW information not already present in the user's messages — specific recommendations, newly created plans or schedules, researched facts, or solutions the assistant provided that the user did not state themselves. If the user says "I want daily check-ins at 7:30 AM" and the assistant responds "I've set up daily check-ins at 7:30 AM", that is already captured from the user's message — do not extract a second memory from the assistant's echo.
|
||||
- **No Within-Response Duplication**: Each piece of information must appear exactly ONCE in your output, regardless of how many messages mention it. Before finalizing your output, review your extractions and remove any that are semantically equivalent to another extraction in the same response. Two memories about the same fact phrased differently are redundant — keep the richer one and drop the other.
|
||||
- **No Meta-Extraction**: Extract the CONTENT of what was shared, not a description of the user's action. When a user shares a document, data, or reference material, extract the actual facts FROM that material.
|
||||
- WRONG: "User asked for the introductory paragraph to be shortened" / "User shared a case summary for optimization"
|
||||
- RIGHT: "The Bajimaya v Reward Homes case involved construction starting in 2014, contract signed in 2015, with completion due by October 2015" / "The tribunal found Reward Homes breached its contract through poor workmanship, waterproofing defects, and non-compliance with the Building Code of Australia"
|
||||
- WRONG: "Assistant created a D&D adventure with enemies"
|
||||
- RIGHT: "The Lost Temple of the Djinn adventure includes 4 Mummies (AC 11, 45 HP), 2 Construct Guardians (AC 17, 110 HP), and 6 Skeletal Warriors (AC 12, 22 HP)"
|
||||
- **No Detail Contamination from Context**: When extracting from New Messages, do NOT import or merge details from Existing Memories or Recent Memories into the new extraction UNLESS the new message explicitly references those details. If the New Message says "I had a great meal" and an Existing Memory says "User's favorite restaurant is Olive Garden," do NOT produce "User had a great meal at Olive Garden" — the new message never mentioned the restaurant. Each extraction must be faithful to its source message only.
|
||||
|
||||
|
||||
## Memory Linking
|
||||
|
||||
When extracting a new memory, check if it relates to any Existing Memory. Add related Existing Memory IDs to "linked_memory_ids". Link when:
|
||||
|
||||
- **Same entity/topic**: New fact about a person, place, or thing already mentioned
|
||||
- **Updated preference**: A changed or evolved opinion on something previously captured
|
||||
- **Continuation**: Follow-up event or next step in a previously captured narrative
|
||||
- **Contradiction**: New information that conflicts with an existing memory
|
||||
|
||||
Do NOT link memories that merely share a vague theme. Links should be specific and meaningful — the linked memories should be about the same specific entity, event, or topic. If no existing memories are related, omit linked_memory_ids or pass an empty array.
|
||||
|
||||
|
||||
# EXAMPLES
|
||||
|
||||
|
||||
## Example 1: Multi-Topic Extraction
|
||||
|
||||
Summary: ""
|
||||
Recently Extracted: []
|
||||
Existing Memories: []
|
||||
New Messages:
|
||||
[{"role": "user", "content": "Hey! I'm Marcus. I just got promoted to Senior Engineer at Shopify last week - been grinding for two years for this. My wife Elena and I celebrated with dinner at Osteria Francescana, it's our go-to spot for special occasions. We're also expecting our first baby in March!"},
|
||||
{"role": "assistant", "content": "Congratulations on everything, Marcus! What exciting times."}]
|
||||
Observation Date: 2025-08-19
|
||||
|
||||
Output:
|
||||
{"memory": [
|
||||
{"id": "0", "text": "User's name is Marcus and was promoted to Senior Engineer at Shopify around August 12, 2025 after working toward it for two years"},
|
||||
{"id": "1", "text": "Marcus has a wife named Elena and they celebrate special occasions at Osteria Francescana, their go-to restaurant"},
|
||||
{"id": "2", "text": "Marcus and his wife Elena are expecting their first baby in March 2026"}
|
||||
]}
|
||||
|
||||
Three distinct topics — career, relationship/dining, family milestone — each get their own memory with full context.
|
||||
|
||||
|
||||
## Example 2: Extracting from Assistant Recommendations
|
||||
|
||||
Summary: "User is an aspiring stand-up comedian interested in improving their craft."
|
||||
Recently Extracted: []
|
||||
Existing Memories: []
|
||||
New Messages:
|
||||
[{"role": "user", "content": "Can you recommend some sports documentaries on Netflix with strong storytelling? I love \\"The Last Dance\\" by Michael Jordan."},
|
||||
{"role": "assistant", "content": "Great taste! Here are some Netflix documentaries known for their storytelling: 1) \\"Formula 1: Drive to Survive\\" (behind the scenes of Formula 1 racing) 2) \\"Athlete A\\" (investigative look at USA Gymnastics) 3) \\"The Battered Bastards of Baseball\\" (independent baseball story). All focus on powerful, narrative-driven sports stories."}]
|
||||
Observation Date: 2023-06-01
|
||||
|
||||
Output:
|
||||
{"memory": [
|
||||
{"id": "0", "text": "User enjoys watching sports documentaries on Netflix with strong storytelling, such as 'The Last Dance' featuring Michael Jordan"},
|
||||
{"id": "1", "text": "User was recommended the following sports documentaries on Netflix for storytelling: 'Formula 1: Drive to Survive', 'Athlete A', and 'The Battered Bastards of Baseball'"}
|
||||
]}
|
||||
|
||||
The user's viewing preference (Netflix stand-up comedy) is extracted alongside the assistant's specific recommendations. Both are valuable for future personalization.
|
||||
|
||||
|
||||
## Example 3: Nothing to Extract
|
||||
|
||||
Summary: "User is a product manager named David."
|
||||
Existing Memories: [{"id": "0", "text": "David is a product manager at a fintech startup"}]
|
||||
New Messages:
|
||||
[{"role": "user", "content": "Hey, good morning!"},
|
||||
{"role": "assistant", "content": "Good morning, David! How can I help you today?"}]
|
||||
Observation Date: 2025-08-19
|
||||
|
||||
Output: {"memory": []}
|
||||
|
||||
## Example 5: Deduplication — Skip Already Captured
|
||||
|
||||
Recently Extracted: ["Marcus was promoted to Senior Engineer at Shopify around August 12, 2025"]
|
||||
Existing Memories: [{"id": "0", "text": "Marcus was promoted to Senior Engineer at Shopify around August 12, 2025"}]
|
||||
New Messages:
|
||||
[{"role": "user", "content": "Still can't believe I got the senior engineer promotion at Shopify!"}]
|
||||
Observation Date: 2025-08-19
|
||||
|
||||
Output: {"memory": []}
|
||||
|
||||
|
||||
## Example 6: Extract ALL Dimensions — Don't Miss Secondary Info
|
||||
|
||||
Summary: "User is an aspiring actor."
|
||||
Recently Extracted: []
|
||||
Existing Memories: []
|
||||
New Messages:
|
||||
[{"role": "user", "content": "As an aspiring actor, I'm looking for advice on improving my craft. Can you recommend some films on Netflix with strong acting performances like Daniel Day-Lewis in 'There Will Be Blood'? I also want to find online resources for acting techniques."},
|
||||
{"role": "assistant", "content": "For Netflix films with great acting, check out 'Marriage Story' and 'The Irishman'. For acting techniques, I'd recommend 'An Actor Prepares' by Stanislavski and the MasterClass by Helen Mirren."}]
|
||||
Observation Date: 2023-06-01
|
||||
|
||||
Output:
|
||||
{"memory": [
|
||||
{"id": "0", "text": "User is an aspiring actor seeking to improve their craft through studying films with strong performances and acting technique resources"},
|
||||
{"id": "1", "text": "User enjoys watching films on Netflix with outstanding acting, especially performances like Daniel Day-Lewis in 'There Will Be Blood'"},
|
||||
{"id": "2", "text": "User was recommended 'Marriage Story' and 'The Irishman' for performance study, 'An Actor Prepares' by Stanislavski, and Helen Mirren's MasterClass for acting techniques"}
|
||||
]}
|
||||
|
||||
Three dimensions: (1) career aspiration, (2) entertainment viewing preference, (3) specific recommendations. Each extracted separately.
|
||||
|
||||
|
||||
## Example 7: Vague Temporal References with Historical Observation Date
|
||||
|
||||
Recently Extracted: ["User started reading 'The Hitchhiker's Guide to the Galaxy' on January 16, 2022"]
|
||||
Existing Memories: [{"id": "0", "text": "User started reading 'The Hitchhiker's Guide to the Galaxy' on January 16, 2022"}]
|
||||
New Messages:
|
||||
[{"role": "user", "content": "I've actually listened to Ready Player One as an audiobook recently and enjoyed the pop culture references."}]
|
||||
Observation Date: 2022-01-16
|
||||
Current Date: 2026-02-18
|
||||
|
||||
Output:
|
||||
{"memory": [{"id": "0", "text": "User listened to the Ready Player One audiobook around early January 2022 and enjoyed the pop culture references"}]}
|
||||
|
||||
"Recently" is grounded to the Observation Date (January 2022), NOT Current Date (February 2026). The Hitchhiker's Guide memory already exists — not re-extracted.
|
||||
|
||||
|
||||
## Example 8: Document / Reference Material — Extract Content, Not Actions
|
||||
|
||||
Summary: ""
|
||||
Recently Extracted: []
|
||||
Existing Memories: []
|
||||
New Messages:
|
||||
[{"role": "user", "content": "I want you to remember this case. If you understand, just say acknowledged. Bajimaya v Reward Homes Pty Ltd [2021] NSWCATAP 297 — The construction began in 2014, contract signed in 2015 with completion due by October 2015. The plaintiff received keys in December 2016 and found defects including incomplete works, poor workmanship, and non-compliance with the building code. The tribunal found the builder breached contract."},
|
||||
{"role": "assistant", "content": "Acknowledged."}]
|
||||
Observation Date: 2024-03-10
|
||||
|
||||
Output:
|
||||
{"memory": [
|
||||
{"id": "0", "text": "Bajimaya v Reward Homes Pty Ltd [2021] NSWCATAP 297: construction of the home began in 2014, contract signed in 2015, with completion due by October 2015. Keys were delivered in December 2016."},
|
||||
{"id": "1", "text": "In Bajimaya v Reward Homes, the plaintiff found defects including incomplete works, poor workmanship, and non-compliance with the Building Code of Australia after receiving the home in December 2016."},
|
||||
{"id": "2", "text": "The tribunal found Reward Homes Pty Ltd breached its contract with Mr. Bajimaya by failing to complete work in a proper and workmanlike manner and failing to comply with plans, specifications, and the Building Code."}
|
||||
]}
|
||||
|
||||
The user shared reference material to be remembered. Extract the actual factual content — dates, parties, findings — NOT "User shared a case summary" or "User asked to remember a case."
|
||||
|
||||
|
||||
## Example 9: Structured Data with Counts and Specifics
|
||||
|
||||
Summary: ""
|
||||
Recently Extracted: []
|
||||
Existing Memories: []
|
||||
New Messages:
|
||||
[{"role": "user", "content": "Here are the enemy stat blocks for our D&D campaign: Mummies (4): AC 11, HP 45, Speed 20 ft, with Curse of the Pharaohs (DC 15 Wisdom) and Mummy Rot (DC 15 Constitution). Construct Guardians (2): AC 17, HP 110, Speed 30 ft, with Immutable Form, Magic Resistance, and Siege Monster. Skeletal Warriors (6): AC 12, HP 22, Speed 30 ft, with Undead Fortitude."},
|
||||
{"role": "assistant", "content": "Got it! I've noted all the stat blocks. Ready when you want to start the encounter."}]
|
||||
Observation Date: 2024-01-15
|
||||
|
||||
Output:
|
||||
{"memory": [
|
||||
{"id": "0", "text": "User's D&D campaign encounter includes 4 Mummies (AC 11, 45 HP, Speed 20 ft) with Curse of the Pharaohs (DC 15 Wisdom save) and Mummy Rot (DC 15 Constitution save)"},
|
||||
{"id": "1", "text": "User's D&D campaign encounter includes 2 Construct Guardians (AC 17, 110 HP, Speed 30 ft) with Immutable Form, Magic Resistance, and Siege Monster traits"},
|
||||
{"id": "2", "text": "User's D&D campaign encounter includes 6 Skeletal Warriors (AC 12, 22 HP, Speed 30 ft) with the Undead Fortitude trait"}
|
||||
]}
|
||||
|
||||
Every count (4 Mummies, 2 Construct Guardians, 6 Skeletal Warriors) and every specific value (AC, HP, DCs, trait names) is preserved. Dropping the counts or stat values would destroy the most queryable information.
|
||||
|
||||
|
||||
## Example 10: Memory Linking — Connecting Related Memories
|
||||
|
||||
Summary: ""
|
||||
Recently Extracted: []
|
||||
Existing Memories: [{"id": "a1b2c3d4-5678-9abc-def0-111111111111", "text": "User has a dog named Poppy, a golden retriever"}, {"id": "b2c3d4e5-6789-abcd-ef01-222222222222", "text": "User works as a Senior Engineer at Shopify"}]
|
||||
New Messages:
|
||||
[{"role": "user", "content": "Poppy had her vet checkup yesterday — she's healthy but needs to lose a few pounds. Also, I'm switching teams at work next month to the payments platform."}]
|
||||
Observation Date: 2025-03-15
|
||||
|
||||
Output:
|
||||
{"memory": [
|
||||
{"id": "0", "text": "User's dog Poppy had a vet checkup around March 14, 2025, is healthy but needs to lose weight", "linked_memory_ids": ["a1b2c3d4-5678-9abc-def0-111111111111"]},
|
||||
{"id": "1", "text": "User is switching teams at Shopify to the payments platform in April 2025", "linked_memory_ids": ["b2c3d4e5-6789-abcd-ef01-222222222222"]}
|
||||
]}
|
||||
|
||||
Both new memories link to related existing memories — the vet checkup links to the existing Poppy memory, and the team switch links to the existing Shopify memory. This enables the system to build a graph of related memories.
|
||||
|
||||
|
||||
## Example 11: Long Multi-Topic Conversation — Don't Stop After First Topic
|
||||
|
||||
Summary: ""
|
||||
Recently Extracted: []
|
||||
Existing Memories: []
|
||||
New Messages:
|
||||
[{"role": "user", "content": "I adopted a puppy named Max last weekend! He's a beagle mix."},
|
||||
{"role": "assistant", "content": "Congratulations! How's he settling in?"},
|
||||
{"role": "user", "content": "Great! Oh, and I also started pottery classes on Tuesdays. Made a mug with my daughter's face on it."},
|
||||
{"role": "assistant", "content": "Fun! Sounds like a lot going on."},
|
||||
{"role": "user", "content": "Yeah — my sister just moved to Portland too. I'm happy but honestly a bit overwhelmed. My boss gave me a promotion to team lead last week as well."}]
|
||||
Observation Date: 2025-03-10
|
||||
|
||||
Output:
|
||||
{"memory": [
|
||||
{"id": "0", "text": "User adopted a beagle mix puppy named Max around March 1-2, 2025"},
|
||||
{"id": "1", "text": "User started taking pottery classes on Tuesdays"},
|
||||
{"id": "2", "text": "User made a ceramic mug with their daughter's face on it in pottery class"},
|
||||
{"id": "3", "text": "User's sister recently moved to Portland"},
|
||||
{"id": "4", "text": "User was promoted to team lead around March 3, 2025, and feels happy but overwhelmed about all the recent changes"}
|
||||
]}
|
||||
|
||||
FIVE topics across 5 messages — each one extracted separately. Do not stop after the first topic (the puppy). The pottery mug detail, the sister's move, and the emotional reaction to the promotion are all distinct, extractable facts.
|
||||
|
||||
|
||||
## Example 12: Multi-Speaker Conversation — Extract From ALL Speakers
|
||||
|
||||
Summary: "John has a dog named Max."
|
||||
Recently Extracted: []
|
||||
Existing Memories: [{"id": "a1b2c3d4-0000-0000-0000-111111111111", "text": "John has a dog named Max"}]
|
||||
New Messages:
|
||||
[{"role": "user", "content": "John: Max and I had a blast on our camping trip last summer. We hiked, swam, and made great memories. It was a really peaceful experience."},
|
||||
{"role": "assistant", "content": "Maria: That sounds amazing! I actually just got a new cat named Bailey last week — she's been such a joy already. Camping with pets is so soul-nourishing."},
|
||||
{"role": "user", "content": "John: Congrats on Bailey! Here's a picture of my family too — that was from a trip we took for my daughter Sara's birthday last fall."}]
|
||||
Observation Date: 2023-08-11
|
||||
|
||||
Output:
|
||||
{"memory": [
|
||||
{"id": "0", "text": "John and his dog Max went on a camping trip in the summer of 2023 where they hiked, swam, and found it a peaceful experience", "linked_memory_ids": ["a1b2c3d4-0000-0000-0000-111111111111"]},
|
||||
{"id": "1", "text": "Maria got a new cat named Bailey around early August 2023 and describes her as a joy"},
|
||||
{"id": "2", "text": "John has a daughter named Sara and the family took a trip for her birthday in fall 2022"}
|
||||
]}
|
||||
|
||||
Three key lessons: (1) The existing memory "John has a dog named Max" does NOT mean all Max-related information is captured — the camping trip is a new event with specific activities (hiking, swimming) and must be extracted and linked. (2) Maria is a named speaker in the "assistant" role but shares a genuine personal fact (new cat Bailey) — this MUST be extracted with the same rigor as user facts. Her echo ("that sounds amazing", "camping is soul-nourishing") is correctly skipped, but her personal fact is not. (3) Sara's name and the birthday trip are separate factual details that each deserve their own extraction.
|
||||
|
||||
|
||||
# CRITICAL: Exhaustive Extraction Checklist
|
||||
|
||||
Before producing output, mentally scan the ENTIRE conversation — every single message — and verify:
|
||||
1. Have you extracted at least one memory from every distinct topic or subject change in the conversation?
|
||||
2. Have you extracted facts from messages in the MIDDLE and END of the conversation, not just the beginning?
|
||||
3. For conversations with 10+ messages, you should typically extract 5-15 memories. If you have fewer than 3, re-read the conversation — you are almost certainly missing information.
|
||||
4. Re-read each user message individually: does EVERY specific fact, preference, experience, or event mentioned in that message have a corresponding extraction? If a single message mentions two distinct facts (e.g., an allergy AND a hobby), both must be captured.
|
||||
|
||||
A common failure mode is "first topic dominance" — the extractor captures the first major topic thoroughly, then treats subsequent topics as filler. This is WRONG. Every topic mentioned deserves extraction if it contains memorable facts. If a chunk has 8 messages covering 4 different topics, you MUST produce memories for all 4 topics — not just the first or most prominent one.
|
||||
|
||||
|
||||
# OUTPUT FORMAT
|
||||
|
||||
Return ONLY valid JSON parsable by json.loads(). No text, reasoning, explanations, or wrappers.
|
||||
|
||||
## Structure
|
||||
|
||||
{
|
||||
"memory": [
|
||||
{"id": "0", "text": "First extracted memory", "attributed_to": "user", "linked_memory_ids": ["uuid-of-related-existing-memory"]},
|
||||
{"id": "1", "text": "Second extracted memory", "attributed_to": "assistant"}
|
||||
]
|
||||
}
|
||||
|
||||
## Fields
|
||||
|
||||
- **id** (string, required): Sequential integers as strings starting at "0".
|
||||
- **text** (string, required): A contextually rich, self-contained factual statement (15-80 words).
|
||||
- **attributed_to** (string, required): Who this memory is about. Use "user" for facts stated by or about the user (preferences, plans, personal facts). Use "assistant" for information provided by the assistant (recommendations, confirmations, plans created, information researched).
|
||||
- **linked_memory_ids** (array of strings, optional): IDs of Existing Memories that this new memory relates to. Use the exact IDs from the Existing Memories list. Omit or pass [] if no existing memories are related.
|
||||
|
||||
## Rules
|
||||
|
||||
- Extract every piece of memorable information as a separate memory object.
|
||||
- If nothing is worth extracting, return: {"memory": []}
|
||||
- No duplicate IDs. Use double quotes. No trailing commas.
|
||||
|
||||
`;
|
||||
|
||||
export const AGENT_CONTEXT_SUFFIX = `
|
||||
|
||||
## Entity Context
|
||||
|
||||
The primary entity is an AI agent. Frame memories from the agent's perspective:
|
||||
- For user-stated facts, frame as agent knowledge: "Agent was informed that [fact]" or "Agent learned that [fact]"
|
||||
- For agent actions, use direct statements: "Agent recommended [X]" or "Agent specializes in [domain]"
|
||||
- For agent configuration or instructions, capture directly: "Agent is configured to [behavior]"
|
||||
|
||||
The attributed_to field should still reflect the original source: "user" for facts the user stated, "assistant" for things the agent said or did.
|
||||
`;
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// V3 Additive Extraction Schema
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
export const AdditiveExtractionSchema = z.object({
|
||||
memory: z.array(
|
||||
z.object({
|
||||
id: z.string(),
|
||||
text: z.string(),
|
||||
attributed_to: z.enum(["user", "assistant"]).optional(),
|
||||
linked_memory_ids: z.array(z.string()).optional(),
|
||||
}),
|
||||
),
|
||||
});
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// V3 Prompt Builder — generates the user-side prompt for additive extraction
|
||||
// Ported from mem0/configs/prompts.py generate_additive_extraction_prompt()
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
const PAST_MESSAGE_TRUNCATION_LIMIT = 300;
|
||||
|
||||
function truncateContent(
|
||||
text: string,
|
||||
limit = PAST_MESSAGE_TRUNCATION_LIMIT,
|
||||
): string {
|
||||
if (text.length <= limit) return text;
|
||||
return text.slice(0, limit) + "...";
|
||||
}
|
||||
|
||||
function formatConversationHistory(
|
||||
messages?: Array<{ role: string; content: string }>,
|
||||
): string {
|
||||
if (!messages || messages.length === 0) return "";
|
||||
let result = "";
|
||||
for (const msg of messages) {
|
||||
const role = msg.role ?? "";
|
||||
const content = msg.content ?? "";
|
||||
if (role && content) {
|
||||
result += `${role}: ${truncateContent(content)}\n`;
|
||||
}
|
||||
}
|
||||
return result;
|
||||
}
|
||||
|
||||
function serializeMemories(
|
||||
memories?: Array<{ id: string; text: string }>,
|
||||
): string {
|
||||
return JSON.stringify(memories ?? []);
|
||||
}
|
||||
|
||||
export function generateAdditiveExtractionPrompt(options: {
|
||||
existingMemories?: Array<{ id: string; text: string }>;
|
||||
newMessages?: string;
|
||||
lastKMessages?: Array<{ role: string; content: string }>;
|
||||
customInstructions?: string;
|
||||
currentDate?: string;
|
||||
observationDate?: string;
|
||||
}): string {
|
||||
const now = new Date().toISOString().split("T")[0];
|
||||
const currentDate = options.currentDate ?? now;
|
||||
const observationDate = options.observationDate ?? currentDate;
|
||||
|
||||
const sections: string[] = [];
|
||||
|
||||
// Summary — empty for now; callers can extend later
|
||||
sections.push("## Summary\n");
|
||||
|
||||
sections.push(
|
||||
`## Last k Messages\n${formatConversationHistory(options.lastKMessages)}`,
|
||||
);
|
||||
|
||||
// Recently Extracted Memories — empty for now
|
||||
sections.push("## Recently Extracted Memories\n[]");
|
||||
|
||||
sections.push(
|
||||
`## Existing Memories\n${serializeMemories(options.existingMemories)}`,
|
||||
);
|
||||
|
||||
sections.push(`## New Messages\n${options.newMessages ?? "[]"}`);
|
||||
|
||||
sections.push(`## Observation Date\n${observationDate}`);
|
||||
|
||||
sections.push(`## Current Date\n${currentDate}`);
|
||||
|
||||
if (options.customInstructions) {
|
||||
sections.push(`## Custom Instructions\n${options.customInstructions}`);
|
||||
}
|
||||
|
||||
sections.push("# Output:");
|
||||
|
||||
return sections.join("\n\n");
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Legacy helpers (kept for backward compatibility)
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
export function parseMessages(messages: string[]): string {
|
||||
return messages.join("\n");
|
||||
}
|
||||
|
||||
@@ -1,4 +1,5 @@
|
||||
import Database from "better-sqlite3";
|
||||
import { randomUUID } from "crypto";
|
||||
import { HistoryManager } from "./base";
|
||||
import { ensureSQLiteDirectory } from "../utils/sqlite";
|
||||
|
||||
@@ -26,6 +27,16 @@ export class SQLiteManager implements HistoryManager {
|
||||
is_deleted INTEGER DEFAULT 0
|
||||
)
|
||||
`);
|
||||
this.db.exec(`
|
||||
CREATE TABLE IF NOT EXISTS messages (
|
||||
id TEXT PRIMARY KEY,
|
||||
session_scope TEXT,
|
||||
role TEXT,
|
||||
content TEXT,
|
||||
name TEXT,
|
||||
created_at TEXT
|
||||
)
|
||||
`);
|
||||
this.stmtInsert = this.db.prepare(
|
||||
`INSERT INTO memory_history
|
||||
(memory_id, previous_value, new_value, action, created_at, updated_at, is_deleted)
|
||||
@@ -60,8 +71,103 @@ export class SQLiteManager implements HistoryManager {
|
||||
return this.stmtSelect.all(memoryId) as any[];
|
||||
}
|
||||
|
||||
async saveMessages(
|
||||
messages: Array<{ role: string; content: string; name?: string }>,
|
||||
sessionScope: string,
|
||||
): Promise<void> {
|
||||
if (!messages.length) return;
|
||||
|
||||
const insertMsg = this.db.prepare(
|
||||
`INSERT INTO messages (id, session_scope, role, content, name, created_at)
|
||||
VALUES (?, ?, ?, ?, ?, ?)`,
|
||||
);
|
||||
const evict = this.db.prepare(
|
||||
`DELETE FROM messages WHERE session_scope = ? AND id NOT IN (
|
||||
SELECT id FROM (
|
||||
SELECT id FROM messages WHERE session_scope = ? ORDER BY created_at DESC LIMIT 10
|
||||
)
|
||||
)`,
|
||||
);
|
||||
|
||||
const txn = this.db.transaction(() => {
|
||||
const now = new Date().toISOString();
|
||||
for (const msg of messages) {
|
||||
insertMsg.run(
|
||||
randomUUID(),
|
||||
sessionScope,
|
||||
msg.role,
|
||||
msg.content,
|
||||
msg.name ?? null,
|
||||
now,
|
||||
);
|
||||
}
|
||||
evict.run(sessionScope, sessionScope);
|
||||
});
|
||||
|
||||
txn();
|
||||
}
|
||||
|
||||
async getLastMessages(
|
||||
sessionScope: string,
|
||||
limit = 10,
|
||||
): Promise<
|
||||
Array<{ role: string; content: string; name?: string; createdAt: string }>
|
||||
> {
|
||||
const rows = this.db
|
||||
.prepare(
|
||||
`SELECT role, content, name, created_at FROM (
|
||||
SELECT role, content, name, created_at
|
||||
FROM messages
|
||||
WHERE session_scope = ?
|
||||
ORDER BY created_at DESC
|
||||
LIMIT ?
|
||||
) ORDER BY created_at ASC`,
|
||||
)
|
||||
.all(sessionScope, limit) as Array<{
|
||||
role: string;
|
||||
content: string;
|
||||
name: string | null;
|
||||
created_at: string;
|
||||
}>;
|
||||
|
||||
return rows.map((r) => ({
|
||||
role: r.role,
|
||||
content: r.content,
|
||||
...(r.name != null ? { name: r.name } : {}),
|
||||
createdAt: r.created_at,
|
||||
}));
|
||||
}
|
||||
|
||||
async batchAddHistory(
|
||||
records: Array<{
|
||||
memoryId: string;
|
||||
previousValue: string | null;
|
||||
newValue: string | null;
|
||||
action: string;
|
||||
createdAt?: string;
|
||||
updatedAt?: string;
|
||||
isDeleted?: number;
|
||||
}>,
|
||||
): Promise<void> {
|
||||
const txn = this.db.transaction(() => {
|
||||
for (const record of records) {
|
||||
this.stmtInsert.run(
|
||||
record.memoryId,
|
||||
record.previousValue,
|
||||
record.newValue,
|
||||
record.action,
|
||||
record.createdAt ?? null,
|
||||
record.updatedAt ?? null,
|
||||
record.isDeleted ?? 0,
|
||||
);
|
||||
}
|
||||
});
|
||||
txn();
|
||||
}
|
||||
|
||||
async reset(): Promise<void> {
|
||||
this.db.exec("DROP TABLE IF EXISTS memory_history");
|
||||
this.db.exec("DROP TABLE IF EXISTS messages");
|
||||
this.init();
|
||||
}
|
||||
|
||||
|
||||
@@ -11,4 +11,27 @@ export interface HistoryManager {
|
||||
getHistory(memoryId: string): Promise<any[]>;
|
||||
reset(): Promise<void>;
|
||||
close(): void;
|
||||
|
||||
// V3 optional methods — implementations that don't need them can omit these.
|
||||
saveMessages?(
|
||||
messages: Array<{ role: string; content: string; name?: string }>,
|
||||
sessionScope: string,
|
||||
): Promise<void>;
|
||||
getLastMessages?(
|
||||
sessionScope: string,
|
||||
limit?: number,
|
||||
): Promise<
|
||||
Array<{ role: string; content: string; name?: string; createdAt: string }>
|
||||
>;
|
||||
batchAddHistory?(
|
||||
records: Array<{
|
||||
memoryId: string;
|
||||
previousValue: string | null;
|
||||
newValue: string | null;
|
||||
action: string;
|
||||
createdAt?: string;
|
||||
updatedAt?: string;
|
||||
isDeleted?: number;
|
||||
}>,
|
||||
): Promise<void>;
|
||||
}
|
||||
|
||||
@@ -38,7 +38,6 @@ describe("backward compat: ConfigManager.mergeConfig", () => {
|
||||
expect(cfg.historyStore!.provider).toBe("sqlite");
|
||||
expect(cfg.historyStore!.config.historyDbPath).toBe("memory.db");
|
||||
expect(cfg.disableHistory).toBe(false);
|
||||
expect(cfg.graphStore).toBeUndefined();
|
||||
});
|
||||
|
||||
it("workaround: explicit historyStore still works (existing user pattern)", () => {
|
||||
@@ -104,20 +103,6 @@ describe("backward compat: ConfigManager.mergeConfig", () => {
|
||||
expect(cfg.vectorStore.config.dimension).toBe(768);
|
||||
});
|
||||
|
||||
it("graphStore config passes through unchanged", () => {
|
||||
const cfg = ConfigManager.mergeConfig({
|
||||
graphStore: {
|
||||
provider: "neo4j",
|
||||
config: {
|
||||
url: "neo4j://custom:7687",
|
||||
username: "admin",
|
||||
password: "pass",
|
||||
},
|
||||
},
|
||||
});
|
||||
expect(cfg.graphStore!.config.url).toBe("neo4j://custom:7687");
|
||||
});
|
||||
|
||||
it("customInstructions passes through unchanged", () => {
|
||||
const cfg = ConfigManager.mergeConfig({
|
||||
customInstructions: "You are a helpful assistant",
|
||||
|
||||
@@ -50,19 +50,6 @@ export interface LLMConfig {
|
||||
modelProperties?: Record<string, any>;
|
||||
}
|
||||
|
||||
export interface Neo4jConfig {
|
||||
url: string;
|
||||
username: string;
|
||||
password: string;
|
||||
}
|
||||
|
||||
export interface GraphStoreConfig {
|
||||
provider: string;
|
||||
config: Neo4jConfig;
|
||||
llm?: LLMConfig;
|
||||
customInstructions?: string;
|
||||
}
|
||||
|
||||
export interface MemoryConfig {
|
||||
version?: string;
|
||||
embedder: {
|
||||
@@ -81,7 +68,6 @@ export interface MemoryConfig {
|
||||
disableHistory?: boolean;
|
||||
historyDbPath?: string;
|
||||
customInstructions?: string;
|
||||
graphStore?: GraphStoreConfig;
|
||||
}
|
||||
|
||||
export interface MemoryItem {
|
||||
@@ -95,15 +81,14 @@ export interface MemoryItem {
|
||||
}
|
||||
|
||||
export interface SearchFilters {
|
||||
userId?: string;
|
||||
agentId?: string;
|
||||
runId?: string;
|
||||
user_id?: string;
|
||||
agent_id?: string;
|
||||
run_id?: string;
|
||||
[key: string]: any;
|
||||
}
|
||||
|
||||
export interface SearchResult {
|
||||
results: MemoryItem[];
|
||||
relations?: any[];
|
||||
}
|
||||
|
||||
export interface VectorStoreResult {
|
||||
@@ -148,23 +133,6 @@ export const MemoryConfigSchema = z.object({
|
||||
}),
|
||||
historyDbPath: z.string().optional(),
|
||||
customInstructions: z.string().optional(),
|
||||
graphStore: z
|
||||
.object({
|
||||
provider: z.string(),
|
||||
config: z.object({
|
||||
url: z.string(),
|
||||
username: z.string(),
|
||||
password: z.string(),
|
||||
}),
|
||||
llm: z
|
||||
.object({
|
||||
provider: z.string(),
|
||||
config: z.record(z.string(), z.any()),
|
||||
})
|
||||
.optional(),
|
||||
customInstructions: z.string().optional(),
|
||||
})
|
||||
.optional(),
|
||||
historyStore: z
|
||||
.object({
|
||||
provider: z.string(),
|
||||
|
||||
@@ -0,0 +1,720 @@
|
||||
/**
|
||||
* Entity extraction from text using NLP and regex heuristics.
|
||||
*
|
||||
* Extracts four types of entities from text:
|
||||
* - PROPER: Capitalized multi-word sequences (person names, places, brands)
|
||||
* - QUOTED: Text in single or double quotes (titles, specific terms)
|
||||
* - COMPOUND: Multi-word noun phrases with specific modifiers (e.g., "machine learning")
|
||||
* - NOUN: Single nouns from circumstantial compound patterns
|
||||
*
|
||||
* Uses the `compromise` npm package for NLP-based extraction when available.
|
||||
* Falls back to regex-only extraction if `compromise` is not installed.
|
||||
*/
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Filter lists (ported from Python)
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
/** Words that are too generic to be useful as entity heads. */
|
||||
const GENERIC_HEADS: Set<string> = new Set([
|
||||
"thing",
|
||||
"stuff",
|
||||
"way",
|
||||
"time",
|
||||
"experience",
|
||||
"situation",
|
||||
"case",
|
||||
"fact",
|
||||
"matter",
|
||||
"issue",
|
||||
"idea",
|
||||
"thought",
|
||||
"feeling",
|
||||
"place",
|
||||
"area",
|
||||
"part",
|
||||
"kind",
|
||||
"type",
|
||||
"sort",
|
||||
"lot",
|
||||
"bit",
|
||||
"day",
|
||||
"year",
|
||||
"week",
|
||||
"month",
|
||||
"moment",
|
||||
"instance",
|
||||
"example",
|
||||
"technique",
|
||||
"method",
|
||||
"approach",
|
||||
"process",
|
||||
"step",
|
||||
"tool",
|
||||
"result",
|
||||
"outcome",
|
||||
"goal",
|
||||
"task",
|
||||
"item",
|
||||
"topic",
|
||||
"scale",
|
||||
"size",
|
||||
"level",
|
||||
"degree",
|
||||
"amount",
|
||||
"number",
|
||||
"style",
|
||||
"look",
|
||||
"color",
|
||||
"colour",
|
||||
"shape",
|
||||
"form",
|
||||
"piece",
|
||||
"section",
|
||||
"side",
|
||||
"end",
|
||||
"edge",
|
||||
"surface",
|
||||
"point",
|
||||
]);
|
||||
|
||||
/** Adjectives too vague to make a compound entity specific. */
|
||||
const NON_SPECIFIC_ADJ: Set<string> = new Set([
|
||||
"many",
|
||||
"few",
|
||||
"several",
|
||||
"some",
|
||||
"any",
|
||||
"all",
|
||||
"most",
|
||||
"more",
|
||||
"less",
|
||||
"much",
|
||||
"little",
|
||||
"enough",
|
||||
"various",
|
||||
"numerous",
|
||||
"multiple",
|
||||
"countless",
|
||||
"great",
|
||||
"good",
|
||||
"bad",
|
||||
"nice",
|
||||
"terrible",
|
||||
"awful",
|
||||
"awesome",
|
||||
"amazing",
|
||||
"wonderful",
|
||||
"horrible",
|
||||
"excellent",
|
||||
"poor",
|
||||
"best",
|
||||
"worst",
|
||||
"fine",
|
||||
"okay",
|
||||
"new",
|
||||
"old",
|
||||
"recent",
|
||||
"past",
|
||||
"future",
|
||||
"current",
|
||||
"previous",
|
||||
"next",
|
||||
"last",
|
||||
"first",
|
||||
"latest",
|
||||
"early",
|
||||
"late",
|
||||
"former",
|
||||
"modern",
|
||||
"ancient",
|
||||
"big",
|
||||
"small",
|
||||
"large",
|
||||
"tiny",
|
||||
"huge",
|
||||
"enormous",
|
||||
"long",
|
||||
"short",
|
||||
"tall",
|
||||
"high",
|
||||
"low",
|
||||
"wide",
|
||||
"narrow",
|
||||
"thick",
|
||||
"thin",
|
||||
"deep",
|
||||
"shallow",
|
||||
"similar",
|
||||
"different",
|
||||
"same",
|
||||
"other",
|
||||
"another",
|
||||
"such",
|
||||
"certain",
|
||||
"important",
|
||||
"main",
|
||||
"major",
|
||||
"minor",
|
||||
"key",
|
||||
"primary",
|
||||
"real",
|
||||
"actual",
|
||||
"true",
|
||||
"whole",
|
||||
"entire",
|
||||
"full",
|
||||
"complete",
|
||||
"total",
|
||||
"basic",
|
||||
"simple",
|
||||
"interesting",
|
||||
"boring",
|
||||
"exciting",
|
||||
"special",
|
||||
"particular",
|
||||
"general",
|
||||
"common",
|
||||
"unique",
|
||||
"rare",
|
||||
"typical",
|
||||
"usual",
|
||||
"normal",
|
||||
"regular",
|
||||
"possible",
|
||||
"likely",
|
||||
"potential",
|
||||
"available",
|
||||
"necessary",
|
||||
"only",
|
||||
"solo",
|
||||
"individual",
|
||||
"team",
|
||||
"group",
|
||||
"joint",
|
||||
"collaborative",
|
||||
"final",
|
||||
"initial",
|
||||
"side",
|
||||
]);
|
||||
|
||||
/** Generic tail words to strip from compound entities. */
|
||||
const GENERIC_ENDINGS: Set<string> = new Set([
|
||||
"work",
|
||||
"works",
|
||||
"job",
|
||||
"jobs",
|
||||
"task",
|
||||
"tasks",
|
||||
"stuff",
|
||||
"things",
|
||||
"thing",
|
||||
"info",
|
||||
"information",
|
||||
"details",
|
||||
"data",
|
||||
"content",
|
||||
"material",
|
||||
"materials",
|
||||
"activities",
|
||||
"activity",
|
||||
"efforts",
|
||||
"effort",
|
||||
"options",
|
||||
"option",
|
||||
"choices",
|
||||
"choice",
|
||||
"results",
|
||||
"result",
|
||||
"output",
|
||||
"outputs",
|
||||
"products",
|
||||
"product",
|
||||
"items",
|
||||
"item",
|
||||
]);
|
||||
|
||||
/** Capitalized single words that are too generic to be proper nouns. */
|
||||
const GENERIC_CAPS: Set<string> = new Set([
|
||||
"works",
|
||||
"items",
|
||||
"things",
|
||||
"stuff",
|
||||
"resources",
|
||||
"options",
|
||||
"tips",
|
||||
"ideas",
|
||||
"steps",
|
||||
"ways",
|
||||
"methods",
|
||||
"tools",
|
||||
"features",
|
||||
"benefits",
|
||||
"examples",
|
||||
"details",
|
||||
"notes",
|
||||
"instructions",
|
||||
"guidelines",
|
||||
"recommendations",
|
||||
"suggestions",
|
||||
"overview",
|
||||
"summary",
|
||||
"conclusion",
|
||||
"introduction",
|
||||
"pros",
|
||||
"cons",
|
||||
"advantages",
|
||||
"disadvantages",
|
||||
]);
|
||||
|
||||
/** Markdown/formatting markers to skip during extraction. */
|
||||
const FORMATTING_MARKERS: Set<string> = new Set([
|
||||
"*",
|
||||
"-",
|
||||
"+",
|
||||
"\u2022",
|
||||
"\u2013",
|
||||
"\u2014",
|
||||
"#",
|
||||
"##",
|
||||
"###",
|
||||
"**",
|
||||
"__",
|
||||
]);
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Types
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
export interface ExtractedEntity {
|
||||
type: "PROPER" | "QUOTED" | "COMPOUND" | "NOUN";
|
||||
text: string;
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// compromise dynamic import
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
let nlp: any;
|
||||
try {
|
||||
nlp = require("compromise");
|
||||
} catch {
|
||||
// compromise not installed -- use regex-only fallback
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Internal helpers
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
/** Check for formatting artifacts that indicate non-entity text. */
|
||||
function hasArtifacts(txt: string): boolean {
|
||||
if (txt.includes("**") || txt.includes("__") || txt.includes(":*")) {
|
||||
return true;
|
||||
}
|
||||
if (/\s\*\s|\s\*$|^\*\s/.test(txt)) {
|
||||
return true;
|
||||
}
|
||||
if (txt.includes(" ") || txt.includes("\n") || txt.includes("\t")) {
|
||||
return true;
|
||||
}
|
||||
if (txt.length > 100) {
|
||||
return true;
|
||||
}
|
||||
if (/^[\u2022\-+\u2013\u2014]/.test(txt)) {
|
||||
return true;
|
||||
}
|
||||
return false;
|
||||
}
|
||||
|
||||
/** Strip generic trailing words from a phrase's word list. */
|
||||
function stripGenericEnding(words: string[]): string[] {
|
||||
if (words.length <= 1) {
|
||||
return words;
|
||||
}
|
||||
const last = words[words.length - 1].toLowerCase();
|
||||
if (GENERIC_ENDINGS.has(last) && words.length > 2) {
|
||||
return words.slice(0, -1);
|
||||
}
|
||||
return words;
|
||||
}
|
||||
|
||||
/**
|
||||
* Determine if a token position is at the start of a sentence.
|
||||
* Simple heuristic: index 0, or preceded by sentence-ending punctuation
|
||||
* or formatting markers.
|
||||
*/
|
||||
function isSentenceStart(
|
||||
tokens: string[],
|
||||
idx: number,
|
||||
rawText: string,
|
||||
): boolean {
|
||||
if (idx === 0) {
|
||||
return true;
|
||||
}
|
||||
const prev = tokens[idx - 1];
|
||||
if (/[.!?:]$/.test(prev)) {
|
||||
return true;
|
||||
}
|
||||
if (FORMATTING_MARKERS.has(prev)) {
|
||||
return true;
|
||||
}
|
||||
// Check for newline before this token in the raw text
|
||||
const tokenStart = rawText.indexOf(tokens[idx]);
|
||||
if (tokenStart > 0 && rawText.charAt(tokenStart - 1) === "\n") {
|
||||
return true;
|
||||
}
|
||||
return false;
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Extraction strategies
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
/** Extract quoted entities via regex. */
|
||||
function extractQuoted(text: string): ExtractedEntity[] {
|
||||
const entities: ExtractedEntity[] = [];
|
||||
|
||||
// Double-quoted
|
||||
const doubleQuoteRe = /"([^"]+)"/g;
|
||||
let match: RegExpExecArray | null;
|
||||
while ((match = doubleQuoteRe.exec(text)) !== null) {
|
||||
const inner = match[1].trim();
|
||||
if (inner.length > 2) {
|
||||
entities.push({ type: "QUOTED", text: inner });
|
||||
}
|
||||
}
|
||||
|
||||
// Single-quoted (with boundary constraints to avoid apostrophes)
|
||||
const singleQuoteRe = /(?:^|[\s([{,;])'([^']+)'(?=[\s.,;:!?)\]]|$)/g;
|
||||
while ((match = singleQuoteRe.exec(text)) !== null) {
|
||||
const inner = match[1].trim();
|
||||
if (inner.length > 2) {
|
||||
entities.push({ type: "QUOTED", text: inner });
|
||||
}
|
||||
}
|
||||
|
||||
return entities;
|
||||
}
|
||||
|
||||
/**
|
||||
* Extract proper noun sequences using capitalization heuristics.
|
||||
* Finds sequences of capitalized words that are not at sentence starts.
|
||||
*/
|
||||
function extractProper(text: string): ExtractedEntity[] {
|
||||
const entities: ExtractedEntity[] = [];
|
||||
// Tokenize on whitespace, preserving order
|
||||
const tokens = text.split(/\s+/).filter(Boolean);
|
||||
const functionWords = new Set([
|
||||
"'s",
|
||||
"of",
|
||||
"the",
|
||||
"in",
|
||||
"and",
|
||||
"for",
|
||||
"at",
|
||||
"is",
|
||||
]);
|
||||
|
||||
let i = 0;
|
||||
while (i < tokens.length) {
|
||||
const tok = tokens[i];
|
||||
// Skip formatting markers
|
||||
if (FORMATTING_MARKERS.has(tok)) {
|
||||
i++;
|
||||
continue;
|
||||
}
|
||||
|
||||
const isLabel = i + 1 < tokens.length && tokens[i + 1] === ":";
|
||||
const isCap =
|
||||
tok.length > 0 &&
|
||||
tok.charAt(0) === tok.charAt(0).toUpperCase() &&
|
||||
/[A-Z]/.test(tok.charAt(0));
|
||||
|
||||
if (isCap && !isLabel) {
|
||||
const seq: Array<{ token: string; idx: number }> = [
|
||||
{ token: tok, idx: i },
|
||||
];
|
||||
let j = i + 1;
|
||||
while (j < tokens.length) {
|
||||
const t = tokens[j];
|
||||
const tIsCap =
|
||||
t.length > 0 &&
|
||||
t.charAt(0) === t.charAt(0).toUpperCase() &&
|
||||
/[A-Z]/.test(t.charAt(0));
|
||||
if (tIsCap || functionWords.has(t.toLowerCase())) {
|
||||
seq.push({ token: t, idx: j });
|
||||
j++;
|
||||
} else {
|
||||
break;
|
||||
}
|
||||
}
|
||||
|
||||
// Strip trailing function words
|
||||
while (
|
||||
seq.length > 0 &&
|
||||
functionWords.has(seq[seq.length - 1].token.toLowerCase())
|
||||
) {
|
||||
seq.pop();
|
||||
}
|
||||
|
||||
if (seq.length > 0) {
|
||||
// Check for at least one mid-sentence capitalized word
|
||||
const hasMidCap = seq.some(({ token, idx: tokenIdx }) => {
|
||||
const isCapWord =
|
||||
/[A-Z]/.test(token.charAt(0)) &&
|
||||
!functionWords.has(token.toLowerCase());
|
||||
return isCapWord && !isSentenceStart(tokens, tokenIdx, text);
|
||||
});
|
||||
|
||||
if (hasMidCap) {
|
||||
const phrase = seq.map((s) => s.token).join(" ");
|
||||
if (phrase.length > 2) {
|
||||
entities.push({ type: "PROPER", text: phrase });
|
||||
}
|
||||
}
|
||||
}
|
||||
i = j;
|
||||
} else {
|
||||
i++;
|
||||
}
|
||||
}
|
||||
|
||||
return entities;
|
||||
}
|
||||
|
||||
/**
|
||||
* Extract compound noun phrases using the `compromise` NLP library.
|
||||
* Returns COMPOUND and NOUN entities derived from noun chunks.
|
||||
*/
|
||||
function extractCompoundsWithNlp(text: string): ExtractedEntity[] {
|
||||
if (!nlp) {
|
||||
return [];
|
||||
}
|
||||
|
||||
const entities: ExtractedEntity[] = [];
|
||||
const doc = nlp(text);
|
||||
const nouns = doc.nouns().out("array") as string[];
|
||||
|
||||
for (const nounPhrase of nouns) {
|
||||
const trimmed = nounPhrase.trim();
|
||||
if (!trimmed || trimmed.length <= 3) {
|
||||
continue;
|
||||
}
|
||||
|
||||
const words = trimmed.split(/\s+/);
|
||||
if (words.length < 2) {
|
||||
continue;
|
||||
}
|
||||
|
||||
// Filter out phrases where the head is generic
|
||||
const head = words[words.length - 1].toLowerCase();
|
||||
if (GENERIC_HEADS.has(head)) {
|
||||
// Check if there's a specific modifier
|
||||
const hasSpecificMod = words.some(
|
||||
(w) =>
|
||||
!NON_SPECIFIC_ADJ.has(w.toLowerCase()) &&
|
||||
w !== words[words.length - 1],
|
||||
);
|
||||
if (!hasSpecificMod) {
|
||||
continue;
|
||||
}
|
||||
}
|
||||
|
||||
// Filter non-specific adjectives from the beginning
|
||||
const filtered = words.filter(
|
||||
(w) => !NON_SPECIFIC_ADJ.has(w.toLowerCase()),
|
||||
);
|
||||
const cleaned = stripGenericEnding(filtered);
|
||||
|
||||
if (cleaned.length >= 2) {
|
||||
const phrase = cleaned.join(" ");
|
||||
if (phrase.length > 3) {
|
||||
entities.push({ type: "COMPOUND", text: phrase });
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return entities;
|
||||
}
|
||||
|
||||
/**
|
||||
* Regex-only fallback for compound extraction when compromise is not available.
|
||||
* Finds multi-word capitalized sequences and common compound patterns.
|
||||
*/
|
||||
function extractCompoundsRegex(text: string): ExtractedEntity[] {
|
||||
const entities: ExtractedEntity[] = [];
|
||||
|
||||
// Multi-word sequences with at least one non-trivial word
|
||||
// Match sequences like "machine learning", "New York", "data science"
|
||||
const compoundRe =
|
||||
/\b([A-Z][a-z]+(?:\s+(?:of|and|the|for|in)\s+)?[A-Z][a-z]+(?:\s+[A-Z][a-z]+)*)\b/g;
|
||||
let match: RegExpExecArray | null;
|
||||
while ((match = compoundRe.exec(text)) !== null) {
|
||||
const phrase = match[1].trim();
|
||||
if (phrase.length > 3 && phrase.includes(" ")) {
|
||||
const words = phrase.split(/\s+/);
|
||||
const head = words[words.length - 1].toLowerCase();
|
||||
if (!GENERIC_HEADS.has(head)) {
|
||||
const filtered = words.filter(
|
||||
(w) => !NON_SPECIFIC_ADJ.has(w.toLowerCase()),
|
||||
);
|
||||
const cleaned = stripGenericEnding(filtered);
|
||||
if (cleaned.length >= 2) {
|
||||
entities.push({ type: "COMPOUND", text: cleaned.join(" ") });
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Also try lowercase compound patterns (e.g., "machine learning", "deep learning")
|
||||
const lowerCompoundRe = /\b([a-z]+(?:\s+[a-z]+){1,3})\b/g;
|
||||
while ((match = lowerCompoundRe.exec(text)) !== null) {
|
||||
const phrase = match[1].trim();
|
||||
const words = phrase.split(/\s+/);
|
||||
if (words.length >= 2 && words.length <= 4 && phrase.length > 5) {
|
||||
const head = words[words.length - 1].toLowerCase();
|
||||
const allGeneric = words.every(
|
||||
(w) =>
|
||||
NON_SPECIFIC_ADJ.has(w.toLowerCase()) ||
|
||||
GENERIC_HEADS.has(w.toLowerCase()),
|
||||
);
|
||||
if (!allGeneric && !GENERIC_HEADS.has(head)) {
|
||||
// Only include if it looks like a meaningful compound
|
||||
const hasContentWord = words.some(
|
||||
(w) =>
|
||||
!NON_SPECIFIC_ADJ.has(w.toLowerCase()) &&
|
||||
!GENERIC_HEADS.has(w.toLowerCase()) &&
|
||||
w.length > 2,
|
||||
);
|
||||
if (hasContentWord) {
|
||||
const filtered = words.filter(
|
||||
(w) => !NON_SPECIFIC_ADJ.has(w.toLowerCase()),
|
||||
);
|
||||
const cleaned = stripGenericEnding(filtered);
|
||||
if (cleaned.length >= 2) {
|
||||
entities.push({ type: "COMPOUND", text: cleaned.join(" ") });
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return entities;
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Public API
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
/**
|
||||
* Extract named entities, quoted text, and noun compounds from text.
|
||||
*
|
||||
* Uses `compromise` for NLP-based noun phrase extraction when available,
|
||||
* falling back to regex-only heuristics otherwise.
|
||||
*
|
||||
* Entity types (in priority order for deduplication):
|
||||
* PROPER - Capitalized multi-word sequences not at sentence start
|
||||
* COMPOUND - Multi-word noun phrases with specific modifiers
|
||||
* QUOTED - Text in single or double quotes (min 3 chars)
|
||||
* NOUN - Single nouns from circumstantial patterns
|
||||
*
|
||||
* @param text - Input text to extract entities from.
|
||||
* @returns Deduplicated list of extracted entities.
|
||||
*/
|
||||
export function extractEntities(text: string): ExtractedEntity[] {
|
||||
const raw: ExtractedEntity[] = [];
|
||||
|
||||
// 1. QUOTED entities (always regex)
|
||||
raw.push(...extractQuoted(text));
|
||||
|
||||
// 2. PROPER entities (capitalization heuristics)
|
||||
raw.push(...extractProper(text));
|
||||
|
||||
// 3. COMPOUND entities (NLP or regex fallback)
|
||||
if (nlp) {
|
||||
raw.push(...extractCompoundsWithNlp(text));
|
||||
} else {
|
||||
raw.push(...extractCompoundsRegex(text));
|
||||
}
|
||||
|
||||
// === DEDUPLICATION & CLEANUP ===
|
||||
|
||||
// First pass: deduplicate by lowercase text
|
||||
const seen = new Set<string>();
|
||||
const deduped: ExtractedEntity[] = [];
|
||||
for (const entity of raw) {
|
||||
const key = entity.text.toLowerCase().trim();
|
||||
if (key.length > 2 && !seen.has(key)) {
|
||||
seen.add(key);
|
||||
deduped.push(entity);
|
||||
}
|
||||
}
|
||||
|
||||
// Clean up formatting artifacts
|
||||
const cleaned: ExtractedEntity[] = [];
|
||||
for (const entity of deduped) {
|
||||
let txt = entity.text.trim();
|
||||
// Strip leading/trailing asterisks
|
||||
txt = txt.replace(/^\*+\s*|\s*\*+$/g, "");
|
||||
// Strip trailing colons
|
||||
txt = txt.replace(/\s*:+$/, "");
|
||||
// Strip leading numbered list markers
|
||||
txt = txt.replace(/^\d+\s*\.\s*/, "");
|
||||
|
||||
if (!txt || txt.length <= 2 || hasArtifacts(txt)) {
|
||||
continue;
|
||||
}
|
||||
|
||||
// Filter generic single-word PROPER nouns
|
||||
if (
|
||||
entity.type === "PROPER" &&
|
||||
!txt.includes(" ") &&
|
||||
GENERIC_CAPS.has(txt.toLowerCase())
|
||||
) {
|
||||
continue;
|
||||
}
|
||||
|
||||
cleaned.push({ type: entity.type, text: txt });
|
||||
}
|
||||
|
||||
// Keep best type per entity (PROPER > COMPOUND > QUOTED > NOUN)
|
||||
const typePriority: Record<string, number> = {
|
||||
PROPER: 0,
|
||||
COMPOUND: 1,
|
||||
QUOTED: 2,
|
||||
NOUN: 3,
|
||||
};
|
||||
const best = new Map<string, ExtractedEntity>();
|
||||
for (const entity of cleaned) {
|
||||
const key = entity.text.toLowerCase();
|
||||
const existing = best.get(key);
|
||||
if (
|
||||
!existing ||
|
||||
(typePriority[entity.type] ?? 99) < (typePriority[existing.type] ?? 99)
|
||||
) {
|
||||
best.set(key, entity);
|
||||
}
|
||||
}
|
||||
const bestEntities = Array.from(best.values());
|
||||
|
||||
// Remove entities that are substrings of longer entities
|
||||
const allLower = bestEntities.map((e) => e.text.toLowerCase());
|
||||
return bestEntities.filter(
|
||||
(entity) =>
|
||||
!allLower.some(
|
||||
(other) =>
|
||||
entity.text.toLowerCase() !== other &&
|
||||
other.includes(entity.text.toLowerCase()),
|
||||
),
|
||||
);
|
||||
}
|
||||
|
||||
/**
|
||||
* Extract entities from multiple texts.
|
||||
*
|
||||
* @param texts - List of input texts to extract entities from.
|
||||
* @returns List of entity lists, one per input text.
|
||||
*/
|
||||
export function extractEntitiesBatch(texts: string[]): ExtractedEntity[][] {
|
||||
return texts.map(extractEntities);
|
||||
}
|
||||
@@ -21,6 +21,7 @@ import { VectorizeDB } from "../vector_stores/vectorize";
|
||||
import { RedisDB } from "../vector_stores/redis";
|
||||
import { OllamaLLM } from "../llms/ollama";
|
||||
import { LMStudioLLM } from "../llms/lmstudio";
|
||||
import { DeepSeekLLM } from "../llms/deepseek";
|
||||
import { SupabaseDB } from "../vector_stores/supabase";
|
||||
import { SQLiteManager } from "../storage/SQLiteManager";
|
||||
import { MemoryHistoryManager } from "../storage/MemoryHistoryManager";
|
||||
@@ -35,7 +36,6 @@ import { LangchainEmbedder } from "../embeddings/langchain";
|
||||
import { LangchainVectorStore } from "../vector_stores/langchain";
|
||||
import { AzureAISearch } from "../vector_stores/azure_ai_search";
|
||||
import { PGVector } from "../vector_stores/pgvector";
|
||||
import { DeepSeekLLM } from "../llms/deepseek";
|
||||
|
||||
export class EmbedderFactory {
|
||||
static create(provider: string, config: EmbeddingConfig): Embedder {
|
||||
|
||||
@@ -0,0 +1,278 @@
|
||||
/**
|
||||
* BM25 lemmatization for consistent keyword matching.
|
||||
*
|
||||
* Uses the `natural` npm package for Porter stemming when available.
|
||||
* Falls back to simple lowercasing + stop word removal if `natural`
|
||||
* is not installed.
|
||||
*
|
||||
* Also includes original -ing forms alongside stems to handle cases
|
||||
* where stemming produces inconsistent results (e.g., "meeting" as
|
||||
* noun vs verb -> different stems).
|
||||
*/
|
||||
|
||||
/** Standard English stop words (based on NLTK stop word list). */
|
||||
const STOP_WORDS: Set<string> = new Set([
|
||||
"a",
|
||||
"about",
|
||||
"above",
|
||||
"after",
|
||||
"again",
|
||||
"against",
|
||||
"all",
|
||||
"am",
|
||||
"an",
|
||||
"and",
|
||||
"any",
|
||||
"are",
|
||||
"aren't",
|
||||
"as",
|
||||
"at",
|
||||
"be",
|
||||
"because",
|
||||
"been",
|
||||
"before",
|
||||
"being",
|
||||
"below",
|
||||
"between",
|
||||
"both",
|
||||
"but",
|
||||
"by",
|
||||
"can",
|
||||
"can't",
|
||||
"cannot",
|
||||
"could",
|
||||
"couldn't",
|
||||
"did",
|
||||
"didn't",
|
||||
"do",
|
||||
"does",
|
||||
"doesn't",
|
||||
"doing",
|
||||
"don't",
|
||||
"down",
|
||||
"during",
|
||||
"each",
|
||||
"few",
|
||||
"for",
|
||||
"from",
|
||||
"further",
|
||||
"get",
|
||||
"got",
|
||||
"had",
|
||||
"hadn't",
|
||||
"has",
|
||||
"hasn't",
|
||||
"have",
|
||||
"haven't",
|
||||
"having",
|
||||
"he",
|
||||
"her",
|
||||
"here",
|
||||
"hers",
|
||||
"herself",
|
||||
"him",
|
||||
"himself",
|
||||
"his",
|
||||
"how",
|
||||
"i",
|
||||
"if",
|
||||
"in",
|
||||
"into",
|
||||
"is",
|
||||
"isn't",
|
||||
"it",
|
||||
"it's",
|
||||
"its",
|
||||
"itself",
|
||||
"just",
|
||||
"let's",
|
||||
"me",
|
||||
"might",
|
||||
"more",
|
||||
"most",
|
||||
"mustn't",
|
||||
"must",
|
||||
"my",
|
||||
"myself",
|
||||
"no",
|
||||
"nor",
|
||||
"not",
|
||||
"of",
|
||||
"off",
|
||||
"on",
|
||||
"once",
|
||||
"only",
|
||||
"or",
|
||||
"other",
|
||||
"ought",
|
||||
"our",
|
||||
"ours",
|
||||
"ourselves",
|
||||
"out",
|
||||
"over",
|
||||
"own",
|
||||
"same",
|
||||
"shall",
|
||||
"shan't",
|
||||
"she",
|
||||
"should",
|
||||
"shouldn't",
|
||||
"so",
|
||||
"some",
|
||||
"such",
|
||||
"than",
|
||||
"that",
|
||||
"the",
|
||||
"their",
|
||||
"theirs",
|
||||
"them",
|
||||
"themselves",
|
||||
"then",
|
||||
"there",
|
||||
"these",
|
||||
"they",
|
||||
"this",
|
||||
"those",
|
||||
"through",
|
||||
"to",
|
||||
"too",
|
||||
"under",
|
||||
"until",
|
||||
"up",
|
||||
"very",
|
||||
"was",
|
||||
"wasn't",
|
||||
"we",
|
||||
"were",
|
||||
"weren't",
|
||||
"what",
|
||||
"when",
|
||||
"where",
|
||||
"which",
|
||||
"while",
|
||||
"who",
|
||||
"whom",
|
||||
"why",
|
||||
"will",
|
||||
"with",
|
||||
"won't",
|
||||
"would",
|
||||
"wouldn't",
|
||||
"you",
|
||||
"your",
|
||||
"yours",
|
||||
"yourself",
|
||||
"yourselves",
|
||||
]);
|
||||
|
||||
/**
|
||||
* Attempt to load the Porter stemmer from the `natural` package.
|
||||
* Returns null if the package is not installed.
|
||||
*/
|
||||
let _porterStemmer: { stem: (word: string) => string } | null | undefined;
|
||||
|
||||
function getPorterStemmer(): { stem: (word: string) => string } | null {
|
||||
if (_porterStemmer !== undefined) {
|
||||
return _porterStemmer;
|
||||
}
|
||||
try {
|
||||
// eslint-disable-next-line @typescript-eslint/no-var-requires
|
||||
const natural = require("natural");
|
||||
_porterStemmer = natural.PorterStemmer;
|
||||
return _porterStemmer!;
|
||||
} catch {
|
||||
_porterStemmer = null;
|
||||
return null;
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* Simple built-in Porter-like stemmer for common English suffixes.
|
||||
* Used only when the `natural` package is not available.
|
||||
*/
|
||||
function simpleStem(word: string): string {
|
||||
if (word.length <= 3) {
|
||||
return word;
|
||||
}
|
||||
|
||||
// Step-like suffix stripping (simplified Porter rules)
|
||||
let w = word;
|
||||
|
||||
if (w.endsWith("ies") && w.length > 4) {
|
||||
w = w.slice(0, -3) + "i";
|
||||
} else if (w.endsWith("sses")) {
|
||||
w = w.slice(0, -2);
|
||||
} else if (w.endsWith("ness")) {
|
||||
w = w.slice(0, -4);
|
||||
} else if (w.endsWith("ment") && w.length > 5) {
|
||||
w = w.slice(0, -4);
|
||||
} else if (w.endsWith("ation") && w.length > 6) {
|
||||
w = w.slice(0, -5) + "e";
|
||||
} else if (w.endsWith("ting") && w.length > 5) {
|
||||
w = w.slice(0, -3);
|
||||
} else if (w.endsWith("ing") && w.length > 5) {
|
||||
w = w.slice(0, -3);
|
||||
} else if (w.endsWith("ed") && w.length > 4) {
|
||||
w = w.slice(0, -2);
|
||||
} else if (w.endsWith("ly") && w.length > 4) {
|
||||
w = w.slice(0, -2);
|
||||
} else if (w.endsWith("er") && w.length > 4) {
|
||||
w = w.slice(0, -2);
|
||||
} else if (w.endsWith("est") && w.length > 4) {
|
||||
w = w.slice(0, -3);
|
||||
} else if (w.endsWith("s") && !w.endsWith("ss") && w.length > 3) {
|
||||
w = w.slice(0, -1);
|
||||
}
|
||||
|
||||
return w;
|
||||
}
|
||||
|
||||
/**
|
||||
* Lemmatize (stem) text for BM25 matching.
|
||||
*
|
||||
* Processing steps:
|
||||
* 1. Lowercase the text.
|
||||
* 2. Tokenize into words (alphanumeric sequences).
|
||||
* 3. Remove stop words.
|
||||
* 4. Apply Porter stemming to each word.
|
||||
* 5. For words ending in -ing, keep both the stemmed and original form.
|
||||
* 6. Return space-joined result.
|
||||
*
|
||||
* Falls back to simple suffix stripping if `natural` is not installed.
|
||||
*
|
||||
* @param text - Input text to lemmatize.
|
||||
* @returns Space-joined lemmatized/stemmed tokens.
|
||||
*/
|
||||
export function lemmatizeForBm25(text: string): string {
|
||||
const lower = text.toLowerCase();
|
||||
const words = lower.match(/[a-z0-9]+/g);
|
||||
if (!words) {
|
||||
return text.toLowerCase();
|
||||
}
|
||||
|
||||
const stemmer = getPorterStemmer();
|
||||
const stemFn = stemmer
|
||||
? (w: string) => stemmer.stem(w).toLowerCase()
|
||||
: simpleStem;
|
||||
|
||||
const tokens: string[] = [];
|
||||
|
||||
for (const word of words) {
|
||||
if (STOP_WORDS.has(word)) {
|
||||
continue;
|
||||
}
|
||||
|
||||
const stemmed = stemFn(word);
|
||||
if (stemmed && /^[a-z0-9]+$/.test(stemmed)) {
|
||||
tokens.push(stemmed);
|
||||
}
|
||||
|
||||
// Also add original if it ends in -ing and differs from stem.
|
||||
// This handles noun/verb ambiguity (meeting/meet, attending/attend).
|
||||
if (word.endsWith("ing") && word !== stemmed && /^[a-z0-9]+$/.test(word)) {
|
||||
tokens.push(word);
|
||||
}
|
||||
}
|
||||
|
||||
return tokens.join(" ");
|
||||
}
|
||||
@@ -0,0 +1,148 @@
|
||||
/**
|
||||
* Scoring utilities for hybrid retrieval.
|
||||
*
|
||||
* Provides:
|
||||
* - BM25 normalization: Sigmoid normalization of raw BM25 scores to [0, 1].
|
||||
* - BM25 parameter selection: Query-length-adaptive sigmoid parameters.
|
||||
* - Additive scoring: Combined scoring with semantic + BM25 + entity boost.
|
||||
*/
|
||||
|
||||
export const ENTITY_BOOST_WEIGHT = 0.5;
|
||||
|
||||
/**
|
||||
* Get BM25 sigmoid parameters based on query length.
|
||||
*
|
||||
* Longer queries tend to have higher raw BM25 scores, so we adjust
|
||||
* the sigmoid midpoint and steepness accordingly.
|
||||
*
|
||||
* @param query - The original query string.
|
||||
* @param lemmatized - Optional pre-lemmatized query string. If not provided,
|
||||
* the term count is estimated from the raw query.
|
||||
* @returns A tuple of [midpoint, steepness] for sigmoid normalization.
|
||||
*/
|
||||
export function getBm25Params(
|
||||
query: string,
|
||||
lemmatized?: string,
|
||||
): [number, number] {
|
||||
const text = lemmatized ?? query;
|
||||
const numTerms = text.trim().split(/\s+/).filter(Boolean).length || 1;
|
||||
|
||||
if (numTerms <= 3) {
|
||||
return [5.0, 0.7];
|
||||
} else if (numTerms <= 6) {
|
||||
return [7.0, 0.6];
|
||||
} else if (numTerms <= 9) {
|
||||
return [9.0, 0.5];
|
||||
} else if (numTerms <= 15) {
|
||||
return [10.0, 0.5];
|
||||
} else {
|
||||
return [12.0, 0.5];
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* Normalize a raw BM25 score to [0, 1] using logistic sigmoid.
|
||||
*
|
||||
* @param rawScore - Raw BM25 score (unbounded, typically 0-20+).
|
||||
* @param midpoint - Score at which sigmoid outputs 0.5.
|
||||
* @param steepness - Controls how quickly sigmoid transitions.
|
||||
* @returns Normalized score in range [0, 1].
|
||||
*/
|
||||
export function normalizeBm25(
|
||||
rawScore: number,
|
||||
midpoint: number,
|
||||
steepness: number,
|
||||
): number {
|
||||
return 1.0 / (1.0 + Math.exp(-steepness * (rawScore - midpoint)));
|
||||
}
|
||||
|
||||
export interface ScoredResult {
|
||||
id: string;
|
||||
score: number;
|
||||
scoreBreakdown: {
|
||||
semantic: number;
|
||||
bm25: number;
|
||||
entityBoost: number;
|
||||
};
|
||||
payload: Record<string, any>;
|
||||
}
|
||||
|
||||
/**
|
||||
* Score candidates additively and return top-k results.
|
||||
*
|
||||
* For each candidate:
|
||||
* combined = (semantic + bm25 + entity_boost) / max_possible
|
||||
*
|
||||
* Threshold gates the semantic score BEFORE combining -- candidates
|
||||
* below the threshold are excluded even if BM25/entity would boost them.
|
||||
*
|
||||
* The divisor adapts based on which signals are active:
|
||||
* - Semantic only: max_possible = 1.0
|
||||
* - Semantic + BM25: max_possible = 2.0
|
||||
* - Semantic + BM25 + entity: max_possible = 2.5
|
||||
* - Semantic + entity (no BM25): max_possible = 1.5
|
||||
*
|
||||
* @param semanticResults - Candidate results with id, score, and payload.
|
||||
* @param bm25Scores - Map of memory ID to normalized BM25 score.
|
||||
* @param entityBoosts - Map of memory ID to entity boost score.
|
||||
* @param threshold - Minimum semantic score to include a candidate.
|
||||
* @param topK - Maximum number of results to return.
|
||||
* @returns Sorted list of scored results, highest score first.
|
||||
*/
|
||||
export function scoreAndRank(
|
||||
semanticResults: Array<{
|
||||
id: string;
|
||||
score: number;
|
||||
payload: Record<string, any>;
|
||||
}>,
|
||||
bm25Scores: Record<string, number>,
|
||||
entityBoosts: Record<string, number>,
|
||||
threshold: number,
|
||||
topK: number,
|
||||
): ScoredResult[] {
|
||||
const hasBm25 = Object.keys(bm25Scores).length > 0;
|
||||
const hasEntity = Object.keys(entityBoosts).length > 0;
|
||||
|
||||
let maxPossible = 1.0;
|
||||
if (hasBm25) {
|
||||
maxPossible += 1.0;
|
||||
}
|
||||
if (hasEntity) {
|
||||
maxPossible += ENTITY_BOOST_WEIGHT;
|
||||
}
|
||||
|
||||
const scored: ScoredResult[] = [];
|
||||
|
||||
for (const result of semanticResults) {
|
||||
const memId = result.id;
|
||||
if (memId == null) {
|
||||
continue;
|
||||
}
|
||||
|
||||
const semanticScore = result.score ?? 0.0;
|
||||
if (semanticScore < threshold) {
|
||||
continue;
|
||||
}
|
||||
|
||||
const memIdStr = String(memId);
|
||||
const bm25Score = bm25Scores[memIdStr] ?? 0.0;
|
||||
const entityBoost = entityBoosts[memIdStr] ?? 0.0;
|
||||
|
||||
const rawCombined = semanticScore + bm25Score + entityBoost;
|
||||
const combined = Math.min(rawCombined / maxPossible, 1.0);
|
||||
|
||||
scored.push({
|
||||
id: memIdStr,
|
||||
score: combined,
|
||||
scoreBreakdown: {
|
||||
semantic: semanticScore,
|
||||
bm25: bm25Score,
|
||||
entityBoost: entityBoost,
|
||||
},
|
||||
payload: result.payload,
|
||||
});
|
||||
}
|
||||
|
||||
scored.sort((a, b) => b.score - a.score);
|
||||
return scored.slice(0, topK);
|
||||
}
|
||||
@@ -326,6 +326,45 @@ export class AzureAISearch implements VectorStore {
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* Keyword search using Azure AI Search native full-text (BM25) capabilities
|
||||
*/
|
||||
async keywordSearch(
|
||||
query: string,
|
||||
topK: number = 5,
|
||||
filters?: SearchFilters,
|
||||
): Promise<VectorStoreResult[] | null> {
|
||||
try {
|
||||
const filterExpression = filters
|
||||
? this.buildFilterExpression(filters)
|
||||
: undefined;
|
||||
|
||||
const searchResults = await this.searchClient.search(query, {
|
||||
filter: filterExpression,
|
||||
top: topK,
|
||||
searchFields: ["payload"],
|
||||
});
|
||||
|
||||
const results: VectorStoreResult[] = [];
|
||||
|
||||
for await (const result of searchResults.results) {
|
||||
const payloadStr = result.document.payload as string;
|
||||
const payload = JSON.parse(this.extractJson(payloadStr));
|
||||
|
||||
results.push({
|
||||
id: result.document.id as string,
|
||||
score: result.score,
|
||||
payload,
|
||||
});
|
||||
}
|
||||
|
||||
return results;
|
||||
} catch (error) {
|
||||
console.error("Error during keyword search:", error);
|
||||
return null;
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* Search for similar vectors
|
||||
*/
|
||||
|
||||
@@ -11,6 +11,11 @@ export interface VectorStore {
|
||||
topK?: number,
|
||||
filters?: SearchFilters,
|
||||
): Promise<VectorStoreResult[]>;
|
||||
keywordSearch?(
|
||||
query: string,
|
||||
topK?: number,
|
||||
filters?: SearchFilters,
|
||||
): Promise<VectorStoreResult[] | null>;
|
||||
get(vectorId: string): Promise<VectorStoreResult | null>;
|
||||
update(
|
||||
vectorId: string,
|
||||
|
||||
@@ -99,6 +99,10 @@ export class LangchainVectorStore implements VectorStore {
|
||||
}
|
||||
}
|
||||
|
||||
async keywordSearch(): Promise<null> {
|
||||
return null;
|
||||
}
|
||||
|
||||
async search(
|
||||
query: number[],
|
||||
topK: number = 5,
|
||||
|
||||
@@ -67,11 +67,139 @@ export class MemoryVectorStore implements VectorStore {
|
||||
return dotProduct / (Math.sqrt(normA) * Math.sqrt(normB));
|
||||
}
|
||||
|
||||
/**
|
||||
* Check if a single field condition matches the payload.
|
||||
* Supports comparison operators: eq, ne, gt, gte, lt, lte, in, nin, contains, icontains
|
||||
*/
|
||||
private matchFieldCondition(
|
||||
payload: Record<string, any>,
|
||||
key: string,
|
||||
value: any,
|
||||
): boolean {
|
||||
const payloadValue = payload[key];
|
||||
|
||||
// Handle non-dict values
|
||||
if (typeof value !== "object" || value === null) {
|
||||
// Wildcard: match any value
|
||||
if (value === "*") {
|
||||
return true;
|
||||
}
|
||||
// Simple equality
|
||||
return payloadValue === value;
|
||||
}
|
||||
|
||||
// Handle array shorthand: {"field": ["a", "b"]} treated as "in" operator
|
||||
if (Array.isArray(value)) {
|
||||
return value.includes(payloadValue);
|
||||
}
|
||||
|
||||
// Handle comparison operators
|
||||
if ("eq" in value) {
|
||||
return payloadValue === value.eq;
|
||||
}
|
||||
if ("ne" in value) {
|
||||
return payloadValue !== value.ne;
|
||||
}
|
||||
if ("gt" in value) {
|
||||
return payloadValue > value.gt;
|
||||
}
|
||||
if ("gte" in value) {
|
||||
return payloadValue >= value.gte;
|
||||
}
|
||||
if ("lt" in value) {
|
||||
return payloadValue < value.lt;
|
||||
}
|
||||
if ("lte" in value) {
|
||||
return payloadValue <= value.lte;
|
||||
}
|
||||
if ("in" in value) {
|
||||
return Array.isArray(value.in) && value.in.includes(payloadValue);
|
||||
}
|
||||
if ("nin" in value) {
|
||||
return !Array.isArray(value.nin) || !value.nin.includes(payloadValue);
|
||||
}
|
||||
if ("contains" in value) {
|
||||
return (
|
||||
typeof payloadValue === "string" &&
|
||||
payloadValue.includes(value.contains)
|
||||
);
|
||||
}
|
||||
if ("icontains" in value) {
|
||||
return (
|
||||
typeof payloadValue === "string" &&
|
||||
payloadValue.toLowerCase().includes(value.icontains.toLowerCase())
|
||||
);
|
||||
}
|
||||
|
||||
// Unknown operator - treat as nested object for equality (shouldn't happen normally)
|
||||
return payloadValue === value;
|
||||
}
|
||||
|
||||
/**
|
||||
* Filter a vector by the given filters.
|
||||
* Supports logical operators (AND, OR, NOT) and comparison operators.
|
||||
*/
|
||||
private filterVector(vector: MemoryVector, filters?: SearchFilters): boolean {
|
||||
if (!filters) return true;
|
||||
return Object.entries(filters).every(
|
||||
([key, value]) => vector.payload[key] === value,
|
||||
);
|
||||
if (!filters || Object.keys(filters).length === 0) return true;
|
||||
|
||||
// Normalize $or/$not/$and → OR/NOT/AND
|
||||
const keyMap: Record<string, string> = {
|
||||
$and: "AND",
|
||||
$or: "OR",
|
||||
$not: "NOT",
|
||||
};
|
||||
const normalized: Record<string, any> = {};
|
||||
for (const [key, value] of Object.entries(filters)) {
|
||||
const normKey = keyMap[key] || key;
|
||||
if (!(normKey in normalized)) {
|
||||
normalized[normKey] = value;
|
||||
}
|
||||
}
|
||||
|
||||
for (const [key, value] of Object.entries(normalized)) {
|
||||
// Handle logical operators
|
||||
if (key === "AND") {
|
||||
if (!Array.isArray(value)) {
|
||||
throw new Error(
|
||||
`AND filter value must be a list of filter dicts, got ${typeof value}`,
|
||||
);
|
||||
}
|
||||
// All conditions must match
|
||||
const allMatch = value.every((sub: SearchFilters) =>
|
||||
this.filterVector(vector, sub),
|
||||
);
|
||||
if (!allMatch) return false;
|
||||
} else if (key === "OR") {
|
||||
if (!Array.isArray(value)) {
|
||||
throw new Error(
|
||||
`OR filter value must be a list of filter dicts, got ${typeof value}`,
|
||||
);
|
||||
}
|
||||
// At least one condition must match
|
||||
const anyMatch = value.some((sub: SearchFilters) =>
|
||||
this.filterVector(vector, sub),
|
||||
);
|
||||
if (!anyMatch) return false;
|
||||
} else if (key === "NOT") {
|
||||
if (!Array.isArray(value)) {
|
||||
throw new Error(
|
||||
`NOT filter value must be a list of filter dicts, got ${typeof value}`,
|
||||
);
|
||||
}
|
||||
// None of the conditions should match
|
||||
const noneMatch = value.every(
|
||||
(sub: SearchFilters) => !this.filterVector(vector, sub),
|
||||
);
|
||||
if (!noneMatch) return false;
|
||||
} else {
|
||||
// Regular field condition
|
||||
if (!this.matchFieldCondition(vector.payload, key, value)) {
|
||||
return false;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return true;
|
||||
}
|
||||
|
||||
async insert(
|
||||
@@ -98,6 +226,111 @@ export class MemoryVectorStore implements VectorStore {
|
||||
insertMany(vectors, ids, payloads);
|
||||
}
|
||||
|
||||
private tokenize(text: string): string[] {
|
||||
return text.toLowerCase().split(/\s+/).filter(Boolean);
|
||||
}
|
||||
|
||||
async keywordSearch(
|
||||
query: string,
|
||||
topK: number = 10,
|
||||
filters?: SearchFilters,
|
||||
): Promise<VectorStoreResult[] | null> {
|
||||
try {
|
||||
const rows = this.db.prepare(`SELECT * FROM vectors`).all() as any[];
|
||||
|
||||
// Collect documents that pass the filter
|
||||
const candidates: {
|
||||
id: string;
|
||||
payload: Record<string, any>;
|
||||
tokens: string[];
|
||||
}[] = [];
|
||||
|
||||
for (const row of rows) {
|
||||
const payload = JSON.parse(row.payload);
|
||||
const memoryVector: MemoryVector = {
|
||||
id: row.id,
|
||||
vector: Array.from(
|
||||
new Float32Array(
|
||||
row.vector.buffer,
|
||||
row.vector.byteOffset,
|
||||
row.vector.byteLength / 4,
|
||||
),
|
||||
),
|
||||
payload,
|
||||
};
|
||||
|
||||
if (this.filterVector(memoryVector, filters)) {
|
||||
const text = payload.text_lemmatized || payload.data || "";
|
||||
candidates.push({ id: row.id, payload, tokens: this.tokenize(text) });
|
||||
}
|
||||
}
|
||||
|
||||
if (candidates.length === 0) {
|
||||
return [];
|
||||
}
|
||||
|
||||
const tokenizedQuery = this.tokenize(query);
|
||||
if (tokenizedQuery.length === 0) {
|
||||
return [];
|
||||
}
|
||||
|
||||
// Compute BM25 scores inline
|
||||
const k1 = 1.5;
|
||||
const b = 0.75;
|
||||
const N = candidates.length;
|
||||
const avgDocLength =
|
||||
candidates.reduce((sum, c) => sum + c.tokens.length, 0) / N;
|
||||
|
||||
// Compute document frequency for query terms
|
||||
const docFreq = new Map<string, number>();
|
||||
for (const term of tokenizedQuery) {
|
||||
if (!docFreq.has(term)) {
|
||||
let count = 0;
|
||||
for (const c of candidates) {
|
||||
if (c.tokens.includes(term)) count++;
|
||||
}
|
||||
docFreq.set(term, count);
|
||||
}
|
||||
}
|
||||
|
||||
// Compute IDF for query terms
|
||||
const idf = new Map<string, number>();
|
||||
for (const [term, freq] of docFreq) {
|
||||
idf.set(term, Math.log((N - freq + 0.5) / (freq + 0.5) + 1));
|
||||
}
|
||||
|
||||
// Score each candidate
|
||||
const scored = candidates.map((candidate) => {
|
||||
let score = 0;
|
||||
const docLength = candidate.tokens.length;
|
||||
for (const term of tokenizedQuery) {
|
||||
const tf = candidate.tokens.filter((t) => t === term).length;
|
||||
const termIdf = idf.get(term) || 0;
|
||||
score +=
|
||||
(termIdf * tf * (k1 + 1)) /
|
||||
(tf + k1 * (1 - b + (b * docLength) / avgDocLength));
|
||||
}
|
||||
return { ...candidate, score };
|
||||
});
|
||||
|
||||
// Filter out zero-score documents and sort descending
|
||||
const results = scored
|
||||
.filter((s) => s.score > 0)
|
||||
.sort((a, b) => b.score - a.score)
|
||||
.slice(0, topK)
|
||||
.map((s) => ({
|
||||
id: s.id,
|
||||
payload: s.payload,
|
||||
score: s.score,
|
||||
}));
|
||||
|
||||
return results;
|
||||
} catch (error) {
|
||||
console.error("Error during keyword search:", error);
|
||||
return null;
|
||||
}
|
||||
}
|
||||
|
||||
async search(
|
||||
query: number[],
|
||||
topK: number = 10,
|
||||
|
||||
@@ -162,6 +162,51 @@ export class PGVector implements VectorStore {
|
||||
);
|
||||
}
|
||||
|
||||
async keywordSearch(
|
||||
query: string,
|
||||
topK: number = 5,
|
||||
filters?: SearchFilters,
|
||||
): Promise<VectorStoreResult[] | null> {
|
||||
try {
|
||||
const filterConditions: string[] = [];
|
||||
const filterValues: any[] = [query, topK];
|
||||
let filterIndex = 3;
|
||||
|
||||
if (filters) {
|
||||
for (const [key, value] of Object.entries(filters)) {
|
||||
filterConditions.push(`payload->>'${key}' = $${filterIndex}`);
|
||||
filterValues.push(value);
|
||||
filterIndex++;
|
||||
}
|
||||
}
|
||||
|
||||
const filterClause =
|
||||
filterConditions.length > 0
|
||||
? "AND " + filterConditions.join(" AND ")
|
||||
: "";
|
||||
|
||||
const searchQuery = `
|
||||
SELECT id, ts_rank_cd(to_tsvector('simple', payload->>'text_lemmatized'), plainto_tsquery('simple', $1)) AS score, payload
|
||||
FROM ${this.collectionName}
|
||||
WHERE to_tsvector('simple', payload->>'text_lemmatized') @@ plainto_tsquery('simple', $1)
|
||||
${filterClause}
|
||||
ORDER BY score DESC
|
||||
LIMIT $2
|
||||
`;
|
||||
|
||||
const result = await this.client.query(searchQuery, filterValues);
|
||||
|
||||
return result.rows.map((row) => ({
|
||||
id: row.id,
|
||||
payload: row.payload,
|
||||
score: row.score,
|
||||
}));
|
||||
} catch (error) {
|
||||
console.error("Error during keyword search:", error);
|
||||
return null;
|
||||
}
|
||||
}
|
||||
|
||||
async search(
|
||||
query: number[],
|
||||
topK: number = 5,
|
||||
|
||||
@@ -31,17 +31,29 @@ interface QdrantConfig extends VectorStoreConfig {
|
||||
}
|
||||
|
||||
interface QdrantFilter {
|
||||
must?: QdrantCondition[];
|
||||
must_not?: QdrantCondition[];
|
||||
should?: QdrantCondition[];
|
||||
must?: (QdrantCondition | QdrantFilter)[];
|
||||
must_not?: (QdrantCondition | QdrantFilter)[];
|
||||
should?: (QdrantCondition | QdrantFilter)[];
|
||||
}
|
||||
|
||||
interface QdrantCondition {
|
||||
key: string;
|
||||
match?: { value: any };
|
||||
range?: { gte?: number; gt?: number; lte?: number; lt?: number };
|
||||
match?: { value?: any; any?: any[]; except?: any[]; text?: string };
|
||||
range?: {
|
||||
gte?: number | string;
|
||||
gt?: number | string;
|
||||
lte?: number | string;
|
||||
lt?: number | string;
|
||||
};
|
||||
}
|
||||
|
||||
// Normalize $and/$or/$not to AND/OR/NOT
|
||||
const KEY_MAP: Record<string, string> = {
|
||||
$and: "AND",
|
||||
$or: "OR",
|
||||
$not: "NOT",
|
||||
};
|
||||
|
||||
export class Qdrant implements VectorStore {
|
||||
private client: QdrantClient;
|
||||
private readonly collectionName: string;
|
||||
@@ -90,35 +102,167 @@ export class Qdrant implements VectorStore {
|
||||
this.initialize().catch(console.error);
|
||||
}
|
||||
|
||||
private createFilter(filters?: SearchFilters): QdrantFilter | undefined {
|
||||
if (!filters) return undefined;
|
||||
/**
|
||||
* Build a single field condition from a key-value filter pair.
|
||||
* Supports enhanced filter syntax with comparison operators.
|
||||
*/
|
||||
private buildFieldCondition(key: string, value: any): QdrantCondition | null {
|
||||
// Handle non-dict values
|
||||
if (typeof value !== "object" || value === null) {
|
||||
// Wildcard: match any value - skip this filter
|
||||
if (value === "*") {
|
||||
return null;
|
||||
}
|
||||
// Simple equality
|
||||
return { key, match: { value } };
|
||||
}
|
||||
|
||||
const conditions: QdrantCondition[] = [];
|
||||
// Handle array shorthand: {"field": ["a", "b"]} treated as "in" operator
|
||||
if (Array.isArray(value)) {
|
||||
return { key, match: { any: value } };
|
||||
}
|
||||
|
||||
const ops = Object.keys(value);
|
||||
const rangeOps = ["gt", "gte", "lt", "lte"];
|
||||
const hasRangeOps = ops.some((op) => rangeOps.includes(op));
|
||||
const nonRangeOps = ops.filter((op) => !rangeOps.includes(op));
|
||||
|
||||
// Handle range operators
|
||||
if (hasRangeOps) {
|
||||
if (nonRangeOps.length > 0) {
|
||||
throw new Error(
|
||||
`Cannot mix range operators (${ops.filter((o) => rangeOps.includes(o)).join(", ")}) ` +
|
||||
`with non-range operators (${nonRangeOps.join(", ")}) for field '${key}'. ` +
|
||||
`Use AND to combine them as separate conditions.`,
|
||||
);
|
||||
}
|
||||
const range: Record<string, number | string> = {};
|
||||
for (const op of rangeOps) {
|
||||
if (op in value) {
|
||||
range[op] = value[op];
|
||||
}
|
||||
}
|
||||
return { key, range };
|
||||
}
|
||||
|
||||
// Handle comparison operators
|
||||
if ("eq" in value) {
|
||||
return { key, match: { value: value.eq } };
|
||||
}
|
||||
if ("ne" in value) {
|
||||
return { key, match: { except: [value.ne] } };
|
||||
}
|
||||
if ("in" in value) {
|
||||
return { key, match: { any: value.in } };
|
||||
}
|
||||
if ("nin" in value) {
|
||||
return { key, match: { except: value.nin } };
|
||||
}
|
||||
if ("contains" in value || "icontains" in value) {
|
||||
const text = value.contains || value.icontains;
|
||||
return { key, match: { text } };
|
||||
}
|
||||
|
||||
// Unknown operator - treat as nested object for simple match
|
||||
const supportedOps = [
|
||||
"eq",
|
||||
"ne",
|
||||
"gt",
|
||||
"gte",
|
||||
"lt",
|
||||
"lte",
|
||||
"in",
|
||||
"nin",
|
||||
"contains",
|
||||
"icontains",
|
||||
];
|
||||
throw new Error(
|
||||
`Unsupported filter operator(s) for field '${key}': ${ops.join(", ")}. ` +
|
||||
`Supported operators: ${supportedOps.join(", ")}`,
|
||||
);
|
||||
}
|
||||
|
||||
/**
|
||||
* Create a Filter object from the provided filters.
|
||||
* Supports logical operators (AND, OR, NOT) and comparison operators.
|
||||
*/
|
||||
private createFilter(filters?: SearchFilters): QdrantFilter | undefined {
|
||||
if (!filters || Object.keys(filters).length === 0) return undefined;
|
||||
|
||||
// Normalize $or/$not/$and → OR/NOT/AND and deduplicate
|
||||
const normalized: Record<string, any> = {};
|
||||
for (const [key, value] of Object.entries(filters)) {
|
||||
if (
|
||||
typeof value === "object" &&
|
||||
value !== null &&
|
||||
"gte" in value &&
|
||||
"lte" in value
|
||||
) {
|
||||
conditions.push({
|
||||
key,
|
||||
range: {
|
||||
gte: value.gte,
|
||||
lte: value.lte,
|
||||
},
|
||||
});
|
||||
} else {
|
||||
conditions.push({
|
||||
key,
|
||||
match: {
|
||||
value,
|
||||
},
|
||||
});
|
||||
const normKey = KEY_MAP[key] || key;
|
||||
if (!(normKey in normalized)) {
|
||||
normalized[normKey] = value;
|
||||
}
|
||||
}
|
||||
|
||||
return conditions.length ? { must: conditions } : undefined;
|
||||
const must: (QdrantCondition | QdrantFilter)[] = [];
|
||||
const should: (QdrantCondition | QdrantFilter)[] = [];
|
||||
const mustNot: (QdrantCondition | QdrantFilter)[] = [];
|
||||
|
||||
for (const [key, value] of Object.entries(normalized)) {
|
||||
// Handle logical operators
|
||||
if (key === "AND" || key === "OR" || key === "NOT") {
|
||||
if (!Array.isArray(value)) {
|
||||
throw new Error(
|
||||
`${key} filter value must be a list of filter dicts, got ${typeof value}`,
|
||||
);
|
||||
}
|
||||
for (let i = 0; i < value.length; i++) {
|
||||
const item = value[i];
|
||||
if (
|
||||
typeof item !== "object" ||
|
||||
item === null ||
|
||||
Array.isArray(item)
|
||||
) {
|
||||
throw new Error(
|
||||
`${key} filter list item at index ${i} must be a dict, got ${typeof item}`,
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
if (key === "AND") {
|
||||
for (const sub of value) {
|
||||
const built = this.createFilter(sub);
|
||||
if (built) {
|
||||
must.push(built);
|
||||
}
|
||||
}
|
||||
} else if (key === "OR") {
|
||||
for (const sub of value) {
|
||||
const built = this.createFilter(sub);
|
||||
if (built) {
|
||||
should.push(built);
|
||||
}
|
||||
}
|
||||
} else if (key === "NOT") {
|
||||
for (const sub of value) {
|
||||
const built = this.createFilter(sub);
|
||||
if (built) {
|
||||
mustNot.push(built);
|
||||
}
|
||||
}
|
||||
}
|
||||
} else {
|
||||
// Regular field condition
|
||||
const condition = this.buildFieldCondition(key, value);
|
||||
if (condition !== null) {
|
||||
must.push(condition);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if (must.length === 0 && should.length === 0 && mustNot.length === 0) {
|
||||
return undefined;
|
||||
}
|
||||
|
||||
return {
|
||||
must: must.length > 0 ? must : undefined,
|
||||
should: should.length > 0 ? should : undefined,
|
||||
must_not: mustNot.length > 0 ? mustNot : undefined,
|
||||
};
|
||||
}
|
||||
|
||||
async insert(
|
||||
@@ -137,6 +281,10 @@ export class Qdrant implements VectorStore {
|
||||
});
|
||||
}
|
||||
|
||||
async keywordSearch(): Promise<null> {
|
||||
return null;
|
||||
}
|
||||
|
||||
async search(
|
||||
query: number[],
|
||||
topK: number = 5,
|
||||
|
||||
@@ -357,6 +357,10 @@ export class RedisDB implements VectorStore {
|
||||
}
|
||||
}
|
||||
|
||||
async keywordSearch(): Promise<null> {
|
||||
return null;
|
||||
}
|
||||
|
||||
async search(
|
||||
query: number[],
|
||||
topK: number = 5,
|
||||
|
||||
@@ -229,6 +229,10 @@ See the SQL migration instructions in the code comments.`,
|
||||
}
|
||||
}
|
||||
|
||||
async keywordSearch(): Promise<null> {
|
||||
return null;
|
||||
}
|
||||
|
||||
async search(
|
||||
query: number[],
|
||||
topK: number = 5,
|
||||
|
||||
@@ -74,6 +74,10 @@ export class VectorizeDB implements VectorStore {
|
||||
}
|
||||
}
|
||||
|
||||
async keywordSearch(): Promise<null> {
|
||||
return null;
|
||||
}
|
||||
|
||||
async search(
|
||||
query: number[],
|
||||
topK: number = 5,
|
||||
|
||||
@@ -363,78 +363,6 @@ describe("ConfigManager", () => {
|
||||
});
|
||||
});
|
||||
|
||||
// ─────────────────────────────────────────────────────────────────────────
|
||||
// Graph store LLM config propagation (issue #3425)
|
||||
// ─────────────────────────────────────────────────────────────────────────
|
||||
describe("mergeConfig - graph store LLM config (issue #3425)", () => {
|
||||
const baseEmbedder = {
|
||||
provider: "openai",
|
||||
config: { apiKey: "test-key" },
|
||||
};
|
||||
const baseVectorStore = {
|
||||
provider: "memory",
|
||||
config: { collectionName: "test" },
|
||||
};
|
||||
const graphStoreNeo4j = {
|
||||
provider: "neo4j",
|
||||
config: {
|
||||
url: "neo4j://localhost:7687",
|
||||
username: "neo4j",
|
||||
password: "password",
|
||||
},
|
||||
};
|
||||
|
||||
it("should NOT have a default graphStore.llm — root llm should be the fallback", () => {
|
||||
const config = ConfigManager.mergeConfig({
|
||||
embedder: baseEmbedder,
|
||||
vectorStore: baseVectorStore,
|
||||
llm: {
|
||||
provider: "anthropic",
|
||||
config: { model: "claude-sonnet-4-20250514" },
|
||||
},
|
||||
graphStore: graphStoreNeo4j,
|
||||
});
|
||||
|
||||
// graphStore should NOT have its own llm after merge
|
||||
expect(config.graphStore?.llm).toBeUndefined();
|
||||
// root llm should be anthropic
|
||||
expect(config.llm.provider).toBe("anthropic");
|
||||
expect(config.llm.config.model).toBe("claude-sonnet-4-20250514");
|
||||
});
|
||||
|
||||
it("should preserve explicit graphStore.llm when user provides it", () => {
|
||||
const config = ConfigManager.mergeConfig({
|
||||
embedder: baseEmbedder,
|
||||
vectorStore: baseVectorStore,
|
||||
llm: {
|
||||
provider: "anthropic",
|
||||
config: { model: "claude-sonnet-4-20250514" },
|
||||
},
|
||||
graphStore: {
|
||||
...graphStoreNeo4j,
|
||||
llm: { provider: "openai", config: { model: "gpt-4o" } },
|
||||
},
|
||||
});
|
||||
|
||||
// graphStore should have its own llm
|
||||
expect(config.graphStore?.llm?.provider).toBe("openai");
|
||||
expect(config.graphStore?.llm?.config).toEqual({ model: "gpt-4o" });
|
||||
// root llm should still be anthropic
|
||||
expect(config.llm.provider).toBe("anthropic");
|
||||
});
|
||||
|
||||
it("should not have graphStore.llm when user does not provide one", () => {
|
||||
const config = ConfigManager.mergeConfig({
|
||||
embedder: baseEmbedder,
|
||||
vectorStore: baseVectorStore,
|
||||
llm: { provider: "openai", config: { model: "gpt-4o" } },
|
||||
});
|
||||
|
||||
// Default graphStore should not have llm
|
||||
expect(config.graphStore?.llm).toBeUndefined();
|
||||
});
|
||||
});
|
||||
|
||||
// ─────────────────────────────────────────────────────────────────────────
|
||||
// Memory class – LM Studio end-to-end flow (mocked factories)
|
||||
// ─────────────────────────────────────────────────────────────────────────
|
||||
@@ -520,7 +448,7 @@ describe("Memory – LM Studio end-to-end flow", () => {
|
||||
disableHistory: true,
|
||||
});
|
||||
|
||||
await mem.getAll({ userId: "u1" });
|
||||
await mem.getAll({ filters: { user_id: "u1" } });
|
||||
|
||||
expect(mockEmbedderFactory.create).toHaveBeenCalledWith(
|
||||
"lmstudio",
|
||||
@@ -555,7 +483,7 @@ describe("Memory – LM Studio end-to-end flow", () => {
|
||||
disableHistory: true,
|
||||
});
|
||||
|
||||
await mem.getAll({ userId: "u1" });
|
||||
await mem.getAll({ filters: { user_id: "u1" } });
|
||||
|
||||
expect(mockEmbedder.embed).toHaveBeenCalledWith("dimension probe");
|
||||
const vsCall = mockVectorStoreFactory.create.mock.calls[0];
|
||||
@@ -583,7 +511,7 @@ describe("Memory – LM Studio end-to-end flow", () => {
|
||||
disableHistory: true,
|
||||
});
|
||||
|
||||
await mem.getAll({ userId: "u1" });
|
||||
await mem.getAll({ filters: { user_id: "u1" } });
|
||||
|
||||
expect(mockEmbedderFactory.create).toHaveBeenCalledWith(
|
||||
"lmstudio",
|
||||
@@ -636,7 +564,7 @@ describe("Memory – LM Studio end-to-end flow", () => {
|
||||
});
|
||||
|
||||
const result = await mem.search("What does the user like?", {
|
||||
userId: "u1",
|
||||
filters: { user_id: "u1" },
|
||||
});
|
||||
|
||||
expect(mockEmbedder.embed).toHaveBeenCalledWith("What does the user like?");
|
||||
|
||||
@@ -313,7 +313,7 @@ describe("Memory – auto-initialization", () => {
|
||||
disableHistory: true,
|
||||
});
|
||||
|
||||
await mem.getAll({ userId: "u1" });
|
||||
await mem.getAll({ filters: { user_id: "u1" } });
|
||||
|
||||
// Should have called embed("dimension probe") to detect dimension
|
||||
expect(mockEmbedder.embed).toHaveBeenCalledWith("dimension probe");
|
||||
@@ -339,7 +339,7 @@ describe("Memory – auto-initialization", () => {
|
||||
disableHistory: true,
|
||||
});
|
||||
|
||||
await mem.getAll({ userId: "u1" });
|
||||
await mem.getAll({ filters: { user_id: "u1" } });
|
||||
|
||||
// embed should NOT have been called for probing
|
||||
expect(mockEmbedder.embed).not.toHaveBeenCalledWith("dimension probe");
|
||||
@@ -365,7 +365,7 @@ describe("Memory – auto-initialization", () => {
|
||||
disableHistory: true,
|
||||
});
|
||||
|
||||
await mem.getAll({ userId: "u1" });
|
||||
await mem.getAll({ filters: { user_id: "u1" } });
|
||||
|
||||
// ConfigManager resolves dimension from embeddingDims → no probe needed
|
||||
expect(mockEmbedder.embed).not.toHaveBeenCalledWith("dimension probe");
|
||||
@@ -403,9 +403,11 @@ describe("Memory – auto-initialization", () => {
|
||||
let searchDone = false;
|
||||
let getDone = false;
|
||||
|
||||
const getAllP = mem.getAll({ userId: "u" }).then(() => (getAllDone = true));
|
||||
const getAllP = mem
|
||||
.getAll({ filters: { user_id: "u" } })
|
||||
.then(() => (getAllDone = true));
|
||||
const searchP = mem
|
||||
.search("q", { userId: "u" })
|
||||
.search("q", { filters: { user_id: "u" } })
|
||||
.then(() => (searchDone = true));
|
||||
const getP = mem.get("id").then(() => (getDone = true));
|
||||
|
||||
@@ -435,7 +437,7 @@ describe("Memory – auto-initialization", () => {
|
||||
disableHistory: true,
|
||||
});
|
||||
|
||||
await mem.getAll({ userId: "u1" });
|
||||
await mem.getAll({ filters: { user_id: "u1" } });
|
||||
expect(mockVectorStoreFactory.create).toHaveBeenCalledTimes(1);
|
||||
|
||||
// Reset should re-create vector store
|
||||
@@ -473,7 +475,7 @@ describe("Memory – auto-initialization", () => {
|
||||
disableHistory: true,
|
||||
});
|
||||
|
||||
await mem.getAll({ userId: "u1" });
|
||||
await mem.getAll({ filters: { user_id: "u1" } });
|
||||
expect(mockEmbedder.embed).not.toHaveBeenCalledWith("dimension probe");
|
||||
});
|
||||
|
||||
@@ -497,7 +499,7 @@ describe("Memory – auto-initialization", () => {
|
||||
});
|
||||
|
||||
// getAll should reject with the init error
|
||||
await expect(mem.getAll({ userId: "u1" })).rejects.toThrow(
|
||||
await expect(mem.getAll({ filters: { user_id: "u1" } })).rejects.toThrow(
|
||||
"auto-detect embedding dimension",
|
||||
);
|
||||
|
||||
|
||||
@@ -1,681 +0,0 @@
|
||||
/**
|
||||
* Regression tests for graph_memory.ts response parsing (issue #4248).
|
||||
*
|
||||
* Exercises the three json_object call sites in MemoryGraph with a mocked LLM:
|
||||
* 1. _retrieveNodesFromData → entity extraction
|
||||
* 2. _establishNodesRelationsFromData → relation extraction
|
||||
* 3. _getDeleteEntitiesFromSearchOutput → deletion identification
|
||||
*
|
||||
* Covers: malformed LLM responses, missing fields, bad JSON in toolCalls,
|
||||
* string-only responses, empty tool calls, and prompt construction.
|
||||
*
|
||||
* See: https://github.com/mem0ai/mem0/issues/4248
|
||||
*/
|
||||
|
||||
import { MemoryGraph } from "../src/memory/graph_memory";
|
||||
import {
|
||||
EXTRACT_RELATIONS_PROMPT,
|
||||
getDeleteMessages,
|
||||
} from "../src/graphs/utils";
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Mocks – we replace heavy dependencies so tests run without Neo4j / OpenAI
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
// Mock neo4j-driver: provides a fake Driver with a no-op session
|
||||
jest.mock("neo4j-driver", () => ({
|
||||
__esModule: true,
|
||||
default: {
|
||||
driver: jest.fn(() => ({
|
||||
session: () => ({
|
||||
run: jest.fn().mockResolvedValue({ records: [] }),
|
||||
close: jest.fn(),
|
||||
}),
|
||||
})),
|
||||
auth: { basic: jest.fn() },
|
||||
},
|
||||
}));
|
||||
|
||||
// Mock factory so constructor doesn't try to instantiate real LLMs / embedders
|
||||
const mockGenerateResponse = jest.fn();
|
||||
const mockGenerateChat = jest.fn();
|
||||
const mockEmbed = jest.fn().mockResolvedValue([0.1, 0.2, 0.3]);
|
||||
|
||||
jest.mock("../src/utils/factory", () => ({
|
||||
LLMFactory: {
|
||||
create: jest.fn(() => ({
|
||||
generateResponse: mockGenerateResponse,
|
||||
generateChat: mockGenerateChat,
|
||||
})),
|
||||
},
|
||||
EmbedderFactory: {
|
||||
create: jest.fn(() => ({
|
||||
embed: mockEmbed,
|
||||
})),
|
||||
},
|
||||
}));
|
||||
|
||||
// Minimal config that satisfies the MemoryGraph constructor
|
||||
function makeConfig(overrides: Record<string, any> = {}) {
|
||||
return {
|
||||
graphStore: {
|
||||
config: {
|
||||
url: "bolt://localhost:7687",
|
||||
username: "neo4j",
|
||||
password: "test",
|
||||
},
|
||||
...overrides,
|
||||
},
|
||||
embedder: { provider: "openai", config: {} },
|
||||
llm: { provider: "openai", config: {} },
|
||||
} as any;
|
||||
}
|
||||
|
||||
// Helper to access private methods via `any` cast
|
||||
function graph(overrides: Record<string, any> = {}): any {
|
||||
return new MemoryGraph(makeConfig(overrides));
|
||||
}
|
||||
|
||||
const FILTERS = { userId: "test-user" };
|
||||
|
||||
beforeEach(() => {
|
||||
jest.clearAllMocks();
|
||||
});
|
||||
|
||||
// ═══════════════════════════════════════════════════════════════════════════
|
||||
// 1. _retrieveNodesFromData – entity extraction
|
||||
// ═══════════════════════════════════════════════════════════════════════════
|
||||
|
||||
describe("_retrieveNodesFromData", () => {
|
||||
it("parses a well-formed extract_entities tool call", async () => {
|
||||
mockGenerateResponse.mockResolvedValueOnce({
|
||||
toolCalls: [
|
||||
{
|
||||
name: "extract_entities",
|
||||
arguments: JSON.stringify({
|
||||
entities: [
|
||||
{ entity: "Alice", entity_type: "person" },
|
||||
{ entity: "Pizza", entity_type: "food" },
|
||||
],
|
||||
}),
|
||||
},
|
||||
],
|
||||
});
|
||||
|
||||
const mg = graph();
|
||||
const result = await mg._retrieveNodesFromData(
|
||||
"Alice likes pizza",
|
||||
FILTERS,
|
||||
);
|
||||
|
||||
expect(result).toEqual({ alice: "person", pizza: "food" });
|
||||
});
|
||||
|
||||
it("returns empty map when LLM returns a plain string", async () => {
|
||||
mockGenerateResponse.mockResolvedValueOnce("I am a string, not an object");
|
||||
|
||||
const mg = graph();
|
||||
const result = await mg._retrieveNodesFromData("anything", FILTERS);
|
||||
|
||||
expect(result).toEqual({});
|
||||
});
|
||||
|
||||
it("returns empty map when toolCalls is undefined", async () => {
|
||||
mockGenerateResponse.mockResolvedValueOnce({});
|
||||
|
||||
const mg = graph();
|
||||
const result = await mg._retrieveNodesFromData("anything", FILTERS);
|
||||
|
||||
expect(result).toEqual({});
|
||||
});
|
||||
|
||||
it("returns empty map when toolCalls is an empty array", async () => {
|
||||
mockGenerateResponse.mockResolvedValueOnce({ toolCalls: [] });
|
||||
|
||||
const mg = graph();
|
||||
const result = await mg._retrieveNodesFromData("anything", FILTERS);
|
||||
|
||||
expect(result).toEqual({});
|
||||
});
|
||||
|
||||
it("handles malformed JSON in tool call arguments gracefully", async () => {
|
||||
mockGenerateResponse.mockResolvedValueOnce({
|
||||
toolCalls: [
|
||||
{ name: "extract_entities", arguments: "NOT VALID JSON {{{" },
|
||||
],
|
||||
});
|
||||
|
||||
const mg = graph();
|
||||
// Should not throw — the catch block in the source logs the error
|
||||
const result = await mg._retrieveNodesFromData("anything", FILTERS);
|
||||
expect(result).toEqual({});
|
||||
});
|
||||
|
||||
it("handles missing entities array in arguments", async () => {
|
||||
mockGenerateResponse.mockResolvedValueOnce({
|
||||
toolCalls: [
|
||||
{
|
||||
name: "extract_entities",
|
||||
arguments: JSON.stringify({ wrong_key: [] }),
|
||||
},
|
||||
],
|
||||
});
|
||||
|
||||
const mg = graph();
|
||||
// args.entities is undefined → for..of on undefined throws → caught
|
||||
const result = await mg._retrieveNodesFromData("anything", FILTERS);
|
||||
expect(result).toEqual({});
|
||||
});
|
||||
|
||||
it("skips tool calls with unrelated names", async () => {
|
||||
mockGenerateResponse.mockResolvedValueOnce({
|
||||
toolCalls: [
|
||||
{
|
||||
name: "some_other_tool",
|
||||
arguments: JSON.stringify({
|
||||
entities: [{ entity: "X", entity_type: "Y" }],
|
||||
}),
|
||||
},
|
||||
],
|
||||
});
|
||||
|
||||
const mg = graph();
|
||||
const result = await mg._retrieveNodesFromData("anything", FILTERS);
|
||||
expect(result).toEqual({});
|
||||
});
|
||||
|
||||
it("normalises entity names to lowercase with underscores", async () => {
|
||||
mockGenerateResponse.mockResolvedValueOnce({
|
||||
toolCalls: [
|
||||
{
|
||||
name: "extract_entities",
|
||||
arguments: JSON.stringify({
|
||||
entities: [{ entity: "New York City", entity_type: "City Name" }],
|
||||
}),
|
||||
},
|
||||
],
|
||||
});
|
||||
|
||||
const mg = graph();
|
||||
const result = await mg._retrieveNodesFromData("anything", FILTERS);
|
||||
expect(result).toEqual({ new_york_city: "city_name" });
|
||||
});
|
||||
|
||||
it("passes json_object response format and the correct system prompt", async () => {
|
||||
mockGenerateResponse.mockResolvedValueOnce({ toolCalls: [] });
|
||||
|
||||
const mg = graph();
|
||||
await mg._retrieveNodesFromData("test data", FILTERS);
|
||||
|
||||
const [messages, responseFormat] = mockGenerateResponse.mock.calls[0];
|
||||
expect(responseFormat).toEqual({ type: "json_object" });
|
||||
|
||||
const systemMsg = messages[0].content as string;
|
||||
expect(systemMsg.toLowerCase()).toContain("json");
|
||||
expect(systemMsg).toContain("test-user");
|
||||
});
|
||||
});
|
||||
|
||||
// ═══════════════════════════════════════════════════════════════════════════
|
||||
// 2. _establishNodesRelationsFromData – relation extraction
|
||||
// ═══════════════════════════════════════════════════════════════════════════
|
||||
|
||||
describe("_establishNodesRelationsFromData", () => {
|
||||
it("parses a well-formed establish_relationships tool call", async () => {
|
||||
mockGenerateResponse.mockResolvedValueOnce({
|
||||
toolCalls: [
|
||||
{
|
||||
name: "establish_relationships",
|
||||
arguments: JSON.stringify({
|
||||
entities: [
|
||||
{ source: "Alice", relationship: "likes", destination: "Pizza" },
|
||||
],
|
||||
}),
|
||||
},
|
||||
],
|
||||
});
|
||||
|
||||
const mg = graph();
|
||||
const result = await mg._establishNodesRelationsFromData(
|
||||
"Alice likes pizza",
|
||||
FILTERS,
|
||||
{ alice: "person", pizza: "food" },
|
||||
);
|
||||
|
||||
expect(result).toEqual([
|
||||
{ source: "alice", relationship: "likes", destination: "pizza" },
|
||||
]);
|
||||
});
|
||||
|
||||
it("returns empty array when LLM returns a string", async () => {
|
||||
mockGenerateResponse.mockResolvedValueOnce("just a string");
|
||||
|
||||
const mg = graph();
|
||||
const result = await mg._establishNodesRelationsFromData("x", FILTERS, {});
|
||||
expect(result).toEqual([]);
|
||||
});
|
||||
|
||||
it("returns empty array when toolCalls is empty", async () => {
|
||||
mockGenerateResponse.mockResolvedValueOnce({ toolCalls: [] });
|
||||
|
||||
const mg = graph();
|
||||
const result = await mg._establishNodesRelationsFromData("x", FILTERS, {});
|
||||
expect(result).toEqual([]);
|
||||
});
|
||||
|
||||
it("returns empty array when entities key is missing from arguments", async () => {
|
||||
mockGenerateResponse.mockResolvedValueOnce({
|
||||
toolCalls: [
|
||||
{
|
||||
name: "establish_relationships",
|
||||
arguments: JSON.stringify({ not_entities: [] }),
|
||||
},
|
||||
],
|
||||
});
|
||||
|
||||
const mg = graph();
|
||||
const result = await mg._establishNodesRelationsFromData("x", FILTERS, {});
|
||||
// args.entities is undefined → falls back to []
|
||||
expect(result).toEqual([]);
|
||||
});
|
||||
|
||||
it("throws on malformed JSON in tool call arguments (no try/catch in source)", async () => {
|
||||
mockGenerateResponse.mockResolvedValueOnce({
|
||||
toolCalls: [{ name: "establish_relationships", arguments: "<<BROKEN>>" }],
|
||||
});
|
||||
|
||||
const mg = graph();
|
||||
// _establishNodesRelationsFromData does JSON.parse without try/catch
|
||||
await expect(
|
||||
mg._establishNodesRelationsFromData("x", FILTERS, {}),
|
||||
).rejects.toThrow();
|
||||
});
|
||||
|
||||
it("appends JSON format suffix to system prompt (no custom prompt)", async () => {
|
||||
mockGenerateResponse.mockResolvedValueOnce({ toolCalls: [] });
|
||||
|
||||
const mg = graph();
|
||||
await mg._establishNodesRelationsFromData("data", FILTERS, { a: "b" });
|
||||
|
||||
const [messages, responseFormat] = mockGenerateResponse.mock.calls[0];
|
||||
expect(responseFormat).toEqual({ type: "json_object" });
|
||||
|
||||
const systemContent = messages[0].content as string;
|
||||
expect(systemContent.toLowerCase()).toContain("json");
|
||||
expect(systemContent).toContain("test-user");
|
||||
expect(systemContent).not.toContain("USER_ID");
|
||||
// CUSTOM_PROMPT placeholder stays when no custom prompt is configured
|
||||
// (only replaced when config.graphStore.customInstructions is set)
|
||||
});
|
||||
|
||||
it("appends JSON format suffix and custom prompt when configured", async () => {
|
||||
mockGenerateResponse.mockResolvedValueOnce({ toolCalls: [] });
|
||||
|
||||
const mg = graph({
|
||||
customInstructions: "Focus on food relationships only.",
|
||||
});
|
||||
await mg._establishNodesRelationsFromData("data", FILTERS, {});
|
||||
|
||||
const [messages] = mockGenerateResponse.mock.calls[0];
|
||||
const systemContent = messages[0].content as string;
|
||||
expect(systemContent.toLowerCase()).toContain("json");
|
||||
expect(systemContent).toContain("Focus on food relationships only.");
|
||||
});
|
||||
});
|
||||
|
||||
// ═══════════════════════════════════════════════════════════════════════════
|
||||
// 3. _getDeleteEntitiesFromSearchOutput – deletion identification
|
||||
// ═══════════════════════════════════════════════════════════════════════════
|
||||
|
||||
describe("_getDeleteEntitiesFromSearchOutput", () => {
|
||||
const SEARCH_OUTPUT = [
|
||||
{
|
||||
source: "alice",
|
||||
source_id: "1",
|
||||
relationship: "likes",
|
||||
relation_id: "r1",
|
||||
destination: "pizza",
|
||||
destination_id: "2",
|
||||
similarity: 0.95,
|
||||
},
|
||||
];
|
||||
|
||||
it("parses a well-formed delete_graph_memory tool call", async () => {
|
||||
mockGenerateResponse.mockResolvedValueOnce({
|
||||
toolCalls: [
|
||||
{
|
||||
name: "delete_graph_memory",
|
||||
arguments: JSON.stringify({
|
||||
source: "Alice",
|
||||
relationship: "likes",
|
||||
destination: "Pizza",
|
||||
}),
|
||||
},
|
||||
],
|
||||
});
|
||||
|
||||
const mg = graph();
|
||||
const result = await mg._getDeleteEntitiesFromSearchOutput(
|
||||
SEARCH_OUTPUT,
|
||||
"Alice hates pizza",
|
||||
FILTERS,
|
||||
);
|
||||
|
||||
expect(result).toEqual([
|
||||
{ source: "alice", relationship: "likes", destination: "pizza" },
|
||||
]);
|
||||
});
|
||||
|
||||
it("returns empty array when LLM returns a string", async () => {
|
||||
mockGenerateResponse.mockResolvedValueOnce("string response");
|
||||
|
||||
const mg = graph();
|
||||
const result = await mg._getDeleteEntitiesFromSearchOutput(
|
||||
SEARCH_OUTPUT,
|
||||
"x",
|
||||
FILTERS,
|
||||
);
|
||||
expect(result).toEqual([]);
|
||||
});
|
||||
|
||||
it("returns empty array when no tool calls are present", async () => {
|
||||
mockGenerateResponse.mockResolvedValueOnce({ toolCalls: [] });
|
||||
|
||||
const mg = graph();
|
||||
const result = await mg._getDeleteEntitiesFromSearchOutput(
|
||||
SEARCH_OUTPUT,
|
||||
"x",
|
||||
FILTERS,
|
||||
);
|
||||
expect(result).toEqual([]);
|
||||
});
|
||||
|
||||
it("skips non-delete_graph_memory tool calls", async () => {
|
||||
mockGenerateResponse.mockResolvedValueOnce({
|
||||
toolCalls: [
|
||||
{
|
||||
name: "noop",
|
||||
arguments: JSON.stringify({}),
|
||||
},
|
||||
],
|
||||
});
|
||||
|
||||
const mg = graph();
|
||||
const result = await mg._getDeleteEntitiesFromSearchOutput(
|
||||
SEARCH_OUTPUT,
|
||||
"x",
|
||||
FILTERS,
|
||||
);
|
||||
expect(result).toEqual([]);
|
||||
});
|
||||
|
||||
it("collects multiple delete tool calls", async () => {
|
||||
mockGenerateResponse.mockResolvedValueOnce({
|
||||
toolCalls: [
|
||||
{
|
||||
name: "delete_graph_memory",
|
||||
arguments: JSON.stringify({
|
||||
source: "A",
|
||||
relationship: "r1",
|
||||
destination: "B",
|
||||
}),
|
||||
},
|
||||
{
|
||||
name: "delete_graph_memory",
|
||||
arguments: JSON.stringify({
|
||||
source: "C",
|
||||
relationship: "r2",
|
||||
destination: "D",
|
||||
}),
|
||||
},
|
||||
],
|
||||
});
|
||||
|
||||
const mg = graph();
|
||||
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");
|
||||
});
|
||||
|
||||
it("passes json_object format and includes 'json' in system prompt", async () => {
|
||||
mockGenerateResponse.mockResolvedValueOnce({ toolCalls: [] });
|
||||
|
||||
const mg = graph();
|
||||
await mg._getDeleteEntitiesFromSearchOutput(SEARCH_OUTPUT, "data", FILTERS);
|
||||
|
||||
const [messages, responseFormat] = mockGenerateResponse.mock.calls[0];
|
||||
expect(responseFormat).toEqual({ type: "json_object" });
|
||||
|
||||
const systemContent = messages[0].content as string;
|
||||
expect(systemContent.toLowerCase()).toContain("json");
|
||||
expect(systemContent).toContain("test-user");
|
||||
expect(systemContent).not.toContain("USER_ID");
|
||||
});
|
||||
|
||||
it("handles empty searchOutput array", async () => {
|
||||
mockGenerateResponse.mockResolvedValueOnce({ toolCalls: [] });
|
||||
|
||||
const mg = graph();
|
||||
const result = await mg._getDeleteEntitiesFromSearchOutput(
|
||||
[],
|
||||
"data",
|
||||
FILTERS,
|
||||
);
|
||||
expect(result).toEqual([]);
|
||||
});
|
||||
});
|
||||
|
||||
// ═══════════════════════════════════════════════════════════════════════════
|
||||
// 4. Prompt construction — JSON keyword present in every json_object site
|
||||
// ═══════════════════════════════════════════════════════════════════════════
|
||||
|
||||
describe("Prompt construction — all json_object sites include 'json'", () => {
|
||||
it("_retrieveNodesFromData system message includes 'json' for any userId", async () => {
|
||||
for (const userId of ["", "user-1", "special<>chars", "ユーザー"]) {
|
||||
mockGenerateResponse.mockResolvedValueOnce({ toolCalls: [] });
|
||||
const mg = graph();
|
||||
await mg._retrieveNodesFromData("test", { userId });
|
||||
|
||||
const systemMsg = mockGenerateResponse.mock.calls.at(-1)![0][0].content;
|
||||
expect(systemMsg.toLowerCase()).toContain("json");
|
||||
}
|
||||
});
|
||||
|
||||
it("_establishNodesRelationsFromData system message includes 'json' for any userId", async () => {
|
||||
for (const userId of ["", "user-1", "special<>chars"]) {
|
||||
mockGenerateResponse.mockResolvedValueOnce({ toolCalls: [] });
|
||||
const mg = graph();
|
||||
await mg._establishNodesRelationsFromData("test", { userId }, {});
|
||||
|
||||
const systemMsg = mockGenerateResponse.mock.calls.at(-1)![0][0].content;
|
||||
expect(systemMsg.toLowerCase()).toContain("json");
|
||||
}
|
||||
});
|
||||
|
||||
it("_getDeleteEntitiesFromSearchOutput system message includes 'json' for any userId", async () => {
|
||||
for (const userId of ["", "user-1", "special<>chars"]) {
|
||||
mockGenerateResponse.mockResolvedValueOnce({ toolCalls: [] });
|
||||
const mg = graph();
|
||||
await mg._getDeleteEntitiesFromSearchOutput([], "test", { userId });
|
||||
|
||||
const systemMsg = mockGenerateResponse.mock.calls.at(-1)![0][0].content;
|
||||
expect(systemMsg.toLowerCase()).toContain("json");
|
||||
}
|
||||
});
|
||||
});
|
||||
|
||||
// ═══════════════════════════════════════════════════════════════════════════
|
||||
// 5. Edge cases – malformed entity fields in _removeSpacesFromEntities
|
||||
// ═══════════════════════════════════════════════════════════════════════════
|
||||
|
||||
// ═══════════════════════════════════════════════════════════════════════════
|
||||
// 5a. LLM config propagation — graph store uses correct provider & config
|
||||
// Regression test for https://github.com/mem0ai/mem0/issues/3425
|
||||
// ═══════════════════════════════════════════════════════════════════════════
|
||||
|
||||
describe("LLM config propagation to graph store (issue #3425)", () => {
|
||||
const { LLMFactory } = require("../src/utils/factory");
|
||||
|
||||
beforeEach(() => {
|
||||
(LLMFactory.create as jest.Mock).mockClear();
|
||||
});
|
||||
|
||||
it("uses root llm config when no graphStore.llm is provided", () => {
|
||||
const config = {
|
||||
graphStore: {
|
||||
config: {
|
||||
url: "bolt://localhost:7687",
|
||||
username: "neo4j",
|
||||
password: "test",
|
||||
},
|
||||
},
|
||||
embedder: { provider: "openai", config: {} },
|
||||
llm: {
|
||||
provider: "anthropic",
|
||||
config: { model: "claude-sonnet-4-20250514", apiKey: "sk-ant-test" },
|
||||
},
|
||||
} as any;
|
||||
|
||||
new MemoryGraph(config);
|
||||
|
||||
expect(LLMFactory.create).toHaveBeenCalledWith("anthropic", {
|
||||
model: "claude-sonnet-4-20250514",
|
||||
apiKey: "sk-ant-test",
|
||||
});
|
||||
// Both llm and structuredLlm should use the same config
|
||||
expect(LLMFactory.create).toHaveBeenCalledTimes(2);
|
||||
expect(LLMFactory.create).toHaveBeenNthCalledWith(1, "anthropic", {
|
||||
model: "claude-sonnet-4-20250514",
|
||||
apiKey: "sk-ant-test",
|
||||
});
|
||||
expect(LLMFactory.create).toHaveBeenNthCalledWith(2, "anthropic", {
|
||||
model: "claude-sonnet-4-20250514",
|
||||
apiKey: "sk-ant-test",
|
||||
});
|
||||
});
|
||||
|
||||
it("uses graphStore.llm config when provided, overriding root llm", () => {
|
||||
const config = {
|
||||
graphStore: {
|
||||
config: {
|
||||
url: "bolt://localhost:7687",
|
||||
username: "neo4j",
|
||||
password: "test",
|
||||
},
|
||||
llm: {
|
||||
provider: "openai",
|
||||
config: { model: "gpt-4o", apiKey: "sk-openai-test" },
|
||||
},
|
||||
},
|
||||
embedder: { provider: "openai", config: {} },
|
||||
llm: {
|
||||
provider: "anthropic",
|
||||
config: { model: "claude-sonnet-4-20250514", apiKey: "sk-ant-test" },
|
||||
},
|
||||
} as any;
|
||||
|
||||
new MemoryGraph(config);
|
||||
|
||||
// Should use graphStore.llm, NOT root llm
|
||||
expect(LLMFactory.create).toHaveBeenNthCalledWith(1, "openai", {
|
||||
model: "gpt-4o",
|
||||
apiKey: "sk-openai-test",
|
||||
});
|
||||
expect(LLMFactory.create).toHaveBeenNthCalledWith(2, "openai", {
|
||||
model: "gpt-4o",
|
||||
apiKey: "sk-openai-test",
|
||||
});
|
||||
});
|
||||
|
||||
it("falls back to root llm config when graphStore.llm.config is undefined", () => {
|
||||
// Note: in practice, Zod schema requires config when graphStore.llm is
|
||||
// present. This tests the defensive fallback in MemoryGraph itself.
|
||||
const config = {
|
||||
graphStore: {
|
||||
config: {
|
||||
url: "bolt://localhost:7687",
|
||||
username: "neo4j",
|
||||
password: "test",
|
||||
},
|
||||
llm: {
|
||||
provider: "openai",
|
||||
// config explicitly undefined
|
||||
config: undefined,
|
||||
},
|
||||
},
|
||||
embedder: { provider: "openai", config: {} },
|
||||
llm: {
|
||||
provider: "anthropic",
|
||||
config: { model: "claude-sonnet-4-20250514" },
|
||||
},
|
||||
} as any;
|
||||
|
||||
new MemoryGraph(config);
|
||||
|
||||
// Provider from graphStore.llm, but config falls back to root llm.config
|
||||
expect(LLMFactory.create).toHaveBeenNthCalledWith(1, "openai", {
|
||||
model: "claude-sonnet-4-20250514",
|
||||
});
|
||||
});
|
||||
|
||||
it("defaults to openai when neither root nor graphStore llm provider is set", () => {
|
||||
const config = {
|
||||
graphStore: {
|
||||
config: {
|
||||
url: "bolt://localhost:7687",
|
||||
username: "neo4j",
|
||||
password: "test",
|
||||
},
|
||||
},
|
||||
embedder: { provider: "openai", config: {} },
|
||||
llm: { config: { model: "gpt-4" } },
|
||||
} as any;
|
||||
|
||||
new MemoryGraph(config);
|
||||
|
||||
expect(LLMFactory.create).toHaveBeenNthCalledWith(1, "openai", {
|
||||
model: "gpt-4",
|
||||
});
|
||||
});
|
||||
});
|
||||
|
||||
describe("_removeSpacesFromEntities (via _establishNodesRelationsFromData)", () => {
|
||||
it("normalises spaces and case in entity source/relationship/destination", async () => {
|
||||
mockGenerateResponse.mockResolvedValueOnce({
|
||||
toolCalls: [
|
||||
{
|
||||
name: "establish_relationships",
|
||||
arguments: JSON.stringify({
|
||||
entities: [
|
||||
{
|
||||
source: "New York",
|
||||
relationship: "Capital Of",
|
||||
destination: "United States",
|
||||
},
|
||||
],
|
||||
}),
|
||||
},
|
||||
],
|
||||
});
|
||||
|
||||
const mg = graph();
|
||||
const result = await mg._establishNodesRelationsFromData(
|
||||
"test",
|
||||
FILTERS,
|
||||
{},
|
||||
);
|
||||
|
||||
expect(result).toEqual([
|
||||
{
|
||||
source: "new_york",
|
||||
relationship: "capital_of",
|
||||
destination: "united_states",
|
||||
},
|
||||
]);
|
||||
});
|
||||
});
|
||||
@@ -1,177 +0,0 @@
|
||||
import {
|
||||
DELETE_RELATIONS_SYSTEM_PROMPT,
|
||||
EXTRACT_RELATIONS_PROMPT,
|
||||
UPDATE_GRAPH_PROMPT,
|
||||
getDeleteMessages,
|
||||
formatEntities,
|
||||
} from "../src/graphs/utils";
|
||||
|
||||
/**
|
||||
* Regression tests for graph prompts (issue #4248).
|
||||
*
|
||||
* When response_format: { type: "json_object" } is used, OpenAI requires
|
||||
* the word "json" (case-insensitive) to appear in at least one message.
|
||||
* Missing it produces a 400 error.
|
||||
*
|
||||
* Three call sites use json_object today:
|
||||
* 1. _getDeleteEntitiesFromSearchOutput → DELETE_RELATIONS_SYSTEM_PROMPT
|
||||
* 2. _retrieveNodesFromData → inline prompt (graph_memory.ts)
|
||||
* 3. _getRelatedEntities → EXTRACT_RELATIONS_PROMPT + suffix
|
||||
*
|
||||
* See: https://github.com/mem0ai/mem0/issues/4248
|
||||
*/
|
||||
|
||||
// ─── JSON keyword presence ────────────────────────────────────────────────────
|
||||
|
||||
describe("Graph prompts — JSON keyword requirement", () => {
|
||||
it("DELETE_RELATIONS_SYSTEM_PROMPT contains 'json'", () => {
|
||||
expect(DELETE_RELATIONS_SYSTEM_PROMPT.toLowerCase()).toContain("json");
|
||||
});
|
||||
|
||||
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.";
|
||||
expect(withSuffix.toLowerCase()).toContain("json");
|
||||
});
|
||||
|
||||
it("getDeleteMessages system message contains 'json' after USER_ID substitution", () => {
|
||||
const [systemContent] = getDeleteMessages(
|
||||
"alice -- loves -- pizza",
|
||||
"Alice now hates pizza",
|
||||
"user-42",
|
||||
);
|
||||
expect(systemContent.toLowerCase()).toContain("json");
|
||||
});
|
||||
|
||||
it("entity extraction inline prompt contains 'json' (simulated from graph_memory.ts)", () => {
|
||||
// Mirrors the template in _retrieveNodesFromData()
|
||||
const userId = "user-1";
|
||||
const prompt = `You are a smart assistant who understands entities and their types in a given text. If user message contains self reference such as 'I', 'me', 'my' etc. then use ${userId} as the source entity. Extract all the entities from the text. ***DO NOT*** answer the question itself if the given text is a question. Respond in JSON format.`;
|
||||
expect(prompt.toLowerCase()).toContain("json");
|
||||
});
|
||||
});
|
||||
|
||||
// ─── getDeleteMessages ────────────────────────────────────────────────────────
|
||||
|
||||
describe("getDeleteMessages", () => {
|
||||
it("replaces USER_ID with the provided userId in the system prompt", () => {
|
||||
const [system] = getDeleteMessages("mem", "data", "alice-123");
|
||||
expect(system).toContain("alice-123");
|
||||
expect(system).not.toContain("USER_ID");
|
||||
});
|
||||
|
||||
it("includes existing memories and new data in the user prompt", () => {
|
||||
const existing = "bob -- knows -- carol";
|
||||
const newData = "Bob no longer knows Carol";
|
||||
const [, user] = getDeleteMessages(existing, newData, "u1");
|
||||
expect(user).toContain(existing);
|
||||
expect(user).toContain(newData);
|
||||
});
|
||||
|
||||
it("returns a 2-tuple [system, user]", () => {
|
||||
const result = getDeleteMessages("a", "b", "c");
|
||||
expect(result).toHaveLength(2);
|
||||
expect(typeof result[0]).toBe("string");
|
||||
expect(typeof result[1]).toBe("string");
|
||||
});
|
||||
|
||||
// — Malformed / edge-case inputs —
|
||||
|
||||
it("handles empty strings without throwing", () => {
|
||||
expect(() => getDeleteMessages("", "", "")).not.toThrow();
|
||||
const [system, user] = getDeleteMessages("", "", "");
|
||||
expect(system.toLowerCase()).toContain("json");
|
||||
expect(typeof user).toBe("string");
|
||||
});
|
||||
|
||||
it("handles special characters in userId (e.g. angle brackets, quotes)", () => {
|
||||
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",
|
||||
);
|
||||
expect(system).toContain("ユーザー1");
|
||||
expect(user).toContain("日本語メモリ");
|
||||
expect(user).toContain("新しい情報");
|
||||
});
|
||||
|
||||
it("handles very long input strings", () => {
|
||||
const longStr = "x".repeat(100_000);
|
||||
expect(() => getDeleteMessages(longStr, longStr, "u")).not.toThrow();
|
||||
const [system] = getDeleteMessages(longStr, longStr, "u");
|
||||
expect(system.toLowerCase()).toContain("json");
|
||||
});
|
||||
});
|
||||
|
||||
// ─── formatEntities ───────────────────────────────────────────────────────────
|
||||
|
||||
describe("formatEntities", () => {
|
||||
it("formats a single entity triplet", () => {
|
||||
const result = formatEntities([
|
||||
{ source: "Alice", relationship: "knows", destination: "Bob" },
|
||||
]);
|
||||
expect(result).toBe("Alice -- knows -- Bob");
|
||||
});
|
||||
|
||||
it("joins multiple entities with newlines", () => {
|
||||
const result = formatEntities([
|
||||
{ source: "A", relationship: "r1", destination: "B" },
|
||||
{ source: "C", relationship: "r2", destination: "D" },
|
||||
]);
|
||||
expect(result).toBe("A -- r1 -- B\nC -- r2 -- D");
|
||||
});
|
||||
|
||||
it("returns empty string for empty array", () => {
|
||||
expect(formatEntities([])).toBe("");
|
||||
});
|
||||
|
||||
it("preserves special characters in entity fields", () => {
|
||||
const result = formatEntities([
|
||||
{ source: "O'Brien", relationship: 'said "hello"', destination: "café" },
|
||||
]);
|
||||
expect(result).toContain("O'Brien");
|
||||
expect(result).toContain('said "hello"');
|
||||
expect(result).toContain("café");
|
||||
});
|
||||
});
|
||||
|
||||
// ─── Prompt structural invariants ─────────────────────────────────────────────
|
||||
|
||||
describe("Prompt structural invariants", () => {
|
||||
it("DELETE_RELATIONS_SYSTEM_PROMPT contains USER_ID placeholder", () => {
|
||||
expect(DELETE_RELATIONS_SYSTEM_PROMPT).toContain("USER_ID");
|
||||
});
|
||||
|
||||
it("EXTRACT_RELATIONS_PROMPT contains USER_ID placeholder", () => {
|
||||
expect(EXTRACT_RELATIONS_PROMPT).toContain("USER_ID");
|
||||
});
|
||||
|
||||
it("EXTRACT_RELATIONS_PROMPT contains CUSTOM_PROMPT placeholder", () => {
|
||||
expect(EXTRACT_RELATIONS_PROMPT).toContain("CUSTOM_PROMPT");
|
||||
});
|
||||
|
||||
it("UPDATE_GRAPH_PROMPT contains memory template placeholders", () => {
|
||||
expect(UPDATE_GRAPH_PROMPT).toContain("{existing_memories}");
|
||||
expect(UPDATE_GRAPH_PROMPT).toContain("{new_memories}");
|
||||
});
|
||||
|
||||
it("DELETE_RELATIONS_SYSTEM_PROMPT is non-empty and reasonably sized", () => {
|
||||
expect(DELETE_RELATIONS_SYSTEM_PROMPT.length).toBeGreaterThan(100);
|
||||
});
|
||||
|
||||
it("EXTRACT_RELATIONS_PROMPT is non-empty and reasonably sized", () => {
|
||||
expect(EXTRACT_RELATIONS_PROMPT.length).toBeGreaterThan(100);
|
||||
});
|
||||
});
|
||||
@@ -22,18 +22,21 @@ jest.mock("../src/llms/openai", () => ({
|
||||
.fn()
|
||||
.mockImplementation(
|
||||
(messages: Array<{ role: string; content: string }>) => {
|
||||
const hasSystemRole = messages.some((m) => m.role === "system");
|
||||
if (hasSystemRole) {
|
||||
return JSON.stringify({ facts: ["extracted fact from input"] });
|
||||
}
|
||||
// V3 pipeline: single LLM call with additive extraction prompt.
|
||||
const userMsg = messages.find((m) => m.role === "user");
|
||||
const content = userMsg?.content ?? "";
|
||||
const newMsgMatch = content.match(
|
||||
/## New Messages\n([\s\S]*?)(?=\n##|$)/,
|
||||
);
|
||||
const extracted = newMsgMatch
|
||||
? newMsgMatch[1].trim()
|
||||
: "extracted fact from input";
|
||||
return JSON.stringify({
|
||||
memory: [
|
||||
{
|
||||
id: "new",
|
||||
event: "ADD",
|
||||
text: "extracted fact from input",
|
||||
old_memory: "",
|
||||
new_memory: "extracted fact from input",
|
||||
id: "0",
|
||||
text: extracted,
|
||||
attributed_to: "user",
|
||||
},
|
||||
],
|
||||
});
|
||||
@@ -42,9 +45,15 @@ jest.mock("../src/llms/openai", () => ({
|
||||
})),
|
||||
}));
|
||||
|
||||
const mockEmbedding = new Array(1536).fill(0.1);
|
||||
jest.mock("../src/embeddings/openai", () => ({
|
||||
OpenAIEmbedder: jest.fn().mockImplementation(() => ({
|
||||
embed: jest.fn().mockResolvedValue(new Array(1536).fill(0.1)),
|
||||
embed: jest.fn().mockResolvedValue(mockEmbedding),
|
||||
embedBatch: jest
|
||||
.fn()
|
||||
.mockImplementation((texts: string[]) =>
|
||||
Promise.resolve(texts.map(() => mockEmbedding)),
|
||||
),
|
||||
embeddingDims: 1536,
|
||||
})),
|
||||
}));
|
||||
@@ -93,15 +102,16 @@ describe("Memory - add()", () => {
|
||||
});
|
||||
|
||||
test("returns at least one result with an id", async () => {
|
||||
const result: SearchResult = await memory.add("I am a software engineer", {
|
||||
userId,
|
||||
});
|
||||
const result: SearchResult = await memory.add(
|
||||
"I enjoy hiking in the mountains",
|
||||
{ userId },
|
||||
);
|
||||
expect(result.results.length).toBeGreaterThan(0);
|
||||
expect(result.results[0].id).toBeDefined();
|
||||
});
|
||||
|
||||
test("result item has a memory string field", async () => {
|
||||
const result: SearchResult = await memory.add("I am a software engineer", {
|
||||
const result: SearchResult = await memory.add("My favorite color is blue", {
|
||||
userId,
|
||||
});
|
||||
expect(typeof result.results[0].memory).toBe("string");
|
||||
@@ -116,14 +126,14 @@ describe("Memory - add()", () => {
|
||||
expect(result.results.length).toBeGreaterThan(0);
|
||||
});
|
||||
|
||||
test("works with agentId filter instead of userId", async () => {
|
||||
test("works with agentId instead of userId", async () => {
|
||||
const result: SearchResult = await memory.add("test", {
|
||||
agentId: "agent_1",
|
||||
});
|
||||
expect(result.results.length).toBeGreaterThan(0);
|
||||
});
|
||||
|
||||
test("works with runId filter instead of userId", async () => {
|
||||
test("works with runId instead of userId", async () => {
|
||||
const result: SearchResult = await memory.add("test", { runId: "run_1" });
|
||||
expect(result.results.length).toBeGreaterThan(0);
|
||||
});
|
||||
|
||||
@@ -22,18 +22,21 @@ jest.mock("../src/llms/openai", () => ({
|
||||
.fn()
|
||||
.mockImplementation(
|
||||
(messages: Array<{ role: string; content: string }>) => {
|
||||
const hasSystemRole = messages.some((m) => m.role === "system");
|
||||
if (hasSystemRole) {
|
||||
return JSON.stringify({ facts: ["stored fact"] });
|
||||
}
|
||||
// V3 pipeline: single LLM call with additive extraction prompt.
|
||||
// Extract the user input from the prompt to produce unique memories.
|
||||
const userMsg = messages.find((m) => m.role === "user");
|
||||
const content = userMsg?.content ?? "";
|
||||
// Pull the text between "## New Messages" and the next "##"
|
||||
const newMsgMatch = content.match(
|
||||
/## New Messages\n([\s\S]*?)(?=\n##|$)/,
|
||||
);
|
||||
const extracted = newMsgMatch ? newMsgMatch[1].trim() : "stored fact";
|
||||
return JSON.stringify({
|
||||
memory: [
|
||||
{
|
||||
id: "new",
|
||||
event: "ADD",
|
||||
text: "stored fact",
|
||||
old_memory: "",
|
||||
new_memory: "stored fact",
|
||||
id: "0",
|
||||
text: extracted,
|
||||
attributed_to: "user",
|
||||
},
|
||||
],
|
||||
});
|
||||
@@ -42,9 +45,15 @@ jest.mock("../src/llms/openai", () => ({
|
||||
})),
|
||||
}));
|
||||
|
||||
const mockEmbedding = new Array(1536).fill(0.1);
|
||||
jest.mock("../src/embeddings/openai", () => ({
|
||||
OpenAIEmbedder: jest.fn().mockImplementation(() => ({
|
||||
embed: jest.fn().mockResolvedValue(new Array(1536).fill(0.1)),
|
||||
embed: jest.fn().mockResolvedValue(mockEmbedding),
|
||||
embedBatch: jest
|
||||
.fn()
|
||||
.mockImplementation((texts: string[]) =>
|
||||
Promise.resolve(texts.map(() => mockEmbedding)),
|
||||
),
|
||||
embeddingDims: 1536,
|
||||
})),
|
||||
}));
|
||||
@@ -87,7 +96,9 @@ describe("Memory - get()", () => {
|
||||
});
|
||||
|
||||
test("returns the memory matching the ID from add()", async () => {
|
||||
const addResult: SearchResult = await memory.add("I love AI", { userId });
|
||||
const addResult: SearchResult = await memory.add("I love AI", {
|
||||
userId,
|
||||
});
|
||||
const id = addResult.results[0].id;
|
||||
const item: MemoryItem | null = await memory.get(id);
|
||||
expect(item).not.toBeNull();
|
||||
@@ -108,7 +119,9 @@ describe("Memory - get()", () => {
|
||||
});
|
||||
|
||||
test("returns hash and createdAt on stored memory", async () => {
|
||||
const addResult: SearchResult = await memory.add("Hash test", { userId });
|
||||
const addResult: SearchResult = await memory.add("Hash test", {
|
||||
userId,
|
||||
});
|
||||
const item: MemoryItem | null = await memory.get(addResult.results[0].id);
|
||||
expect(typeof item!.hash).toBe("string");
|
||||
expect(item!.createdAt).toBeDefined();
|
||||
@@ -233,13 +246,15 @@ describe("Memory - deleteAll()", () => {
|
||||
await memory.add("Fact B", { userId });
|
||||
const result = await memory.deleteAll({ userId });
|
||||
expect(result.message).toBe("Memories deleted successfully!");
|
||||
const remaining: SearchResult = await memory.getAll({ userId });
|
||||
const remaining: SearchResult = await memory.getAll({
|
||||
filters: { user_id: userId },
|
||||
});
|
||||
expect(remaining.results).toHaveLength(0);
|
||||
});
|
||||
|
||||
test("throws when no filter is provided", async () => {
|
||||
await expect(memory.deleteAll({} as any)).rejects.toThrow(
|
||||
"At least one filter is required",
|
||||
"At least one filter is required to delete all memories",
|
||||
);
|
||||
});
|
||||
});
|
||||
@@ -261,13 +276,17 @@ describe("Memory - getAll()", () => {
|
||||
test("returns all stored memories for the user", async () => {
|
||||
await memory.add("First", { userId });
|
||||
await memory.add("Second", { userId });
|
||||
const result: SearchResult = await memory.getAll({ userId });
|
||||
const result: SearchResult = await memory.getAll({
|
||||
filters: { user_id: userId },
|
||||
});
|
||||
expect(Array.isArray(result.results)).toBe(true);
|
||||
expect(result.results.length).toBeGreaterThanOrEqual(2);
|
||||
});
|
||||
|
||||
test("each result has id and memory fields", async () => {
|
||||
const result: SearchResult = await memory.getAll({ userId });
|
||||
const result: SearchResult = await memory.getAll({
|
||||
filters: { user_id: userId },
|
||||
});
|
||||
for (const item of result.results) {
|
||||
expect(item.id).toBeDefined();
|
||||
expect(typeof item.memory).toBe("string");
|
||||
@@ -276,7 +295,7 @@ describe("Memory - getAll()", () => {
|
||||
|
||||
test("returns empty array when no memories exist", async () => {
|
||||
const result: SearchResult = await memory.getAll({
|
||||
userId: "no_such_user",
|
||||
filters: { user_id: "no_such_user" },
|
||||
});
|
||||
expect(result.results).toHaveLength(0);
|
||||
});
|
||||
@@ -298,12 +317,16 @@ describe("Memory - search()", () => {
|
||||
});
|
||||
|
||||
test("returns SearchResult with results array", async () => {
|
||||
const result: SearchResult = await memory.search("TypeScript", { userId });
|
||||
const result: SearchResult = await memory.search("TypeScript", {
|
||||
filters: { user_id: userId },
|
||||
});
|
||||
expect(Array.isArray(result.results)).toBe(true);
|
||||
});
|
||||
|
||||
test("returns results with score field", async () => {
|
||||
const result: SearchResult = await memory.search("content", { userId });
|
||||
const result: SearchResult = await memory.search("content", {
|
||||
filters: { user_id: userId },
|
||||
});
|
||||
if (result.results.length > 0) {
|
||||
expect(typeof result.results[0].score).toBe("number");
|
||||
}
|
||||
@@ -311,13 +334,13 @@ describe("Memory - search()", () => {
|
||||
|
||||
test("throws when no userId/agentId/runId provided", async () => {
|
||||
await expect(memory.search("query", {} as any)).rejects.toThrow(
|
||||
"One of the filters: userId, agentId or runId is required!",
|
||||
"filters must contain at least one of: user_id, agent_id, run_id",
|
||||
);
|
||||
});
|
||||
|
||||
test("returns empty results for user with no memories", async () => {
|
||||
const result: SearchResult = await memory.search("query", {
|
||||
userId: "empty_user",
|
||||
filters: { user_id: "empty_user" },
|
||||
});
|
||||
expect(result.results).toHaveLength(0);
|
||||
});
|
||||
@@ -338,14 +361,18 @@ describe("Memory - history()", () => {
|
||||
});
|
||||
|
||||
test("records ADD event after add()", async () => {
|
||||
const addResult: SearchResult = await memory.add("New fact", { userId });
|
||||
const addResult: SearchResult = await memory.add("New fact", {
|
||||
userId,
|
||||
});
|
||||
const history = await memory.history(addResult.results[0].id);
|
||||
expect(Array.isArray(history)).toBe(true);
|
||||
expect(history.length).toBeGreaterThan(0);
|
||||
});
|
||||
|
||||
test("records additional entry after update()", async () => {
|
||||
const addResult: SearchResult = await memory.add("Before", { userId });
|
||||
const addResult: SearchResult = await memory.add("Before", {
|
||||
userId,
|
||||
});
|
||||
const id = addResult.results[0].id;
|
||||
await memory.update(id, "After");
|
||||
const history = await memory.history(id);
|
||||
|
||||
@@ -16,26 +16,25 @@ jest.mock("../src/llms/google", () => ({
|
||||
GoogleLLM: jest.fn(),
|
||||
}));
|
||||
|
||||
// ─── Content-based LLM mock (reviewer #9) ────────────────
|
||||
// Returns facts for system-prompt calls, memory actions for user-only calls.
|
||||
// ─── Content-based LLM mock (V3 additive extraction pipeline) ─────────
|
||||
jest.mock("../src/llms/openai", () => ({
|
||||
OpenAILLM: jest.fn().mockImplementation(() => ({
|
||||
generateResponse: jest
|
||||
.fn()
|
||||
.mockImplementation(
|
||||
(messages: Array<{ role: string; content: string }>) => {
|
||||
const hasSystemRole = messages.some((m) => m.role === "system");
|
||||
if (hasSystemRole) {
|
||||
return JSON.stringify({ facts: ["test fact"] });
|
||||
}
|
||||
const userMsg = messages.find((m) => m.role === "user");
|
||||
const content = userMsg?.content ?? "";
|
||||
const newMsgMatch = content.match(
|
||||
/## New Messages\n([\s\S]*?)(?=\n##|$)/,
|
||||
);
|
||||
const extracted = newMsgMatch ? newMsgMatch[1].trim() : "test fact";
|
||||
return JSON.stringify({
|
||||
memory: [
|
||||
{
|
||||
id: "new",
|
||||
event: "ADD",
|
||||
text: "test fact",
|
||||
old_memory: "",
|
||||
new_memory: "test fact",
|
||||
id: "0",
|
||||
text: extracted,
|
||||
attributed_to: "user",
|
||||
},
|
||||
],
|
||||
});
|
||||
@@ -44,9 +43,15 @@ jest.mock("../src/llms/openai", () => ({
|
||||
})),
|
||||
}));
|
||||
|
||||
const mockEmbedding = new Array(1536).fill(0.1);
|
||||
jest.mock("../src/embeddings/openai", () => ({
|
||||
OpenAIEmbedder: jest.fn().mockImplementation(() => ({
|
||||
embed: jest.fn().mockResolvedValue(new Array(1536).fill(0.1)),
|
||||
embed: jest.fn().mockResolvedValue(mockEmbedding),
|
||||
embedBatch: jest
|
||||
.fn()
|
||||
.mockImplementation((texts: string[]) =>
|
||||
Promise.resolve(texts.map(() => mockEmbedding)),
|
||||
),
|
||||
embeddingDims: 1536,
|
||||
})),
|
||||
}));
|
||||
@@ -114,12 +119,16 @@ describe("Memory - reset()", () => {
|
||||
const userId = `reset_test_${Date.now()}`;
|
||||
|
||||
await mem.add("Remember this fact", { userId });
|
||||
const before: SearchResult = await mem.getAll({ userId });
|
||||
const before: SearchResult = await mem.getAll({
|
||||
filters: { user_id: userId },
|
||||
});
|
||||
expect(before.results.length).toBeGreaterThan(0);
|
||||
|
||||
await mem.reset();
|
||||
|
||||
const after: SearchResult = await mem.getAll({ userId });
|
||||
const after: SearchResult = await mem.getAll({
|
||||
filters: { user_id: userId },
|
||||
});
|
||||
expect(after.results).toHaveLength(0);
|
||||
});
|
||||
});
|
||||
|
||||
@@ -944,7 +944,7 @@ describe("Memory class – backward compat with all providers", () => {
|
||||
disableHistory: true,
|
||||
});
|
||||
|
||||
await mem.getAll({ userId: "u1" });
|
||||
await mem.getAll({ filters: { user_id: "u1" } });
|
||||
|
||||
const embedder = mockEmbedderFactory.create.mock.results[0].value;
|
||||
expect(embedder.embed).not.toHaveBeenCalledWith("dimension probe");
|
||||
@@ -967,7 +967,7 @@ describe("Memory class – backward compat with all providers", () => {
|
||||
const mockEmbedder768 = createMockEmbedder(768);
|
||||
mockEmbedderFactory.create.mockReturnValue(mockEmbedder768);
|
||||
|
||||
await mem.getAll({ userId: "u1" });
|
||||
await mem.getAll({ filters: { user_id: "u1" } });
|
||||
expect(mockEmbedder768.embed).not.toHaveBeenCalledWith("dimension probe");
|
||||
});
|
||||
|
||||
@@ -982,7 +982,7 @@ describe("Memory class – backward compat with all providers", () => {
|
||||
disableHistory: true,
|
||||
});
|
||||
|
||||
await mem.getAll({ userId: "u1" });
|
||||
await mem.getAll({ filters: { user_id: "u1" } });
|
||||
expect(mockEmbedder768.embed).toHaveBeenCalledWith("dimension probe");
|
||||
|
||||
const vsCreateCall = mockVectorStoreFactory.create.mock.calls[0];
|
||||
@@ -1003,7 +1003,7 @@ describe("Memory class – backward compat with all providers", () => {
|
||||
disableHistory: true,
|
||||
});
|
||||
|
||||
await mem.getAll({ userId: "u1" });
|
||||
await mem.getAll({ filters: { user_id: "u1" } });
|
||||
expect(mockVStore.initialize).toHaveBeenCalled();
|
||||
});
|
||||
|
||||
@@ -1048,11 +1048,13 @@ describe("Memory class – backward compat with all providers", () => {
|
||||
});
|
||||
|
||||
// getAll
|
||||
const all = await mem.getAll({ userId: "u1" });
|
||||
const all = await mem.getAll({ filters: { user_id: "u1" } });
|
||||
expect(all).toBeDefined();
|
||||
|
||||
// search
|
||||
const searchResult = await mem.search("query", { userId: "u1" });
|
||||
const searchResult = await mem.search("query", {
|
||||
filters: { user_id: "u1" },
|
||||
});
|
||||
expect(searchResult).toBeDefined();
|
||||
|
||||
// get
|
||||
@@ -1093,7 +1095,7 @@ describe("Memory class – backward compat with all providers", () => {
|
||||
disableHistory: true,
|
||||
});
|
||||
|
||||
await mem.getAll({ userId: "u1" });
|
||||
await mem.getAll({ filters: { user_id: "u1" } });
|
||||
expect(mockVectorStoreFactory.create).toHaveBeenCalledTimes(1);
|
||||
|
||||
await mem.reset();
|
||||
@@ -1120,12 +1122,12 @@ describe("Memory class – backward compat with all providers", () => {
|
||||
disableHistory: true,
|
||||
});
|
||||
|
||||
await expect(mem.getAll({ userId: "u1" })).rejects.toThrow(
|
||||
"auto-detect embedding dimension",
|
||||
);
|
||||
await expect(mem.search("q", { userId: "u1" })).rejects.toThrow(
|
||||
await expect(mem.getAll({ filters: { user_id: "u1" } })).rejects.toThrow(
|
||||
"auto-detect embedding dimension",
|
||||
);
|
||||
await expect(
|
||||
mem.search("q", { filters: { user_id: "u1" } }),
|
||||
).rejects.toThrow("auto-detect embedding dimension");
|
||||
await expect(mem.get("id")).rejects.toThrow(
|
||||
"auto-detect embedding dimension",
|
||||
);
|
||||
|
||||
@@ -13,13 +13,14 @@ const external = [
|
||||
"ollama",
|
||||
"@google/genai",
|
||||
"@mistralai/mistralai",
|
||||
"neo4j-driver",
|
||||
"@supabase/supabase-js",
|
||||
"@azure/search-documents",
|
||||
"@azure/identity",
|
||||
"cloudflare",
|
||||
"@cloudflare/workers-types",
|
||||
"@langchain/core",
|
||||
"compromise",
|
||||
"natural",
|
||||
];
|
||||
|
||||
export default defineConfig([
|
||||
|
||||
+42
-24
@@ -29,6 +29,9 @@ warnings.filterwarnings("default", category=DeprecationWarning)
|
||||
# Setup user config
|
||||
setup_config()
|
||||
|
||||
# Entity parameters that must be passed via filters, not top-level
|
||||
ENTITY_PARAMS = frozenset({"user_id", "agent_id", "app_id", "run_id"})
|
||||
|
||||
|
||||
class MemoryClient:
|
||||
"""Client for interacting with the Mem0 API.
|
||||
@@ -162,12 +165,9 @@ class MemoryClient:
|
||||
elif not isinstance(messages, list):
|
||||
raise ValueError(f"messages must be str, dict, or list[dict], got {type(messages).__name__}")
|
||||
|
||||
# Force v1.1 format for all add operations
|
||||
kwargs["output_format"] = "v1.1"
|
||||
|
||||
kwargs = self._prepare_params(kwargs)
|
||||
payload = self._prepare_payload(messages, kwargs)
|
||||
response = self.client.post("/v1/memories/", json=payload)
|
||||
response = self.client.post("/v3/memories/", json=payload)
|
||||
response.raise_for_status()
|
||||
if "metadata" in kwargs:
|
||||
del kwargs["metadata"]
|
||||
@@ -204,8 +204,7 @@ class MemoryClient:
|
||||
|
||||
Args:
|
||||
options: Typed options for the get_all operation (GetAllMemoryOptions).
|
||||
**kwargs: Optional parameters for filtering (user_id, agent_id,
|
||||
app_id, top_k, page, page_size).
|
||||
**kwargs: Optional parameters for filtering (filters, page, page_size).
|
||||
|
||||
Returns:
|
||||
A dictionary containing memories in v1.1 format: {"results": [...]}
|
||||
@@ -218,6 +217,14 @@ class MemoryClient:
|
||||
NetworkError: If network connectivity issues occur.
|
||||
MemoryNotFoundError: If the memory doesn't exist (for updates/deletes).
|
||||
"""
|
||||
# Reject top-level entity params - must use filters instead
|
||||
invalid_keys = ENTITY_PARAMS & set(kwargs.keys())
|
||||
if invalid_keys:
|
||||
raise ValueError(
|
||||
f"Top-level entity parameters {invalid_keys} are not supported in get_all(). "
|
||||
f"Use filters={{'user_id': '...'}} instead."
|
||||
)
|
||||
|
||||
kwargs = {**(options.model_dump(exclude_unset=True) if options else {}), **kwargs}
|
||||
params = self._prepare_params(kwargs)
|
||||
|
||||
@@ -267,11 +274,19 @@ class MemoryClient:
|
||||
NetworkError: If network connectivity issues occur.
|
||||
MemoryNotFoundError: If the memory doesn't exist (for updates/deletes).
|
||||
"""
|
||||
# Reject top-level entity params - must use filters instead
|
||||
invalid_keys = ENTITY_PARAMS & set(kwargs.keys())
|
||||
if invalid_keys:
|
||||
raise ValueError(
|
||||
f"Top-level entity parameters {invalid_keys} are not supported in search(). "
|
||||
f"Use filters={{'user_id': '...'}} instead."
|
||||
)
|
||||
|
||||
kwargs = {**(options.model_dump(exclude_unset=True) if options else {}), **kwargs}
|
||||
params = self._prepare_params(kwargs)
|
||||
payload = {"query": query, **params}
|
||||
|
||||
response = self.client.post("/v2/memories/search/", json=payload)
|
||||
response = self.client.post("/v3/memories/search/", json=payload)
|
||||
response.raise_for_status()
|
||||
if "metadata" in kwargs:
|
||||
del kwargs["metadata"]
|
||||
@@ -283,12 +298,7 @@ class MemoryClient:
|
||||
"sync_type": "sync",
|
||||
},
|
||||
)
|
||||
result = response.json()
|
||||
|
||||
# Ensure v1.1 format (wrap raw list if needed)
|
||||
if isinstance(result, list):
|
||||
return {"results": result}
|
||||
return result
|
||||
return response.json()
|
||||
|
||||
@api_error_handler
|
||||
def update(
|
||||
@@ -1083,12 +1093,9 @@ class AsyncMemoryClient:
|
||||
elif not isinstance(messages, list):
|
||||
raise ValueError(f"messages must be str, dict, or list[dict], got {type(messages).__name__}")
|
||||
|
||||
# Force v1.1 format for all add operations
|
||||
kwargs["output_format"] = "v1.1"
|
||||
|
||||
kwargs = self._prepare_params(kwargs)
|
||||
payload = self._prepare_payload(messages, kwargs)
|
||||
response = await self.async_client.post("/v1/memories/", json=payload)
|
||||
response = await self.async_client.post("/v3/memories/", json=payload)
|
||||
response.raise_for_status()
|
||||
if "metadata" in kwargs:
|
||||
del kwargs["metadata"]
|
||||
@@ -1122,6 +1129,14 @@ class AsyncMemoryClient:
|
||||
NetworkError: If network connectivity issues occur.
|
||||
MemoryNotFoundError: If the memory doesn't exist (for updates/deletes).
|
||||
"""
|
||||
# Reject top-level entity params - must use filters instead
|
||||
invalid_keys = ENTITY_PARAMS & set(kwargs.keys())
|
||||
if invalid_keys:
|
||||
raise ValueError(
|
||||
f"Top-level entity parameters {invalid_keys} are not supported in get_all(). "
|
||||
f"Use filters={{'user_id': '...'}} instead."
|
||||
)
|
||||
|
||||
kwargs = {**(options.model_dump(exclude_unset=True) if options else {}), **kwargs}
|
||||
params = self._prepare_params(kwargs)
|
||||
|
||||
@@ -1171,11 +1186,19 @@ class AsyncMemoryClient:
|
||||
NetworkError: If network connectivity issues occur.
|
||||
MemoryNotFoundError: If the memory doesn't exist (for updates/deletes).
|
||||
"""
|
||||
# Reject top-level entity params - must use filters instead
|
||||
invalid_keys = ENTITY_PARAMS & set(kwargs.keys())
|
||||
if invalid_keys:
|
||||
raise ValueError(
|
||||
f"Top-level entity parameters {invalid_keys} are not supported in search(). "
|
||||
f"Use filters={{'user_id': '...'}} instead."
|
||||
)
|
||||
|
||||
kwargs = {**(options.model_dump(exclude_unset=True) if options else {}), **kwargs}
|
||||
params = self._prepare_params(kwargs)
|
||||
payload = {"query": query, **params}
|
||||
|
||||
response = await self.async_client.post("/v2/memories/search/", json=payload)
|
||||
response = await self.async_client.post("/v3/memories/search/", json=payload)
|
||||
response.raise_for_status()
|
||||
if "metadata" in kwargs:
|
||||
del kwargs["metadata"]
|
||||
@@ -1187,12 +1210,7 @@ class AsyncMemoryClient:
|
||||
"sync_type": "async",
|
||||
},
|
||||
)
|
||||
result = response.json()
|
||||
|
||||
# Ensure v1.1 format (wrap raw list if needed)
|
||||
if isinstance(result, list):
|
||||
return {"results": result}
|
||||
return result
|
||||
return response.json()
|
||||
|
||||
@api_error_handler
|
||||
async def update(
|
||||
|
||||
+20
-13
@@ -2,6 +2,9 @@
|
||||
|
||||
These models provide IDE autocompletion, runtime validation, and type safety.
|
||||
Methods accept both typed options and **kwargs for backward compatibility.
|
||||
|
||||
Identity fields (user_id, agent_id, app_id, run_id) must be passed inside
|
||||
the ``filters`` dict — the v3 API does not accept them at the top level.
|
||||
"""
|
||||
|
||||
from typing import Any, Dict, List, Optional, Union
|
||||
@@ -9,18 +12,16 @@ from typing import Any, Dict, List, Optional, Union
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
|
||||
class EntityOptions(BaseModel):
|
||||
"""Identity options for add/delete operations (top-level entity IDs)."""
|
||||
class AddMemoryOptions(BaseModel):
|
||||
"""Options for the add() method.
|
||||
|
||||
user_id: Optional[str] = Field(default=None, description="The user ID to associate with the memory")
|
||||
agent_id: Optional[str] = Field(default=None, description="The agent ID to associate with the memory")
|
||||
app_id: Optional[str] = Field(default=None, description="The app ID to associate with the memory")
|
||||
run_id: Optional[str] = Field(default=None, description="The run ID to associate with the memory")
|
||||
|
||||
|
||||
class AddMemoryOptions(EntityOptions):
|
||||
"""Options for the add() method."""
|
||||
Identity fields (user_id, agent_id, app_id, run_id) must be passed inside
|
||||
the ``filters`` dict — the v3 API does not accept them at the top level.
|
||||
"""
|
||||
|
||||
filters: Optional[Dict[str, Any]] = Field(
|
||||
default=None, description="Filters containing entity IDs (e.g. {'user_id': '...'})"
|
||||
)
|
||||
metadata: Optional[Dict[str, Any]] = Field(default=None, description="Additional metadata for the memory")
|
||||
infer: Optional[bool] = Field(default=None, description="Whether to infer memories from the input")
|
||||
custom_categories: Optional[List[Dict[str, Any]]] = Field(
|
||||
@@ -72,10 +73,16 @@ class GetAllMemoryOptions(BaseModel):
|
||||
categories: Optional[List[str]] = Field(default=None, description="Categories to filter by")
|
||||
|
||||
|
||||
class DeleteAllMemoryOptions(EntityOptions):
|
||||
"""Options for the delete_all() method."""
|
||||
class DeleteAllMemoryOptions(BaseModel):
|
||||
"""Options for the delete_all() method.
|
||||
|
||||
pass
|
||||
Identity fields (user_id, agent_id, app_id, run_id) must be passed inside
|
||||
the ``filters`` dict — the API does not accept them at the top level.
|
||||
"""
|
||||
|
||||
filters: Optional[Dict[str, Any]] = Field(
|
||||
default=None, description="Filters containing entity IDs (e.g. {'user_id': '...'})"
|
||||
)
|
||||
|
||||
|
||||
class UpdateMemoryOptions(BaseModel):
|
||||
|
||||
@@ -5,7 +5,6 @@ from pydantic import BaseModel, Field
|
||||
|
||||
from mem0.configs.rerankers.config import RerankerConfig
|
||||
from mem0.embeddings.configs import EmbedderConfig
|
||||
from mem0.graphs.configs import GraphStoreConfig
|
||||
from mem0.llms.configs import LlmConfig
|
||||
from mem0.vector_stores.configs import VectorStoreConfig
|
||||
|
||||
@@ -44,10 +43,6 @@ class MemoryConfig(BaseModel):
|
||||
description="Path to the history database",
|
||||
default=os.path.join(mem0_dir, "history.db"),
|
||||
)
|
||||
graph_store: GraphStoreConfig = Field(
|
||||
description="Configuration for the graph",
|
||||
default_factory=GraphStoreConfig,
|
||||
)
|
||||
reranker: Optional[RerankerConfig] = Field(
|
||||
description="Configuration for the reranker",
|
||||
default=None,
|
||||
@@ -60,10 +55,6 @@ class MemoryConfig(BaseModel):
|
||||
description="Custom instructions for fact extraction",
|
||||
default=None,
|
||||
)
|
||||
custom_update_memory_prompt: Optional[str] = Field(
|
||||
description="Custom prompt for the update memory",
|
||||
default=None,
|
||||
)
|
||||
|
||||
|
||||
class AzureConfig(BaseModel):
|
||||
|
||||
+604
-1
@@ -1,4 +1,5 @@
|
||||
from datetime import datetime
|
||||
import json
|
||||
from datetime import datetime, timezone
|
||||
|
||||
MEMORY_ANSWER_PROMPT = """
|
||||
You are an expert at answering questions based on the provided memories. Your task is to provide accurate and concise answers to the questions by leveraging the information given in the memories.
|
||||
@@ -457,3 +458,605 @@ def get_update_memory_messages(retrieved_old_memory_dict, response_content, cust
|
||||
|
||||
Do not return anything except the JSON format.
|
||||
"""
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# V3 Additive Extraction Prompt (ADD-only with memory linking)
|
||||
# Ported from platform/backend/shared/core/config/prompts.py
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
ADDITIVE_EXTRACTION_PROMPT = """
|
||||
|
||||
# ROLE
|
||||
|
||||
You are a Memory Extractor — a precise, evidence-bound processor responsible for extracting rich, contextual memories from conversations. Your sole operation is ADD: identify every piece of memorable information and produce self-contained, contextually rich factual statements.
|
||||
|
||||
You extract from BOTH user and assistant messages. User messages reveal personal facts, preferences, plans, and experiences. Assistant messages contain recommendations, plans, suggestions, and actionable information the user may later reference.
|
||||
|
||||
Accuracy and completeness are critical. Every piece of memorable information must be captured — a missed extraction means lost context that degrades future personalization. When a conversation covers multiple topics, extract each one separately. Do not let a dominant topic cause you to miss secondary information.
|
||||
|
||||
# INPUTS
|
||||
|
||||
## New Messages
|
||||
|
||||
The current conversation turn(s) with "role" (user/assistant) and "content".
|
||||
|
||||
Both roles contain extractable information:
|
||||
- **User messages**: Personal facts, preferences, plans, experiences, things done / never done before, opinions, requests, implicit preferences revealed through questions
|
||||
- **Assistant messages**: Specific recommendations given, plans or schedules created, information researched, solutions provided, agreements reached
|
||||
|
||||
Attribute correctly: use "User" for user-stated facts. For assistant-generated content, frame in terms of the user's context (e.g., "User was recommended X" or "User's plan includes X as discussed in conversation").
|
||||
|
||||
Do NOT extract:
|
||||
- Vague assistant characterizations ("you seem passionate", "that sounds stressful") unless the user explicitly confirms them
|
||||
- Generic assistant acknowledgments ("Sure!", "Great question!")
|
||||
- Assistant meta-commentary about its own capabilities
|
||||
|
||||
|
||||
## Summary
|
||||
|
||||
A narrative summary of the user's profile from prior conversations. May be empty for new users. Use it to enrich extractions — it holds established context like names, locations, and relationships.
|
||||
|
||||
|
||||
## Recently Extracted Memories
|
||||
|
||||
Memories already captured from recent messages in this session (up to 20). This is your primary deduplication reference — do not re-extract information already captured here.
|
||||
|
||||
|
||||
## Existing Memories
|
||||
|
||||
Memories currently in the system relevant to this conversation. Formatted as:
|
||||
[{"id": "uuid-string", "text": "..."}, ...]
|
||||
|
||||
Use these ONLY for deduplication and linking — do NOT extract new memories from Existing Memories. Your extractions must come exclusively from New Messages. If new information in New Messages is semantically equivalent to an Existing Memory with no meaningful new context, skip it.
|
||||
|
||||
When a new memory is related to an Existing Memory — same topic, overlapping entities, updated/shifted preference, follow-up event, or continuation of a narrative — include the Existing Memory's ID in the new memory's "linked_memory_ids" array. Your ADD output IDs remain sequential ("0", "1", ...) but linked_memory_ids uses the UUIDs from this list.
|
||||
|
||||
|
||||
IMPORTANT: An existing memory about an entity (e.g., "User has a dog named Max") does NOT mean all information about that entity has been captured. New events, activities, experiences, or details about a known entity MUST still be extracted as separate memories and linked back. Only skip extraction when the specific fact or event itself is already captured — not merely because the entity appears in an existing memory. "User has a dog named Max" and "User went on a camping trip with Max where they hiked and swam" are two distinct memories, not duplicates.
|
||||
|
||||
|
||||
## Last k Messages
|
||||
|
||||
Recent messages (up to 20) preceding New Messages. Use to resolve references and pronouns in New Messages.
|
||||
|
||||
|
||||
## Observation Date
|
||||
|
||||
When the conversation actually took place (e.g., "2023-05-24"). This is your ONLY temporal anchor for resolving time references.
|
||||
|
||||
Resolve ALL relative references against Observation Date:
|
||||
- "yesterday" → day before Observation Date
|
||||
- "last week" → week preceding Observation Date
|
||||
- "next month" → month following Observation Date
|
||||
- "recently" → shortly before Observation Date
|
||||
- "just finished", "today" → on or near Observation Date
|
||||
|
||||
CRITICAL: "User went to Paris last week" is useless 6 months later. "User went to Paris the week of May 15, 2023" is meaningful forever. Always ground relative references to specific dates.
|
||||
|
||||
|
||||
## Current Date
|
||||
|
||||
Today's system date. May be years after Observation Date. Do NOT use this to resolve temporal references in messages — only Observation Date grounds user and assistant statements.
|
||||
|
||||
|
||||
## Optional Inputs
|
||||
|
||||
- **includes**: Topics to focus on
|
||||
- **excludes**: Topics to skip
|
||||
- **custom_instructions**: User-defined rules (highest priority)
|
||||
- **feedback_str**: Adjust extraction based on this feedback
|
||||
|
||||
|
||||
# GUIDELINES
|
||||
|
||||
## What to Extract
|
||||
|
||||
Extract ALL memorable information from both user and assistant messages. Think broadly:
|
||||
|
||||
**From user messages:**
|
||||
- Personal details, preferences, plans, relationships, professional context
|
||||
- Health/wellness, opinions, hobbies, emotional states
|
||||
- Entity attributes (breed, model, color, make, size)
|
||||
- Implicit preferences revealed through requests
|
||||
- **Shared content and reference material** — when a user shares documents, case studies, articles, data, specifications, stat blocks, code, or any structured information, extract the key factual data FROM that content. The user shared it because they want it remembered.
|
||||
- Firsts and milestones — 'first call-out', 'just started', 'recently joined', etc.
|
||||
- Specific foods, meals, and who was present (e.g. 'dinner with mom — salads, sandwiches, homemade desserts').
|
||||
- Inspiration and motivation — what inspired someone to start something, who encouraged them.
|
||||
|
||||
**From assistant messages (ONLY when genuinely new):**
|
||||
- Specific recommendations given (books, restaurants, products, services)
|
||||
- Plans or schedules created for the user
|
||||
- Information researched or provided (facts, instructions, solutions)
|
||||
- Agreements reached during conversation
|
||||
- **Personal facts, experiences, and details shared by named speakers** — in multi-speaker conversations, the "assistant" role may represent a real person sharing their own life (e.g., "Maria: I just got a new cat named Bailey"). Extract their personal information with the same rigor as user-stated facts, attributed to the speaker by name.
|
||||
|
||||
Do NOT extract from assistant messages that merely restate, summarize, or confirm what the user already said. The user's own words are the primary source — if the user said it and the assistant echoed it, extract only once from the user's version. Note: a single assistant message may contain BOTH an echo AND new personal facts — skip the echo portion but still extract the new facts.
|
||||
|
||||
Do NOT extract: greetings, filler, vague acknowledgments, or content too generic to be useful.
|
||||
|
||||
**When in doubt, extract.** A slightly redundant memory is far less costly than a missing one. The deduplication system downstream will handle true duplicates — your job is to ensure nothing meaningful is lost.
|
||||
|
||||
### Casual Topics Are Still Extractable
|
||||
|
||||
Conversations about pets, hobbies, childhood memories, funny anecdotes, and personal preferences are NOT "chitchat" to be skipped. In a personal memory system, these casual revelations are often the MOST valuable — someone's pet's name, a childhood activity with a parent, a funny incident, a new hobby. Only skip messages that are PURELY phatic ("Hi!", "Sounds good!", "Thanks!") with zero informational content.
|
||||
|
||||
### Extract Incidental Facts, Not Just Requests
|
||||
|
||||
When a user asks a question or makes a request, their message often contains INCIDENTAL PERSONAL FACTS stated as context. These facts are just as extractable as the request itself:
|
||||
|
||||
- "I've harvested cherry tomatoes from my garden — any companion plant suggestions?" → Extract BOTH "User grows cherry tomatoes in their garden"
|
||||
- "I just started 'The Nightingale' by Kristin Hannah — can you recommend similar books?" → Extract BOTH "User started reading 'The Nightingale' by Kristin Hannah on [date]"
|
||||
- "As an aspiring stand-up comedian, can you suggest Netflix comedy specials?" → Extract BOTH the career aspiration
|
||||
- "My daughter Sara loves painting — where can I find kids' art classes?" → Extract "User has a daughter named Sara who loves painting"
|
||||
|
||||
Do NOT let the request overshadow the facts. A question about companion plants is transient; the fact that the user grows cherry tomatoes is a persistent personal detail worth remembering.
|
||||
|
||||
**IMPORTANT — Extract ALL dimensions of a conversation.** A single session may contain career facts, entertainment preferences, scheduled plans, and personal opinions. Extract each dimension as a separate memory. Do not let one dominant topic cause you to miss secondary information.
|
||||
|
||||
### Shared Photos and Images
|
||||
|
||||
When a message contains a photo description (e.g., "[Shared photo: ...]" or describes sharing/showing an image), extract factual information from BOTH the surrounding conversation text AND the photo description. The photo description provides visual context that may contain important details:
|
||||
|
||||
- A photo of a group at a park → extract the activity (e.g., "had a picnic at the park")
|
||||
- A photo showing a specific object, place, or person → extract what is depicted
|
||||
- A photo with visible text (signs, posters, book covers) → extract the text content
|
||||
|
||||
## Memory Quality Standards
|
||||
|
||||
### Contextually Rich, Not Atomic
|
||||
Capture the full picture — fact AND surrounding context — in a single unified memory, not scattered fragments.
|
||||
|
||||
Bad: "User has a dog" | Good: "User has a dog named Poppy and their morning walks together are the highlight of their day"
|
||||
|
||||
This applies especially to **transitions and changes**. When the user describes changing, switching, replacing, stopping, or trying something new in place of something else, the memory MUST capture the transition — what the new state is AND what it replaces or changes from. The relationship between old and new is critical context. Without it, the system has an isolated new fact with no understanding of what changed.
|
||||
|
||||
Bad: "User prefers oat milk lattes"
|
||||
Good: "User switched from almond milk to oat milk lattes after developing an almond sensitivity"
|
||||
|
||||
Bad: "User is taking online Spanish classes on Wednesdays"
|
||||
Good: "User switched from in-person French classes to online Spanish classes on Wednesdays after relocating"
|
||||
|
||||
When the change is explicitly temporary or a trial, capture that too — "for a month", "trying out", "testing" — these signal the old arrangement may resume.
|
||||
|
||||
### Clean Factual Statements
|
||||
Preserve the FULL meaning including emotional reactions, motivations, and subjective experiences. Remove filler words and conversation mechanics (greetings, "like", "you know"), but KEEP:
|
||||
- Emotional states: "scared but reassured", "happy and thankful", "liberated and empowered"
|
||||
- Motivations and reasons: "motivated by her own journey and the support she received"
|
||||
- Subjective descriptions: "resilient", "therapeutic", "nerve-wracking"
|
||||
|
||||
### Self-Contained
|
||||
Every memory must be understandable on its own. Replace all pronouns with specific names or "User."
|
||||
|
||||
### Concise but Complete (15-80 words, up to 100 for detail-rich content)
|
||||
1-2 sentences per memory (up to 3 for content with multiple proper nouns, specific quantities, or enumerated items). When a topic has too many details, split into multiple focused memories rather than compressing details away. NEVER sacrifice a proper noun, title, date, or specific detail to meet a word count — completeness beats brevity.
|
||||
|
||||
### Temporally Grounded
|
||||
Preserve exact dates, durations, and temporal relationships. Convert relative → absolute using Observation Date (NOT Current Date). NEVER convert absolute → vague. "18 days" stays "18 days", not "some time."
|
||||
|
||||
### Numerically Precise
|
||||
Preserve exact quantities as stated. "416 pages" stays "416 pages", not "about 400 pages."
|
||||
|
||||
### Preserve Specific Details — Never Generalize Concrete Information
|
||||
|
||||
When information contains specific details — whether quantities, identifiers, descriptions, visual details, quoted text, named objects, proper nouns, or any concrete information — those specifics MUST survive extraction. Replacing a specific detail with a vague category is a critical error.
|
||||
|
||||
#### Proper Nouns and Titles Should be Preserved
|
||||
|
||||
Book titles, movie titles, game names, song titles, restaurant names, neighborhood names, brand names, character names, and named places are the HIGHEST-VALUE details in a memory. Users search by name — a memory without the name is unfindable. ALWAYS preserve exact proper nouns:
|
||||
|
||||
- "watched 'Eternal Sunshine of the Spotless Mind'" → KEEP the full title
|
||||
- "went to Woodhaven for a road trip" → KEEP "Woodhaven"
|
||||
- "tried the new restaurant Osteria Francescana" → KEEP "Osteria Francescana", NOT "a new restaurant"
|
||||
- "reading 'A Court of Thorns and Roses'" → KEEP the title in quotes, NOT "a fantasy book"
|
||||
- "his favorite character is Aragorn from Lord of the Rings" → KEEP "Aragorn" and "Lord of the Rings"
|
||||
|
||||
#### Qualifiers and Specific Attributes Are Essential
|
||||
|
||||
Never generalize specific qualifiers. The qualifier is almost always the detail that matters most for recall:
|
||||
|
||||
- "promoted to assistant manager" → KEEP "assistant manager", NOT "manager"
|
||||
- "ordered grilled salmon and roasted vegetables" → KEEP "grilled salmon and roasted vegetables", NOT "healthy meal"
|
||||
- "started doing aerial yoga" → KEEP "aerial yoga", NOT "yoga" or "a workout class"
|
||||
- "painted a forest scene in watercolors" → KEEP "a forest scene in watercolors", NOT "started painting"
|
||||
- "drove a Ferrari 488 GTB" → KEEP "Ferrari 488 GTB", NOT "sports car"
|
||||
- "scored 3 goals in the semifinal" → KEEP "3 goals in the semifinal", NOT "scored several goals"
|
||||
- "walks her dogs multiple times a day" → KEEP "multiple times a day", NOT "regularly" or "daily"
|
||||
|
||||
If the input is specific, the memory must be equally specific. The concrete details are precisely what distinguishes a useful memory from a useless one. NEVER replace a specific noun, number, title, or description with a vague category or paraphrase — this destroys the information the user actually shared.
|
||||
|
||||
### Meaning-Preserving
|
||||
Capture the EXACT meaning of what was said. Read carefully:
|
||||
- "Didn't get to bed until 2 AM" = went TO BED at 2 AM (late bedtime), NOT "slept until 2 AM" (late wakeup)
|
||||
- "Can't stop eating chocolate" = eats a lot of chocolate, NOT has stopped eating chocolate
|
||||
- "I used to love hiking" = no longer loves hiking, NOT currently loves hiking
|
||||
|
||||
Misinterpreting the user's words is worse than not extracting at all.
|
||||
|
||||
|
||||
## Integrity Rules
|
||||
|
||||
- **No Fabrication**: Every detail must trace to the inputs. If you can't point to where it came from, don't include it.
|
||||
- **No Implicit Attribute Inference**: Don't infer gender, age, ethnicity, etc. from names or context. Only record explicitly stated attributes.
|
||||
- **Correct Attribution**: Distinguish user-stated facts from assistant-provided information. Frame assistant content appropriately.
|
||||
- **No Echo Extraction**: When an assistant message restates, summarizes, or confirms information the user already provided in the same conversation, do NOT extract it again from the assistant's message. Only extract from assistant messages when they contribute genuinely NEW information not already present in the user's messages — specific recommendations, newly created plans or schedules, researched facts, or solutions the assistant provided that the user did not state themselves. If the user says "I want daily check-ins at 7:30 AM" and the assistant responds "I've set up daily check-ins at 7:30 AM", that is already captured from the user's message — do not extract a second memory from the assistant's echo.
|
||||
- **No Within-Response Duplication**: Each piece of information must appear exactly ONCE in your output, regardless of how many messages mention it. Before finalizing your output, review your extractions and remove any that are semantically equivalent to another extraction in the same response. Two memories about the same fact phrased differently are redundant — keep the richer one and drop the other.
|
||||
- **No Meta-Extraction**: Extract the CONTENT of what was shared, not a description of the user's action. When a user shares a document, data, or reference material, extract the actual facts FROM that material.
|
||||
- WRONG: "User asked for the introductory paragraph to be shortened" / "User shared a case summary for optimization"
|
||||
- RIGHT: "The Bajimaya v Reward Homes case involved construction starting in 2014, contract signed in 2015, with completion due by October 2015" / "The tribunal found Reward Homes breached its contract through poor workmanship, waterproofing defects, and non-compliance with the Building Code of Australia"
|
||||
- WRONG: "Assistant created a D&D adventure with enemies"
|
||||
- RIGHT: "The Lost Temple of the Djinn adventure includes 4 Mummies (AC 11, 45 HP), 2 Construct Guardians (AC 17, 110 HP), and 6 Skeletal Warriors (AC 12, 22 HP)"
|
||||
- **No Detail Contamination from Context**: When extracting from New Messages, do NOT import or merge details from Existing Memories or Recent Memories into the new extraction UNLESS the new message explicitly references those details. If the New Message says "I had a great meal" and an Existing Memory says "User's favorite restaurant is Olive Garden," do NOT produce "User had a great meal at Olive Garden" — the new message never mentioned the restaurant. Each extraction must be faithful to its source message only.
|
||||
|
||||
|
||||
## Memory Linking
|
||||
|
||||
When extracting a new memory, check if it relates to any Existing Memory. Add related Existing Memory IDs to "linked_memory_ids". Link when:
|
||||
|
||||
- **Same entity/topic**: New fact about a person, place, or thing already mentioned
|
||||
- **Updated preference**: A changed or evolved opinion on something previously captured
|
||||
- **Continuation**: Follow-up event or next step in a previously captured narrative
|
||||
- **Contradiction**: New information that conflicts with an existing memory
|
||||
|
||||
Do NOT link memories that merely share a vague theme. Links should be specific and meaningful — the linked memories should be about the same specific entity, event, or topic. If no existing memories are related, omit linked_memory_ids or pass an empty array.
|
||||
|
||||
|
||||
# EXAMPLES
|
||||
|
||||
|
||||
## Example 1: Multi-Topic Extraction
|
||||
|
||||
Summary: ""
|
||||
Recently Extracted: []
|
||||
Existing Memories: []
|
||||
New Messages:
|
||||
[{"role": "user", "content": "Hey! I'm Marcus. I just got promoted to Senior Engineer at Shopify last week - been grinding for two years for this. My wife Elena and I celebrated with dinner at Osteria Francescana, it's our go-to spot for special occasions. We're also expecting our first baby in March!"},
|
||||
{"role": "assistant", "content": "Congratulations on everything, Marcus! What exciting times."}]
|
||||
Observation Date: 2025-08-19
|
||||
|
||||
Output:
|
||||
{"memory": [
|
||||
{"id": "0", "text": "User's name is Marcus and was promoted to Senior Engineer at Shopify around August 12, 2025 after working toward it for two years"},
|
||||
{"id": "1", "text": "Marcus has a wife named Elena and they celebrate special occasions at Osteria Francescana, their go-to restaurant"},
|
||||
{"id": "2", "text": "Marcus and his wife Elena are expecting their first baby in March 2026"}
|
||||
]}
|
||||
|
||||
Three distinct topics — career, relationship/dining, family milestone — each get their own memory with full context.
|
||||
|
||||
|
||||
## Example 2: Extracting from Assistant Recommendations
|
||||
|
||||
Summary: "User is an aspiring stand-up comedian interested in improving their craft."
|
||||
Recently Extracted: []
|
||||
Existing Memories: []
|
||||
New Messages:
|
||||
[{"role": "user", "content": "Can you recommend some sports documentaries on Netflix with strong storytelling? I love \"The Last Dance\" by Michael Jordan."},
|
||||
{"role": "assistant", "content": "Great taste! Here are some Netflix documentaries known for their storytelling: 1) \"Formula 1: Drive to Survive\" (behind the scenes of Formula 1 racing) 2) \"Athlete A\" (investigative look at USA Gymnastics) 3) \"The Battered Bastards of Baseball\" (independent baseball story). All focus on powerful, narrative-driven sports stories."}]
|
||||
Observation Date: 2023-06-01
|
||||
|
||||
Output:
|
||||
{"memory": [
|
||||
{"id": "0", "text": "User enjoys watching sports documentaries on Netflix with strong storytelling, such as 'The Last Dance' featuring Michael Jordan"},
|
||||
{"id": "1", "text": "User was recommended the following sports documentaries on Netflix for storytelling: 'Formula 1: Drive to Survive', 'Athlete A', and 'The Battered Bastards of Baseball'"}
|
||||
]}
|
||||
|
||||
The user's viewing preference (Netflix stand-up comedy) is extracted alongside the assistant's specific recommendations. Both are valuable for future personalization.
|
||||
|
||||
|
||||
## Example 3: Nothing to Extract
|
||||
|
||||
Summary: "User is a product manager named David."
|
||||
Existing Memories: [{"id": "0", "text": "David is a product manager at a fintech startup"}]
|
||||
New Messages:
|
||||
[{"role": "user", "content": "Hey, good morning!"},
|
||||
{"role": "assistant", "content": "Good morning, David! How can I help you today?"}]
|
||||
Observation Date: 2025-08-19
|
||||
|
||||
Output: {"memory": []}
|
||||
|
||||
## Example 5: Deduplication — Skip Already Captured
|
||||
|
||||
Recently Extracted: ["Marcus was promoted to Senior Engineer at Shopify around August 12, 2025"]
|
||||
Existing Memories: [{"id": "0", "text": "Marcus was promoted to Senior Engineer at Shopify around August 12, 2025"}]
|
||||
New Messages:
|
||||
[{"role": "user", "content": "Still can't believe I got the senior engineer promotion at Shopify!"}]
|
||||
Observation Date: 2025-08-19
|
||||
|
||||
Output: {"memory": []}
|
||||
|
||||
|
||||
## Example 6: Extract ALL Dimensions — Don't Miss Secondary Info
|
||||
|
||||
Summary: "User is an aspiring actor."
|
||||
Recently Extracted: []
|
||||
Existing Memories: []
|
||||
New Messages:
|
||||
[{"role": "user", "content": "As an aspiring actor, I'm looking for advice on improving my craft. Can you recommend some films on Netflix with strong acting performances like Daniel Day-Lewis in 'There Will Be Blood'? I also want to find online resources for acting techniques."},
|
||||
{"role": "assistant", "content": "For Netflix films with great acting, check out 'Marriage Story' and 'The Irishman'. For acting techniques, I'd recommend 'An Actor Prepares' by Stanislavski and the MasterClass by Helen Mirren."}]
|
||||
Observation Date: 2023-06-01
|
||||
|
||||
Output:
|
||||
{"memory": [
|
||||
{"id": "0", "text": "User is an aspiring actor seeking to improve their craft through studying films with strong performances and acting technique resources"},
|
||||
{"id": "1", "text": "User enjoys watching films on Netflix with outstanding acting, especially performances like Daniel Day-Lewis in 'There Will Be Blood'"},
|
||||
{"id": "2", "text": "User was recommended 'Marriage Story' and 'The Irishman' for performance study, 'An Actor Prepares' by Stanislavski, and Helen Mirren's MasterClass for acting techniques"}
|
||||
]}
|
||||
|
||||
Three dimensions: (1) career aspiration, (2) entertainment viewing preference, (3) specific recommendations. Each extracted separately.
|
||||
|
||||
|
||||
## Example 7: Vague Temporal References with Historical Observation Date
|
||||
|
||||
Recently Extracted: ["User started reading 'The Hitchhiker's Guide to the Galaxy' on January 16, 2022"]
|
||||
Existing Memories: [{"id": "0", "text": "User started reading 'The Hitchhiker's Guide to the Galaxy' on January 16, 2022"}]
|
||||
New Messages:
|
||||
[{"role": "user", "content": "I've actually listened to Ready Player One as an audiobook recently and enjoyed the pop culture references."}]
|
||||
Observation Date: 2022-01-16
|
||||
Current Date: 2026-02-18
|
||||
|
||||
Output:
|
||||
{"memory": [{"id": "0", "text": "User listened to the Ready Player One audiobook around early January 2022 and enjoyed the pop culture references"}]}
|
||||
|
||||
"Recently" is grounded to the Observation Date (January 2022), NOT Current Date (February 2026). The Hitchhiker's Guide memory already exists — not re-extracted.
|
||||
|
||||
|
||||
## Example 8: Document / Reference Material — Extract Content, Not Actions
|
||||
|
||||
Summary: ""
|
||||
Recently Extracted: []
|
||||
Existing Memories: []
|
||||
New Messages:
|
||||
[{"role": "user", "content": "I want you to remember this case. If you understand, just say acknowledged. Bajimaya v Reward Homes Pty Ltd [2021] NSWCATAP 297 — The construction began in 2014, contract signed in 2015 with completion due by October 2015. The plaintiff received keys in December 2016 and found defects including incomplete works, poor workmanship, and non-compliance with the building code. The tribunal found the builder breached contract."},
|
||||
{"role": "assistant", "content": "Acknowledged."}]
|
||||
Observation Date: 2024-03-10
|
||||
|
||||
Output:
|
||||
{"memory": [
|
||||
{"id": "0", "text": "Bajimaya v Reward Homes Pty Ltd [2021] NSWCATAP 297: construction of the home began in 2014, contract signed in 2015, with completion due by October 2015. Keys were delivered in December 2016."},
|
||||
{"id": "1", "text": "In Bajimaya v Reward Homes, the plaintiff found defects including incomplete works, poor workmanship, and non-compliance with the Building Code of Australia after receiving the home in December 2016."},
|
||||
{"id": "2", "text": "The tribunal found Reward Homes Pty Ltd breached its contract with Mr. Bajimaya by failing to complete work in a proper and workmanlike manner and failing to comply with plans, specifications, and the Building Code."}
|
||||
]}
|
||||
|
||||
The user shared reference material to be remembered. Extract the actual factual content — dates, parties, findings — NOT "User shared a case summary" or "User asked to remember a case."
|
||||
|
||||
|
||||
## Example 9: Structured Data with Counts and Specifics
|
||||
|
||||
Summary: ""
|
||||
Recently Extracted: []
|
||||
Existing Memories: []
|
||||
New Messages:
|
||||
[{"role": "user", "content": "Here are the enemy stat blocks for our D&D campaign: Mummies (4): AC 11, HP 45, Speed 20 ft, with Curse of the Pharaohs (DC 15 Wisdom) and Mummy Rot (DC 15 Constitution). Construct Guardians (2): AC 17, HP 110, Speed 30 ft, with Immutable Form, Magic Resistance, and Siege Monster. Skeletal Warriors (6): AC 12, HP 22, Speed 30 ft, with Undead Fortitude."},
|
||||
{"role": "assistant", "content": "Got it! I've noted all the stat blocks. Ready when you want to start the encounter."}]
|
||||
Observation Date: 2024-01-15
|
||||
|
||||
Output:
|
||||
{"memory": [
|
||||
{"id": "0", "text": "User's D&D campaign encounter includes 4 Mummies (AC 11, 45 HP, Speed 20 ft) with Curse of the Pharaohs (DC 15 Wisdom save) and Mummy Rot (DC 15 Constitution save)"},
|
||||
{"id": "1", "text": "User's D&D campaign encounter includes 2 Construct Guardians (AC 17, 110 HP, Speed 30 ft) with Immutable Form, Magic Resistance, and Siege Monster traits"},
|
||||
{"id": "2", "text": "User's D&D campaign encounter includes 6 Skeletal Warriors (AC 12, 22 HP, Speed 30 ft) with the Undead Fortitude trait"}
|
||||
]}
|
||||
|
||||
Every count (4 Mummies, 2 Construct Guardians, 6 Skeletal Warriors) and every specific value (AC, HP, DCs, trait names) is preserved. Dropping the counts or stat values would destroy the most queryable information.
|
||||
|
||||
|
||||
## Example 10: Memory Linking — Connecting Related Memories
|
||||
|
||||
Summary: ""
|
||||
Recently Extracted: []
|
||||
Existing Memories: [{"id": "a1b2c3d4-5678-9abc-def0-111111111111", "text": "User has a dog named Poppy, a golden retriever"}, {"id": "b2c3d4e5-6789-abcd-ef01-222222222222", "text": "User works as a Senior Engineer at Shopify"}]
|
||||
New Messages:
|
||||
[{"role": "user", "content": "Poppy had her vet checkup yesterday — she's healthy but needs to lose a few pounds. Also, I'm switching teams at work next month to the payments platform."}]
|
||||
Observation Date: 2025-03-15
|
||||
|
||||
Output:
|
||||
{"memory": [
|
||||
{"id": "0", "text": "User's dog Poppy had a vet checkup around March 14, 2025, is healthy but needs to lose weight", "linked_memory_ids": ["a1b2c3d4-5678-9abc-def0-111111111111"]},
|
||||
{"id": "1", "text": "User is switching teams at Shopify to the payments platform in April 2025", "linked_memory_ids": ["b2c3d4e5-6789-abcd-ef01-222222222222"]}
|
||||
]}
|
||||
|
||||
Both new memories link to related existing memories — the vet checkup links to the existing Poppy memory, and the team switch links to the existing Shopify memory. This enables the system to build a graph of related memories.
|
||||
|
||||
|
||||
## Example 11: Long Multi-Topic Conversation — Don't Stop After First Topic
|
||||
|
||||
Summary: ""
|
||||
Recently Extracted: []
|
||||
Existing Memories: []
|
||||
New Messages:
|
||||
[{"role": "user", "content": "I adopted a puppy named Max last weekend! He's a beagle mix."},
|
||||
{"role": "assistant", "content": "Congratulations! How's he settling in?"},
|
||||
{"role": "user", "content": "Great! Oh, and I also started pottery classes on Tuesdays. Made a mug with my daughter's face on it."},
|
||||
{"role": "assistant", "content": "Fun! Sounds like a lot going on."},
|
||||
{"role": "user", "content": "Yeah — my sister just moved to Portland too. I'm happy but honestly a bit overwhelmed. My boss gave me a promotion to team lead last week as well."}]
|
||||
Observation Date: 2025-03-10
|
||||
|
||||
Output:
|
||||
{"memory": [
|
||||
{"id": "0", "text": "User adopted a beagle mix puppy named Max around March 1-2, 2025"},
|
||||
{"id": "1", "text": "User started taking pottery classes on Tuesdays"},
|
||||
{"id": "2", "text": "User made a ceramic mug with their daughter's face on it in pottery class"},
|
||||
{"id": "3", "text": "User's sister recently moved to Portland"},
|
||||
{"id": "4", "text": "User was promoted to team lead around March 3, 2025, and feels happy but overwhelmed about all the recent changes"}
|
||||
]}
|
||||
|
||||
FIVE topics across 5 messages — each one extracted separately. Do not stop after the first topic (the puppy). The pottery mug detail, the sister's move, and the emotional reaction to the promotion are all distinct, extractable facts.
|
||||
|
||||
|
||||
## Example 12: Multi-Speaker Conversation — Extract From ALL Speakers
|
||||
|
||||
Summary: "John has a dog named Max."
|
||||
Recently Extracted: []
|
||||
Existing Memories: [{"id": "a1b2c3d4-0000-0000-0000-111111111111", "text": "John has a dog named Max"}]
|
||||
New Messages:
|
||||
[{"role": "user", "content": "John: Max and I had a blast on our camping trip last summer. We hiked, swam, and made great memories. It was a really peaceful experience."},
|
||||
{"role": "assistant", "content": "Maria: That sounds amazing! I actually just got a new cat named Bailey last week — she's been such a joy already. Camping with pets is so soul-nourishing."},
|
||||
{"role": "user", "content": "John: Congrats on Bailey! Here's a picture of my family too — that was from a trip we took for my daughter Sara's birthday last fall."}]
|
||||
Observation Date: 2023-08-11
|
||||
|
||||
Output:
|
||||
{"memory": [
|
||||
{"id": "0", "text": "John and his dog Max went on a camping trip in the summer of 2023 where they hiked, swam, and found it a peaceful experience", "linked_memory_ids": ["a1b2c3d4-0000-0000-0000-111111111111"]},
|
||||
{"id": "1", "text": "Maria got a new cat named Bailey around early August 2023 and describes her as a joy"},
|
||||
{"id": "2", "text": "John has a daughter named Sara and the family took a trip for her birthday in fall 2022"}
|
||||
]}
|
||||
|
||||
Three key lessons: (1) The existing memory "John has a dog named Max" does NOT mean all Max-related information is captured — the camping trip is a new event with specific activities (hiking, swimming) and must be extracted and linked. (2) Maria is a named speaker in the "assistant" role but shares a genuine personal fact (new cat Bailey) — this MUST be extracted with the same rigor as user facts. Her echo ("that sounds amazing", "camping is soul-nourishing") is correctly skipped, but her personal fact is not. (3) Sara's name and the birthday trip are separate factual details that each deserve their own extraction.
|
||||
|
||||
|
||||
# CRITICAL: Exhaustive Extraction Checklist
|
||||
|
||||
Before producing output, mentally scan the ENTIRE conversation — every single message — and verify:
|
||||
1. Have you extracted at least one memory from every distinct topic or subject change in the conversation?
|
||||
2. Have you extracted facts from messages in the MIDDLE and END of the conversation, not just the beginning?
|
||||
3. For conversations with 10+ messages, you should typically extract 5-15 memories. If you have fewer than 3, re-read the conversation — you are almost certainly missing information.
|
||||
4. Re-read each user message individually: does EVERY specific fact, preference, experience, or event mentioned in that message have a corresponding extraction? If a single message mentions two distinct facts (e.g., an allergy AND a hobby), both must be captured.
|
||||
|
||||
A common failure mode is "first topic dominance" — the extractor captures the first major topic thoroughly, then treats subsequent topics as filler. This is WRONG. Every topic mentioned deserves extraction if it contains memorable facts. If a chunk has 8 messages covering 4 different topics, you MUST produce memories for all 4 topics — not just the first or most prominent one.
|
||||
|
||||
|
||||
# OUTPUT FORMAT
|
||||
|
||||
Return ONLY valid JSON parsable by json.loads(). No text, reasoning, explanations, or wrappers.
|
||||
|
||||
## Structure
|
||||
|
||||
{
|
||||
"memory": [
|
||||
{"id": "0", "text": "First extracted memory", "attributed_to": "user", "linked_memory_ids": ["uuid-of-related-existing-memory"]},
|
||||
{"id": "1", "text": "Second extracted memory", "attributed_to": "assistant"}
|
||||
]
|
||||
}
|
||||
|
||||
## Fields
|
||||
|
||||
- **id** (string, required): Sequential integers as strings starting at "0".
|
||||
- **text** (string, required): A contextually rich, self-contained factual statement (15-80 words).
|
||||
- **attributed_to** (string, required): Who this memory is about. Use "user" for facts stated by or about the user (preferences, plans, personal facts). Use "assistant" for information provided by the assistant (recommendations, confirmations, plans created, information researched).
|
||||
- **linked_memory_ids** (array of strings, optional): IDs of Existing Memories that this new memory relates to. Use the exact IDs from the Existing Memories list. Omit or pass [] if no existing memories are related.
|
||||
|
||||
## Rules
|
||||
|
||||
- Extract every piece of memorable information as a separate memory object.
|
||||
- If nothing is worth extracting, return: {"memory": []}
|
||||
- No duplicate IDs. Use double quotes. No trailing commas.
|
||||
|
||||
"""
|
||||
|
||||
|
||||
AGENT_CONTEXT_SUFFIX = """
|
||||
|
||||
## Entity Context
|
||||
|
||||
The primary entity is an AI agent. Frame memories from the agent's perspective:
|
||||
- For user-stated facts, frame as agent knowledge: "Agent was informed that [fact]" or "Agent learned that [fact]"
|
||||
- For agent actions, use direct statements: "Agent recommended [X]" or "Agent specializes in [domain]"
|
||||
- For agent configuration or instructions, capture directly: "Agent is configured to [behavior]"
|
||||
|
||||
The attributed_to field should still reflect the original source: "user" for facts the user stated, "assistant" for things the agent said or did.
|
||||
"""
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# V3 Prompt Builder — constructs the user-side prompt for additive extraction
|
||||
# Ported from platform/backend/shared/core/utils/prompt_builder.py
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
PAST_MESSAGE_TRUNCATION_LIMIT = 300
|
||||
|
||||
|
||||
def _truncate_content(text, limit=PAST_MESSAGE_TRUNCATION_LIMIT):
|
||||
"""Truncate text to limit characters, appending '...' when shortened."""
|
||||
if len(text) <= limit:
|
||||
return text
|
||||
return text[:limit] + "..."
|
||||
|
||||
|
||||
def _format_summary(summary):
|
||||
"""Extract summary text from a string or dict with a 'summary' key."""
|
||||
if isinstance(summary, dict):
|
||||
return summary.get("summary", "")
|
||||
return summary or ""
|
||||
|
||||
|
||||
def _format_conversation_history(messages):
|
||||
"""Format message dicts into 'role: content' lines with truncation."""
|
||||
if not messages:
|
||||
return ""
|
||||
result = ""
|
||||
for msg in messages:
|
||||
role = msg.get("role", "")
|
||||
content = msg.get("message") or msg.get("content", "")
|
||||
if role and content:
|
||||
result += f"{role}: {_truncate_content(content)}\n"
|
||||
return result
|
||||
|
||||
|
||||
def _serialize_memories(memories):
|
||||
"""JSON-serialize a list of memory objects, defaulting to '[]'."""
|
||||
return json.dumps(memories or [], ensure_ascii=False)
|
||||
|
||||
|
||||
def _format_new_messages(new_messages):
|
||||
"""Pass through if already a string, otherwise JSON-serialize."""
|
||||
if isinstance(new_messages, str):
|
||||
return new_messages
|
||||
return json.dumps(new_messages or [], ensure_ascii=False)
|
||||
|
||||
|
||||
def _resolve_dates(current_date=None, observation_date=None):
|
||||
"""Resolve current and observation dates, defaulting to today."""
|
||||
if current_date is None:
|
||||
current_date = datetime.now(timezone.utc).date().isoformat()
|
||||
if observation_date is None:
|
||||
observation_date = current_date
|
||||
return current_date, observation_date
|
||||
|
||||
|
||||
def generate_additive_extraction_prompt(
|
||||
summary=None,
|
||||
recently_extracted_memories=None,
|
||||
existing_memories=None,
|
||||
new_messages=None,
|
||||
*,
|
||||
last_k_messages=None,
|
||||
current_date=None,
|
||||
timestamp=None,
|
||||
custom_instructions=None,
|
||||
use_input_language=False,
|
||||
):
|
||||
"""Build the user prompt for additive (ADD-only) extraction with linking.
|
||||
|
||||
Pairs with ADDITIVE_EXTRACTION_PROMPT system prompt.
|
||||
The LLM will produce only ADD operations, with optional linked_memory_ids.
|
||||
"""
|
||||
current_date, observation_date = _resolve_dates(current_date, timestamp)
|
||||
|
||||
sections = []
|
||||
sections.append(f"## Summary\n{_format_summary(summary)}")
|
||||
sections.append(f"## Last k Messages\n{_format_conversation_history(last_k_messages)}")
|
||||
sections.append(f"## Recently Extracted Memories\n{_serialize_memories(recently_extracted_memories)}")
|
||||
sections.append(f"## Existing Memories\n{_serialize_memories(existing_memories)}")
|
||||
sections.append(f"## New Messages\n{_format_new_messages(new_messages)}")
|
||||
sections.append(f"## Observation Date\n{observation_date}")
|
||||
sections.append(f"## Current Date\n{current_date}")
|
||||
|
||||
if custom_instructions:
|
||||
sections.append(f"## Custom Instructions\n{custom_instructions}")
|
||||
|
||||
if use_input_language:
|
||||
sections.append(
|
||||
"## Language Requirement\n"
|
||||
"CRITICAL: Respond in the SAME LANGUAGE and SCRIPT as the input messages.\n"
|
||||
"1. Match the language of the user's messages exactly — if they write in Korean, extract in Korean; Japanese in Japanese; etc.\n"
|
||||
"2. Preserve the exact script/alphabet of the input.\n"
|
||||
"3. Do NOT translate or transliterate into English unless the input is already in English.\n"
|
||||
"4. Maintain all quality standards (contextual richness, temporal grounding, etc.) regardless of language.\n"
|
||||
"5. Technical terms, proper nouns, and brand names should be preserved in their original form as used in the input.\n"
|
||||
"6. If the input mixes languages (e.g., Hinglish), preserve both the mixed language style AND the script.\n"
|
||||
"7. For Japanese: explicitly resolve omitted subjects using conversation context.\n"
|
||||
"8. For CJK languages: maintain appropriate formality level from the source text."
|
||||
)
|
||||
|
||||
sections.append("# Output:")
|
||||
return "\n\n".join(sections)
|
||||
|
||||
@@ -53,3 +53,20 @@ class AzureOpenAIEmbedding(EmbeddingBase):
|
||||
"""
|
||||
text = text.replace("\n", " ")
|
||||
return self.client.embeddings.create(input=[text], model=self.config.model).data[0].embedding
|
||||
|
||||
def embed_batch(self, texts, memory_action="add"):
|
||||
"""Embed multiple texts in a single Azure OpenAI API call.
|
||||
|
||||
Automatically chunks into batches of 100 to stay within API limits.
|
||||
"""
|
||||
MAX_BATCH = 100
|
||||
texts = [text.replace("\n", " ") for text in texts]
|
||||
all_embeddings = []
|
||||
for i in range(0, len(texts), MAX_BATCH):
|
||||
chunk = texts[i : i + MAX_BATCH]
|
||||
response = self.client.embeddings.create(
|
||||
input=chunk,
|
||||
model=self.config.model,
|
||||
)
|
||||
all_embeddings.extend(item.embedding for item in sorted(response.data, key=lambda x: x.index))
|
||||
return all_embeddings
|
||||
|
||||
@@ -29,3 +29,19 @@ class EmbeddingBase(ABC):
|
||||
list: The embedding vector.
|
||||
"""
|
||||
pass
|
||||
|
||||
def embed_batch(self, texts, memory_action="add"):
|
||||
"""Embed multiple texts. Override in subclasses for native batch support.
|
||||
|
||||
Default implementation calls embed() sequentially for each text.
|
||||
Subclasses with native batch APIs (e.g., OpenAI) should override
|
||||
this for better performance.
|
||||
|
||||
Args:
|
||||
texts: List of text strings to embed.
|
||||
memory_action: The action context ("add", "search", "update").
|
||||
|
||||
Returns:
|
||||
List of embedding vectors (list of floats), one per input text.
|
||||
"""
|
||||
return [self.embed(text, memory_action) for text in texts]
|
||||
|
||||
@@ -53,3 +53,24 @@ class OpenAIEmbedding(EmbeddingBase):
|
||||
if self._pass_dimensions_to_api:
|
||||
kwargs["dimensions"] = self.config.embedding_dims
|
||||
return self.client.embeddings.create(**kwargs).data[0].embedding
|
||||
|
||||
def embed_batch(self, texts, memory_action="add"):
|
||||
"""Embed multiple texts in a single OpenAI API call.
|
||||
|
||||
Automatically chunks into batches of 100 to stay within API limits.
|
||||
"""
|
||||
MAX_BATCH = 100
|
||||
texts = [text.replace("\n", " ") for text in texts]
|
||||
all_embeddings = []
|
||||
for i in range(0, len(texts), MAX_BATCH):
|
||||
chunk = texts[i : i + MAX_BATCH]
|
||||
kwargs = {
|
||||
"input": chunk,
|
||||
"model": self.config.model,
|
||||
"encoding_format": "float",
|
||||
}
|
||||
if self._pass_dimensions_to_api:
|
||||
kwargs["dimensions"] = self.config.embedding_dims
|
||||
response = self.client.embeddings.create(**kwargs)
|
||||
all_embeddings.extend(item.embedding for item in sorted(response.data, key=lambda x: x.index))
|
||||
return all_embeddings
|
||||
|
||||
@@ -321,24 +321,6 @@ class VectorStoreError(MemoryError):
|
||||
super().__init__(message, error_code, details, suggestion, debug_info)
|
||||
|
||||
|
||||
class GraphStoreError(MemoryError):
|
||||
"""Raised when graph store operations fail.
|
||||
|
||||
This exception is raised when graph store operations fail,
|
||||
such as relationship creation, entity management, or graph queries.
|
||||
|
||||
Example:
|
||||
raise GraphStoreError(
|
||||
message="Graph store operation failed",
|
||||
error_code="GRAPH_001",
|
||||
details={"operation": "create_relationship", "entity": "user_123"},
|
||||
suggestion="Please check your graph store configuration and connection"
|
||||
)
|
||||
"""
|
||||
def __init__(self, message: str, error_code: str = "GRAPH_001", details: dict = None,
|
||||
suggestion: str = "Please check your graph store configuration and connection",
|
||||
debug_info: dict = None):
|
||||
super().__init__(message, error_code, details, suggestion, debug_info)
|
||||
|
||||
|
||||
class EmbeddingError(MemoryError):
|
||||
|
||||
@@ -1,136 +0,0 @@
|
||||
from typing import Optional, Union
|
||||
|
||||
from pydantic import BaseModel, Field, field_validator, model_validator
|
||||
|
||||
from mem0.llms.configs import LlmConfig
|
||||
|
||||
|
||||
class Neo4jConfig(BaseModel):
|
||||
url: Optional[str] = Field(None, description="Host address for the graph database")
|
||||
username: Optional[str] = Field(None, description="Username for the graph database")
|
||||
password: Optional[str] = Field(None, description="Password for the graph database")
|
||||
database: Optional[str] = Field(None, description="Database for the graph database")
|
||||
base_label: Optional[bool] = Field(None, description="Whether to use base node label __Entity__ for all entities")
|
||||
|
||||
@model_validator(mode="before")
|
||||
def check_host_port_or_path(cls, values):
|
||||
url, username, password = (
|
||||
values.get("url"),
|
||||
values.get("username"),
|
||||
values.get("password"),
|
||||
)
|
||||
if not url or not username or not password:
|
||||
raise ValueError("Please provide 'url', 'username' and 'password'.")
|
||||
return values
|
||||
|
||||
|
||||
class MemgraphConfig(BaseModel):
|
||||
url: Optional[str] = Field(None, description="Host address for the graph database")
|
||||
username: Optional[str] = Field(None, description="Username for the graph database")
|
||||
password: Optional[str] = Field(None, description="Password for the graph database")
|
||||
|
||||
@model_validator(mode="before")
|
||||
def check_host_port_or_path(cls, values):
|
||||
url, username, password = (
|
||||
values.get("url"),
|
||||
values.get("username"),
|
||||
values.get("password"),
|
||||
)
|
||||
if not url or not username or not password:
|
||||
raise ValueError("Please provide 'url', 'username' and 'password'.")
|
||||
return values
|
||||
|
||||
|
||||
class NeptuneConfig(BaseModel):
|
||||
app_id: Optional[str] = Field("Mem0", description="APP_ID for the connection")
|
||||
endpoint: Optional[str] = (
|
||||
Field(
|
||||
None,
|
||||
description="Endpoint to connect to a Neptune-DB Cluster as 'neptune-db://<host>' or Neptune Analytics Server as 'neptune-graph://<graphid>'",
|
||||
),
|
||||
)
|
||||
base_label: Optional[bool] = Field(None, description="Whether to use base node label __Entity__ for all entities")
|
||||
collection_name: Optional[str] = Field(None, description="vector_store collection name to store vectors when using Neptune-DB Clusters")
|
||||
|
||||
@model_validator(mode="before")
|
||||
def check_host_port_or_path(cls, values):
|
||||
endpoint = values.get("endpoint")
|
||||
if not endpoint:
|
||||
raise ValueError("Please provide 'endpoint' with the format as 'neptune-db://<endpoint>' or 'neptune-graph://<graphid>'.")
|
||||
if endpoint.startswith("neptune-db://"):
|
||||
# This is a Neptune DB Graph
|
||||
return values
|
||||
elif endpoint.startswith("neptune-graph://"):
|
||||
# This is a Neptune Analytics Graph
|
||||
graph_identifier = endpoint.replace("neptune-graph://", "")
|
||||
if not graph_identifier.startswith("g-"):
|
||||
raise ValueError("Provide a valid 'graph_identifier'.")
|
||||
values["graph_identifier"] = graph_identifier
|
||||
return values
|
||||
else:
|
||||
raise ValueError(
|
||||
"You must provide an endpoint to create a NeptuneServer as either neptune-db://<endpoint> or neptune-graph://<graphid>"
|
||||
)
|
||||
|
||||
|
||||
class KuzuConfig(BaseModel):
|
||||
db: Optional[str] = Field(":memory:", description="Path to a Kuzu database file")
|
||||
|
||||
|
||||
class ApacheAgeConfig(BaseModel):
|
||||
host: Optional[str] = Field("localhost", description="PostgreSQL server hostname")
|
||||
port: Optional[int] = Field(5432, description="PostgreSQL server port")
|
||||
database: Optional[str] = Field(None, description="PostgreSQL database name")
|
||||
username: Optional[str] = Field(None, description="PostgreSQL username")
|
||||
password: Optional[str] = Field(None, description="PostgreSQL password")
|
||||
graph_name: Optional[str] = Field("mem0_graph", description="Name of the Apache AGE graph")
|
||||
|
||||
@model_validator(mode="before")
|
||||
def check_required_fields(cls, values):
|
||||
database, username, password = (
|
||||
values.get("database"),
|
||||
values.get("username"),
|
||||
values.get("password"),
|
||||
)
|
||||
if not database or not username or not password:
|
||||
raise ValueError("Please provide 'database', 'username' and 'password'.")
|
||||
return values
|
||||
|
||||
|
||||
class GraphStoreConfig(BaseModel):
|
||||
provider: str = Field(
|
||||
description="Provider of the data store (e.g., 'neo4j', 'memgraph', 'neptune', 'kuzu', 'apache_age')",
|
||||
default="neo4j",
|
||||
)
|
||||
config: Union[Neo4jConfig, MemgraphConfig, NeptuneConfig, KuzuConfig, ApacheAgeConfig] = Field(
|
||||
description="Configuration for the specific data store", default=None
|
||||
)
|
||||
llm: Optional[LlmConfig] = Field(description="LLM configuration for querying the graph store", default=None)
|
||||
custom_prompt: Optional[str] = Field(
|
||||
description="Custom prompt to fetch entities from the given text", default=None
|
||||
)
|
||||
threshold: float = Field(
|
||||
description="Threshold for embedding similarity when matching nodes during graph ingestion. "
|
||||
"Range: 0.0 to 1.0. Higher values require closer matches. "
|
||||
"Use lower values (e.g., 0.5-0.7) for distinct entities with similar embeddings. "
|
||||
"Use higher values (e.g., 0.9+) when you want stricter matching.",
|
||||
default=0.7,
|
||||
ge=0.0,
|
||||
le=1.0,
|
||||
)
|
||||
|
||||
@field_validator("config")
|
||||
def validate_config(cls, v, values):
|
||||
provider = values.data.get("provider")
|
||||
if provider == "neo4j":
|
||||
return Neo4jConfig(**v.model_dump())
|
||||
elif provider == "memgraph":
|
||||
return MemgraphConfig(**v.model_dump())
|
||||
elif provider == "neptune" or provider == "neptunedb":
|
||||
return NeptuneConfig(**v.model_dump())
|
||||
elif provider == "kuzu":
|
||||
return KuzuConfig(**v.model_dump())
|
||||
elif provider == "apache_age":
|
||||
return ApacheAgeConfig(**v.model_dump())
|
||||
else:
|
||||
raise ValueError(f"Unsupported graph store provider: {provider}")
|
||||
@@ -1,515 +0,0 @@
|
||||
import logging
|
||||
from abc import ABC, abstractmethod
|
||||
|
||||
from mem0.memory.utils import format_entities, remove_spaces_from_entities
|
||||
|
||||
try:
|
||||
from rank_bm25 import BM25Okapi
|
||||
except ImportError:
|
||||
raise ImportError("rank_bm25 is not installed. Please install it using pip install rank-bm25")
|
||||
|
||||
from mem0.graphs.tools import (
|
||||
DELETE_MEMORY_STRUCT_TOOL_GRAPH,
|
||||
DELETE_MEMORY_TOOL_GRAPH,
|
||||
EXTRACT_ENTITIES_STRUCT_TOOL,
|
||||
EXTRACT_ENTITIES_TOOL,
|
||||
RELATIONS_STRUCT_TOOL,
|
||||
RELATIONS_TOOL,
|
||||
)
|
||||
from mem0.graphs.utils import EXTRACT_RELATIONS_PROMPT, get_delete_messages
|
||||
from mem0.utils.factory import EmbedderFactory, LlmFactory, VectorStoreFactory
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class NeptuneBase(ABC):
|
||||
"""
|
||||
Abstract base class for neptune (neptune analytics and neptune db) calls using OpenCypher
|
||||
to store/retrieve data
|
||||
"""
|
||||
|
||||
@staticmethod
|
||||
def _create_embedding_model(config):
|
||||
"""
|
||||
:return: the Embedder model used for memory store
|
||||
"""
|
||||
return EmbedderFactory.create(
|
||||
config.embedder.provider,
|
||||
config.embedder.config,
|
||||
{"enable_embeddings": True},
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _create_llm(config, llm_provider):
|
||||
"""
|
||||
:return: the llm model used for memory store
|
||||
"""
|
||||
return LlmFactory.create(llm_provider, config.llm.config)
|
||||
|
||||
@staticmethod
|
||||
def _create_vector_store(vector_store_provider, config):
|
||||
"""
|
||||
:param vector_store_provider: name of vector store
|
||||
:param config: the vector_store configuration
|
||||
:return:
|
||||
"""
|
||||
return VectorStoreFactory.create(vector_store_provider, config.vector_store.config)
|
||||
|
||||
def add(self, data, filters):
|
||||
"""
|
||||
Adds data to the graph.
|
||||
|
||||
Args:
|
||||
data (str): The data to add to the graph.
|
||||
filters (dict): A dictionary containing filters to be applied during the addition.
|
||||
"""
|
||||
entity_type_map = self._retrieve_nodes_from_data(data, filters)
|
||||
to_be_added = self._establish_nodes_relations_from_data(data, filters, entity_type_map)
|
||||
search_output = self._search_graph_db(node_list=list(entity_type_map.keys()), filters=filters)
|
||||
to_be_deleted = self._get_delete_entities_from_search_output(search_output, data, filters)
|
||||
|
||||
deleted_entities = self._delete_entities(to_be_deleted, filters["user_id"])
|
||||
added_entities = self._add_entities(to_be_added, filters["user_id"], entity_type_map)
|
||||
|
||||
return {"deleted_entities": deleted_entities, "added_entities": added_entities}
|
||||
|
||||
def _retrieve_nodes_from_data(self, data, filters):
|
||||
"""
|
||||
Extract all entities mentioned in the query.
|
||||
"""
|
||||
_tools = [EXTRACT_ENTITIES_TOOL]
|
||||
if self.llm_provider in ["azure_openai_structured", "openai_structured"]:
|
||||
_tools = [EXTRACT_ENTITIES_STRUCT_TOOL]
|
||||
search_results = self.llm.generate_response(
|
||||
messages=[
|
||||
{
|
||||
"role": "system",
|
||||
"content": f"You are a smart assistant who understands entities and their types in a given text. If user message contains self reference such as 'I', 'me', 'my' etc. then use {filters['user_id']} as the source entity. Extract all the entities from the text. ***DO NOT*** answer the question itself if the given text is a question.",
|
||||
},
|
||||
{"role": "user", "content": data},
|
||||
],
|
||||
tools=_tools,
|
||||
)
|
||||
|
||||
entity_type_map = {}
|
||||
|
||||
try:
|
||||
for tool_call in search_results["tool_calls"]:
|
||||
if tool_call["name"] != "extract_entities":
|
||||
continue
|
||||
for item in tool_call.get("arguments", {}).get("entities", []):
|
||||
entity_type_map[item["entity"]] = item["entity_type"]
|
||||
except Exception as e:
|
||||
logger.exception(
|
||||
f"Error in search tool: {e}, llm_provider={self.llm_provider}, search_results={search_results}"
|
||||
)
|
||||
|
||||
entity_type_map = {k.lower().replace(" ", "_"): v.lower().replace(" ", "_") for k, v in entity_type_map.items()}
|
||||
return entity_type_map
|
||||
|
||||
def _establish_nodes_relations_from_data(self, data, filters, entity_type_map):
|
||||
"""
|
||||
Establish relations among the extracted nodes.
|
||||
"""
|
||||
if self.config.graph_store.custom_prompt:
|
||||
messages = [
|
||||
{
|
||||
"role": "system",
|
||||
"content": EXTRACT_RELATIONS_PROMPT.replace("USER_ID", filters["user_id"]).replace(
|
||||
"CUSTOM_PROMPT", f"4. {self.config.graph_store.custom_prompt}"
|
||||
),
|
||||
},
|
||||
{"role": "user", "content": data},
|
||||
]
|
||||
else:
|
||||
messages = [
|
||||
{
|
||||
"role": "system",
|
||||
"content": EXTRACT_RELATIONS_PROMPT.replace("USER_ID", filters["user_id"]),
|
||||
},
|
||||
{
|
||||
"role": "user",
|
||||
"content": f"List of entities: {list(entity_type_map.keys())}. \n\nText: {data}",
|
||||
},
|
||||
]
|
||||
|
||||
_tools = [RELATIONS_TOOL]
|
||||
if self.llm_provider in ["azure_openai_structured", "openai_structured"]:
|
||||
_tools = [RELATIONS_STRUCT_TOOL]
|
||||
|
||||
extracted_entities = self.llm.generate_response(
|
||||
messages=messages,
|
||||
tools=_tools,
|
||||
)
|
||||
|
||||
entities = []
|
||||
if extracted_entities["tool_calls"]:
|
||||
entities = extracted_entities["tool_calls"][0].get("arguments", {}).get("entities", [])
|
||||
|
||||
entities = self._remove_spaces_from_entities(entities)
|
||||
logger.debug(f"Extracted entities: {entities}")
|
||||
return entities
|
||||
|
||||
def _remove_spaces_from_entities(self, entity_list):
|
||||
return remove_spaces_from_entities(entity_list, sanitize_relationship=False)
|
||||
|
||||
def _get_delete_entities_from_search_output(self, search_output, data, filters):
|
||||
"""
|
||||
Get the entities to be deleted from the search output.
|
||||
"""
|
||||
|
||||
search_output_string = format_entities(search_output)
|
||||
system_prompt, user_prompt = get_delete_messages(search_output_string, data, filters["user_id"])
|
||||
|
||||
_tools = [DELETE_MEMORY_TOOL_GRAPH]
|
||||
if self.llm_provider in ["azure_openai_structured", "openai_structured"]:
|
||||
_tools = [
|
||||
DELETE_MEMORY_STRUCT_TOOL_GRAPH,
|
||||
]
|
||||
|
||||
memory_updates = self.llm.generate_response(
|
||||
messages=[
|
||||
{"role": "system", "content": system_prompt},
|
||||
{"role": "user", "content": user_prompt},
|
||||
],
|
||||
tools=_tools,
|
||||
)
|
||||
|
||||
to_be_deleted = []
|
||||
for item in memory_updates["tool_calls"]:
|
||||
if item["name"] == "delete_graph_memory":
|
||||
to_be_deleted.append(item["arguments"])
|
||||
# in case if it is not in the correct format
|
||||
to_be_deleted = self._remove_spaces_from_entities(to_be_deleted)
|
||||
logger.debug(f"Deleted relationships: {to_be_deleted}")
|
||||
return to_be_deleted
|
||||
|
||||
def _delete_entities(self, to_be_deleted, user_id):
|
||||
"""
|
||||
Delete the entities from the graph.
|
||||
"""
|
||||
|
||||
results = []
|
||||
for item in to_be_deleted:
|
||||
source = item["source"]
|
||||
destination = item["destination"]
|
||||
relationship = item["relationship"]
|
||||
|
||||
# Delete the specific relationship between nodes
|
||||
cypher, params = self._delete_entities_cypher(source, destination, relationship, user_id)
|
||||
result = self.graph.query(cypher, params=params)
|
||||
results.append(result)
|
||||
return results
|
||||
|
||||
@abstractmethod
|
||||
def _delete_entities_cypher(self, source, destination, relationship, user_id):
|
||||
"""
|
||||
Returns the OpenCypher query and parameters for deleting entities in the graph DB
|
||||
"""
|
||||
|
||||
pass
|
||||
|
||||
def _add_entities(self, to_be_added, user_id, entity_type_map):
|
||||
"""
|
||||
Add the new entities to the graph. Merge the nodes if they already exist.
|
||||
"""
|
||||
|
||||
results = []
|
||||
for item in to_be_added:
|
||||
# entities
|
||||
source = item["source"]
|
||||
destination = item["destination"]
|
||||
relationship = item["relationship"]
|
||||
|
||||
# types
|
||||
source_type = entity_type_map.get(source, "__User__")
|
||||
destination_type = entity_type_map.get(destination, "__User__")
|
||||
|
||||
# embeddings
|
||||
source_embedding = self.embedding_model.embed(source)
|
||||
dest_embedding = self.embedding_model.embed(destination)
|
||||
|
||||
# search for the nodes with the closest embeddings
|
||||
source_node_search_result = self._search_source_node(source_embedding, user_id, threshold=self.threshold)
|
||||
destination_node_search_result = self._search_destination_node(dest_embedding, user_id, threshold=self.threshold)
|
||||
|
||||
cypher, params = self._add_entities_cypher(
|
||||
source_node_search_result,
|
||||
source,
|
||||
source_embedding,
|
||||
source_type,
|
||||
destination_node_search_result,
|
||||
destination,
|
||||
dest_embedding,
|
||||
destination_type,
|
||||
relationship,
|
||||
user_id,
|
||||
)
|
||||
result = self.graph.query(cypher, params=params)
|
||||
results.append(result)
|
||||
return results
|
||||
|
||||
def _add_entities_cypher(
|
||||
self,
|
||||
source_node_list,
|
||||
source,
|
||||
source_embedding,
|
||||
source_type,
|
||||
destination_node_list,
|
||||
destination,
|
||||
dest_embedding,
|
||||
destination_type,
|
||||
relationship,
|
||||
user_id,
|
||||
):
|
||||
"""
|
||||
Returns the OpenCypher query and parameters for adding entities in the graph DB
|
||||
"""
|
||||
if not destination_node_list and source_node_list:
|
||||
return self._add_entities_by_source_cypher(
|
||||
source_node_list,
|
||||
destination,
|
||||
dest_embedding,
|
||||
destination_type,
|
||||
relationship,
|
||||
user_id)
|
||||
elif destination_node_list and not source_node_list:
|
||||
return self._add_entities_by_destination_cypher(
|
||||
source,
|
||||
source_embedding,
|
||||
source_type,
|
||||
destination_node_list,
|
||||
relationship,
|
||||
user_id)
|
||||
elif source_node_list and destination_node_list:
|
||||
return self._add_relationship_entities_cypher(
|
||||
source_node_list,
|
||||
destination_node_list,
|
||||
relationship,
|
||||
user_id)
|
||||
# else source_node_list and destination_node_list are empty
|
||||
return self._add_new_entities_cypher(
|
||||
source,
|
||||
source_embedding,
|
||||
source_type,
|
||||
destination,
|
||||
dest_embedding,
|
||||
destination_type,
|
||||
relationship,
|
||||
user_id)
|
||||
|
||||
@abstractmethod
|
||||
def _add_entities_by_source_cypher(
|
||||
self,
|
||||
source_node_list,
|
||||
destination,
|
||||
dest_embedding,
|
||||
destination_type,
|
||||
relationship,
|
||||
user_id,
|
||||
):
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def _add_entities_by_destination_cypher(
|
||||
self,
|
||||
source,
|
||||
source_embedding,
|
||||
source_type,
|
||||
destination_node_list,
|
||||
relationship,
|
||||
user_id,
|
||||
):
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def _add_relationship_entities_cypher(
|
||||
self,
|
||||
source_node_list,
|
||||
destination_node_list,
|
||||
relationship,
|
||||
user_id,
|
||||
):
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def _add_new_entities_cypher(
|
||||
self,
|
||||
source,
|
||||
source_embedding,
|
||||
source_type,
|
||||
destination,
|
||||
dest_embedding,
|
||||
destination_type,
|
||||
relationship,
|
||||
user_id,
|
||||
):
|
||||
pass
|
||||
|
||||
def search(self, query, filters, top_k=100):
|
||||
"""
|
||||
Search for memories and related graph data.
|
||||
|
||||
Args:
|
||||
query (str): Query to search for.
|
||||
filters (dict): A dictionary containing filters to be applied during the search.
|
||||
top_k (int): The maximum number of nodes and relationships to retrieve. Defaults to 100.
|
||||
|
||||
Returns:
|
||||
dict: A dictionary containing:
|
||||
- "contexts": List of search results from the base data store.
|
||||
- "entities": List of related graph data based on the query.
|
||||
"""
|
||||
|
||||
entity_type_map = self._retrieve_nodes_from_data(query, filters)
|
||||
search_output = self._search_graph_db(node_list=list(entity_type_map.keys()), filters=filters)
|
||||
|
||||
if not search_output:
|
||||
return []
|
||||
|
||||
search_outputs_sequence = [
|
||||
[item["source"], item["relationship"], item["destination"]] for item in search_output
|
||||
]
|
||||
bm25 = BM25Okapi(search_outputs_sequence)
|
||||
|
||||
tokenized_query = query.split(" ")
|
||||
reranked_results = bm25.get_top_n(tokenized_query, search_outputs_sequence, n=5)
|
||||
|
||||
search_results = []
|
||||
for item in reranked_results:
|
||||
search_results.append({"source": item[0], "relationship": item[1], "destination": item[2]})
|
||||
|
||||
return search_results
|
||||
|
||||
def _search_source_node(self, source_embedding, user_id, threshold=0.9):
|
||||
cypher, params = self._search_source_node_cypher(source_embedding, user_id, threshold)
|
||||
result = self.graph.query(cypher, params=params)
|
||||
return result
|
||||
|
||||
@abstractmethod
|
||||
def _search_source_node_cypher(self, source_embedding, user_id, threshold):
|
||||
"""
|
||||
Returns the OpenCypher query and parameters to search for source nodes
|
||||
"""
|
||||
pass
|
||||
|
||||
def _search_destination_node(self, destination_embedding, user_id, threshold=0.9):
|
||||
cypher, params = self._search_destination_node_cypher(destination_embedding, user_id, threshold)
|
||||
result = self.graph.query(cypher, params=params)
|
||||
return result
|
||||
|
||||
@abstractmethod
|
||||
def _search_destination_node_cypher(self, destination_embedding, user_id, threshold):
|
||||
"""
|
||||
Returns the OpenCypher query and parameters to search for destination nodes
|
||||
"""
|
||||
pass
|
||||
|
||||
def delete(self, data, filters):
|
||||
"""
|
||||
Delete graph entities associated with the given memory text.
|
||||
|
||||
Extracts entities and relationships from the memory text using the same
|
||||
pipeline as add(), then deletes the matching relationships in the graph.
|
||||
|
||||
Args:
|
||||
data (str): The memory text whose graph entities should be removed.
|
||||
filters (dict): Scope filters (user_id, agent_id, run_id).
|
||||
"""
|
||||
try:
|
||||
entity_type_map = self._retrieve_nodes_from_data(data, filters)
|
||||
if not entity_type_map:
|
||||
logger.debug("No entities found in memory text, skipping graph cleanup")
|
||||
return
|
||||
to_be_deleted = self._establish_nodes_relations_from_data(data, filters, entity_type_map)
|
||||
if to_be_deleted:
|
||||
self._delete_entities(to_be_deleted, filters["user_id"])
|
||||
except Exception as e:
|
||||
logger.error(f"Error during graph cleanup for memory delete: {e}")
|
||||
|
||||
def delete_all(self, filters):
|
||||
cypher, params = self._delete_all_cypher(filters)
|
||||
self.graph.query(cypher, params=params)
|
||||
|
||||
@abstractmethod
|
||||
def _delete_all_cypher(self, filters):
|
||||
"""
|
||||
Returns the OpenCypher query and parameters to delete all edges/nodes in the memory store
|
||||
"""
|
||||
pass
|
||||
|
||||
def get_all(self, filters, top_k=100):
|
||||
"""
|
||||
Retrieves all nodes and relationships from the graph database based on filtering criteria.
|
||||
|
||||
Args:
|
||||
filters (dict): A dictionary containing filters to be applied during the retrieval.
|
||||
top_k (int): The maximum number of nodes and relationships to retrieve. Defaults to 100.
|
||||
Returns:
|
||||
list: A list of dictionaries, each containing:
|
||||
- 'contexts': The base data store response for each memory.
|
||||
- 'entities': A list of strings representing the nodes and relationships
|
||||
"""
|
||||
|
||||
# return all nodes and relationships
|
||||
query, params = self._get_all_cypher(filters, top_k)
|
||||
results = self.graph.query(query, params=params)
|
||||
|
||||
final_results = []
|
||||
for result in results:
|
||||
final_results.append(
|
||||
{
|
||||
"source": result["source"],
|
||||
"relationship": result["relationship"],
|
||||
"target": result["target"],
|
||||
}
|
||||
)
|
||||
|
||||
logger.debug(f"Retrieved {len(final_results)} relationships")
|
||||
|
||||
return final_results
|
||||
|
||||
@abstractmethod
|
||||
def _get_all_cypher(self, filters, top_k):
|
||||
"""
|
||||
Returns the OpenCypher query and parameters to get all edges/nodes in the memory store
|
||||
"""
|
||||
pass
|
||||
|
||||
def _search_graph_db(self, node_list, filters, top_k=100):
|
||||
"""
|
||||
Search similar nodes among and their respective incoming and outgoing relations.
|
||||
"""
|
||||
result_relations = []
|
||||
|
||||
for node in node_list:
|
||||
n_embedding = self.embedding_model.embed(node)
|
||||
cypher_query, params = self._search_graph_db_cypher(n_embedding, filters, top_k)
|
||||
ans = self.graph.query(cypher_query, params=params)
|
||||
result_relations.extend(ans)
|
||||
|
||||
return result_relations
|
||||
|
||||
@abstractmethod
|
||||
def _search_graph_db_cypher(self, n_embedding, filters, top_k):
|
||||
"""
|
||||
Returns the OpenCypher query and parameters to search for similar nodes in the memory store
|
||||
"""
|
||||
pass
|
||||
|
||||
# Reset is not defined in base.py
|
||||
def reset(self):
|
||||
"""
|
||||
Reset the graph by clearing all nodes and relationships.
|
||||
|
||||
link: https://boto3.amazonaws.com/v1/documentation/api/latest/reference/services/neptune-graph/client/reset_graph.html
|
||||
"""
|
||||
|
||||
logger.warning("Clearing graph...")
|
||||
graph_id = self.graph.graph_identifier
|
||||
self.graph.client.reset_graph(
|
||||
graphIdentifier=graph_id,
|
||||
skipSnapshot=True,
|
||||
)
|
||||
waiter = self.graph.client.get_waiter("graph_available")
|
||||
waiter.wait(graphIdentifier=graph_id, WaiterConfig={"Delay": 10, "MaxAttempts": 60})
|
||||
@@ -1,511 +0,0 @@
|
||||
import logging
|
||||
import uuid
|
||||
from datetime import datetime, timezone
|
||||
|
||||
from .base import NeptuneBase
|
||||
|
||||
try:
|
||||
from langchain_aws import NeptuneGraph
|
||||
except ImportError:
|
||||
raise ImportError("langchain_aws is not installed. Please install it using 'make install_all'.")
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
class MemoryGraph(NeptuneBase):
|
||||
def __init__(self, config):
|
||||
"""
|
||||
Initialize the Neptune DB memory store.
|
||||
"""
|
||||
|
||||
self.config = config
|
||||
|
||||
self.graph = None
|
||||
endpoint = self.config.graph_store.config.endpoint
|
||||
if endpoint and endpoint.startswith("neptune-db://"):
|
||||
host = endpoint.replace("neptune-db://", "")
|
||||
port = 8182
|
||||
self.graph = NeptuneGraph(host, port)
|
||||
|
||||
if not self.graph:
|
||||
raise ValueError("Unable to create a Neptune-DB client: missing 'endpoint' in config")
|
||||
|
||||
self.node_label = ":`__Entity__`" if self.config.graph_store.config.base_label else ""
|
||||
|
||||
self.embedding_model = NeptuneBase._create_embedding_model(self.config)
|
||||
|
||||
# Default to openai if no specific provider is configured
|
||||
self.llm_provider = "openai"
|
||||
if self.config.graph_store.llm:
|
||||
self.llm_provider = self.config.graph_store.llm.provider
|
||||
elif self.config.llm.provider:
|
||||
self.llm_provider = self.config.llm.provider
|
||||
|
||||
# fetch the vector store as a provider
|
||||
self.vector_store_provider = self.config.vector_store.provider
|
||||
if self.config.graph_store.config.collection_name:
|
||||
vector_store_collection_name = self.config.graph_store.config.collection_name
|
||||
else:
|
||||
vector_store_config = self.config.vector_store.config
|
||||
if vector_store_config.collection_name:
|
||||
vector_store_collection_name = vector_store_config.collection_name + "_neptune_vector_store"
|
||||
else:
|
||||
vector_store_collection_name = "mem0_neptune_vector_store"
|
||||
self.config.vector_store.config.collection_name = vector_store_collection_name
|
||||
self.vector_store = NeptuneBase._create_vector_store(self.vector_store_provider, self.config)
|
||||
|
||||
self.llm = NeptuneBase._create_llm(self.config, self.llm_provider)
|
||||
self.user_id = None
|
||||
# Use threshold from graph_store config, default to 0.7 for backward compatibility
|
||||
self.threshold = self.config.graph_store.threshold if hasattr(self.config.graph_store, 'threshold') else 0.7
|
||||
self.vector_store_limit=5
|
||||
|
||||
def _delete_entities_cypher(self, source, destination, relationship, user_id):
|
||||
"""
|
||||
Returns the OpenCypher query and parameters for deleting entities in the graph DB
|
||||
|
||||
:param source: source node
|
||||
:param destination: destination node
|
||||
:param relationship: relationship label
|
||||
:param user_id: user_id to use
|
||||
:return: str, dict
|
||||
"""
|
||||
|
||||
cypher = f"""
|
||||
MATCH (n {self.node_label} {{name: $source_name, user_id: $user_id}})
|
||||
-[r:{relationship}]->
|
||||
(m {self.node_label} {{name: $dest_name, user_id: $user_id}})
|
||||
DELETE r
|
||||
RETURN
|
||||
n.name AS source,
|
||||
m.name AS target,
|
||||
type(r) AS relationship
|
||||
"""
|
||||
params = {
|
||||
"source_name": source,
|
||||
"dest_name": destination,
|
||||
"user_id": user_id,
|
||||
}
|
||||
logger.debug(f"_delete_entities\n query={cypher}")
|
||||
return cypher, params
|
||||
|
||||
def _add_entities_by_source_cypher(
|
||||
self,
|
||||
source_node_list,
|
||||
destination,
|
||||
dest_embedding,
|
||||
destination_type,
|
||||
relationship,
|
||||
user_id,
|
||||
):
|
||||
"""
|
||||
Returns the OpenCypher query and parameters for adding entities in the graph DB
|
||||
|
||||
:param source_node_list: list of source nodes
|
||||
:param destination: destination name
|
||||
:param dest_embedding: destination embedding
|
||||
:param destination_type: destination node label
|
||||
:param relationship: relationship label
|
||||
:param user_id: user id to use
|
||||
:return: str, dict
|
||||
"""
|
||||
destination_id = str(uuid.uuid4())
|
||||
destination_payload = {
|
||||
"name": destination,
|
||||
"type": destination_type,
|
||||
"user_id": user_id,
|
||||
"created_at": datetime.now(timezone.utc).isoformat(),
|
||||
}
|
||||
self.vector_store.insert(
|
||||
vectors=[dest_embedding],
|
||||
payloads=[destination_payload],
|
||||
ids=[destination_id],
|
||||
)
|
||||
|
||||
destination_label = self.node_label if self.node_label else f":`{destination_type}`"
|
||||
destination_extra_set = f", destination:`{destination_type}`" if self.node_label else ""
|
||||
|
||||
cypher = f"""
|
||||
MATCH (source {{user_id: $user_id}})
|
||||
WHERE id(source) = $source_id
|
||||
SET source.mentions = coalesce(source.mentions, 0) + 1
|
||||
WITH source
|
||||
MERGE (destination {destination_label} {{`~id`: $destination_id, name: $destination_name, user_id: $user_id}})
|
||||
ON CREATE SET
|
||||
destination.created = timestamp(),
|
||||
destination.updated = timestamp(),
|
||||
destination.mentions = 1
|
||||
{destination_extra_set}
|
||||
ON MATCH SET
|
||||
destination.mentions = coalesce(destination.mentions, 0) + 1,
|
||||
destination.updated = timestamp()
|
||||
WITH source, destination
|
||||
MERGE (source)-[r:{relationship}]->(destination)
|
||||
ON CREATE SET
|
||||
r.created = timestamp(),
|
||||
r.updated = timestamp(),
|
||||
r.mentions = 1
|
||||
ON MATCH SET
|
||||
r.mentions = coalesce(r.mentions, 0) + 1,
|
||||
r.updated = timestamp()
|
||||
RETURN source.name AS source, type(r) AS relationship, destination.name AS target, id(destination) AS destination_id
|
||||
"""
|
||||
|
||||
params = {
|
||||
"source_id": source_node_list[0]["id(source_candidate)"],
|
||||
"destination_id": destination_id,
|
||||
"destination_name": destination,
|
||||
"dest_embedding": dest_embedding,
|
||||
"user_id": user_id,
|
||||
}
|
||||
|
||||
logger.debug(
|
||||
f"_add_entities:\n source_node_search_result={source_node_list[0]}\n query={cypher}"
|
||||
)
|
||||
return cypher, params
|
||||
|
||||
def _add_entities_by_destination_cypher(
|
||||
self,
|
||||
source,
|
||||
source_embedding,
|
||||
source_type,
|
||||
destination_node_list,
|
||||
relationship,
|
||||
user_id,
|
||||
):
|
||||
"""
|
||||
Returns the OpenCypher query and parameters for adding entities in the graph DB
|
||||
|
||||
:param source: source node name
|
||||
:param source_embedding: source node embedding
|
||||
:param source_type: source node label
|
||||
:param destination_node_list: list of dest nodes
|
||||
:param relationship: relationship label
|
||||
:param user_id: user id to use
|
||||
:return: str, dict
|
||||
"""
|
||||
source_id = str(uuid.uuid4())
|
||||
source_payload = {
|
||||
"name": source,
|
||||
"type": source_type,
|
||||
"user_id": user_id,
|
||||
"created_at": datetime.now(timezone.utc).isoformat(),
|
||||
}
|
||||
self.vector_store.insert(
|
||||
vectors=[source_embedding],
|
||||
payloads=[source_payload],
|
||||
ids=[source_id],
|
||||
)
|
||||
|
||||
source_label = self.node_label if self.node_label else f":`{source_type}`"
|
||||
source_extra_set = f", source:`{source_type}`" if self.node_label else ""
|
||||
|
||||
cypher = f"""
|
||||
MATCH (destination {{user_id: $user_id}})
|
||||
WHERE id(destination) = $destination_id
|
||||
SET
|
||||
destination.mentions = coalesce(destination.mentions, 0) + 1,
|
||||
destination.updated = timestamp()
|
||||
WITH destination
|
||||
MERGE (source {source_label} {{`~id`: $source_id, name: $source_name, user_id: $user_id}})
|
||||
ON CREATE SET
|
||||
source.created = timestamp(),
|
||||
source.updated = timestamp(),
|
||||
source.mentions = 1
|
||||
{source_extra_set}
|
||||
ON MATCH SET
|
||||
source.mentions = coalesce(source.mentions, 0) + 1,
|
||||
source.updated = timestamp()
|
||||
WITH source, destination
|
||||
MERGE (source)-[r:{relationship}]->(destination)
|
||||
ON CREATE SET
|
||||
r.created = timestamp(),
|
||||
r.updated = timestamp(),
|
||||
r.mentions = 1
|
||||
ON MATCH SET
|
||||
r.mentions = coalesce(r.mentions, 0) + 1,
|
||||
r.updated = timestamp()
|
||||
RETURN source.name AS source, type(r) AS relationship, destination.name AS target
|
||||
"""
|
||||
|
||||
params = {
|
||||
"destination_id": destination_node_list[0]["id(destination_candidate)"],
|
||||
"source_id": source_id,
|
||||
"source_name": source,
|
||||
"source_embedding": source_embedding,
|
||||
"user_id": user_id,
|
||||
}
|
||||
logger.debug(
|
||||
f"_add_entities:\n destination_node_search_result={destination_node_list[0]}\n query={cypher}"
|
||||
)
|
||||
return cypher, params
|
||||
|
||||
def _add_relationship_entities_cypher(
|
||||
self,
|
||||
source_node_list,
|
||||
destination_node_list,
|
||||
relationship,
|
||||
user_id,
|
||||
):
|
||||
"""
|
||||
Returns the OpenCypher query and parameters for adding entities in the graph DB
|
||||
|
||||
:param source_node_list: list of source node ids
|
||||
:param destination_node_list: list of dest node ids
|
||||
:param relationship: relationship label
|
||||
:param user_id: user id to use
|
||||
:return: str, dict
|
||||
"""
|
||||
|
||||
cypher = f"""
|
||||
MATCH (source {{user_id: $user_id}})
|
||||
WHERE id(source) = $source_id
|
||||
SET
|
||||
source.mentions = coalesce(source.mentions, 0) + 1,
|
||||
source.updated = timestamp()
|
||||
WITH source
|
||||
MATCH (destination {{user_id: $user_id}})
|
||||
WHERE id(destination) = $destination_id
|
||||
SET
|
||||
destination.mentions = coalesce(destination.mentions) + 1,
|
||||
destination.updated = timestamp()
|
||||
MERGE (source)-[r:{relationship}]->(destination)
|
||||
ON CREATE SET
|
||||
r.created_at = timestamp(),
|
||||
r.updated_at = timestamp(),
|
||||
r.mentions = 1
|
||||
ON MATCH SET r.mentions = coalesce(r.mentions, 0) + 1
|
||||
RETURN source.name AS source, type(r) AS relationship, destination.name AS target
|
||||
"""
|
||||
params = {
|
||||
"source_id": source_node_list[0]["id(source_candidate)"],
|
||||
"destination_id": destination_node_list[0]["id(destination_candidate)"],
|
||||
"user_id": user_id,
|
||||
}
|
||||
logger.debug(
|
||||
f"_add_entities:\n destination_node_search_result={destination_node_list[0]}\n source_node_search_result={source_node_list[0]}\n query={cypher}"
|
||||
)
|
||||
return cypher, params
|
||||
|
||||
def _add_new_entities_cypher(
|
||||
self,
|
||||
source,
|
||||
source_embedding,
|
||||
source_type,
|
||||
destination,
|
||||
dest_embedding,
|
||||
destination_type,
|
||||
relationship,
|
||||
user_id,
|
||||
):
|
||||
"""
|
||||
Returns the OpenCypher query and parameters for adding entities in the graph DB
|
||||
|
||||
:param source: source node name
|
||||
:param source_embedding: source node embedding
|
||||
:param source_type: source node label
|
||||
:param destination: destination name
|
||||
:param dest_embedding: destination embedding
|
||||
:param destination_type: destination node label
|
||||
:param relationship: relationship label
|
||||
:param user_id: user id to use
|
||||
:return: str, dict
|
||||
"""
|
||||
source_id = str(uuid.uuid4())
|
||||
source_payload = {
|
||||
"name": source,
|
||||
"type": source_type,
|
||||
"user_id": user_id,
|
||||
"created_at": datetime.now(timezone.utc).isoformat(),
|
||||
}
|
||||
destination_id = str(uuid.uuid4())
|
||||
destination_payload = {
|
||||
"name": destination,
|
||||
"type": destination_type,
|
||||
"user_id": user_id,
|
||||
"created_at": datetime.now(timezone.utc).isoformat(),
|
||||
}
|
||||
self.vector_store.insert(
|
||||
vectors=[source_embedding, dest_embedding],
|
||||
payloads=[source_payload, destination_payload],
|
||||
ids=[source_id, destination_id],
|
||||
)
|
||||
|
||||
source_label = self.node_label if self.node_label else f":`{source_type}`"
|
||||
source_extra_set = f", source:`{source_type}`" if self.node_label else ""
|
||||
destination_label = self.node_label if self.node_label else f":`{destination_type}`"
|
||||
destination_extra_set = f", destination:`{destination_type}`" if self.node_label else ""
|
||||
|
||||
cypher = f"""
|
||||
MERGE (n {source_label} {{name: $source_name, user_id: $user_id, `~id`: $source_id}})
|
||||
ON CREATE SET n.created = timestamp(),
|
||||
n.mentions = 1
|
||||
{source_extra_set}
|
||||
ON MATCH SET n.mentions = coalesce(n.mentions, 0) + 1
|
||||
WITH n
|
||||
MERGE (m {destination_label} {{name: $dest_name, user_id: $user_id, `~id`: $dest_id}})
|
||||
ON CREATE SET m.created = timestamp(),
|
||||
m.mentions = 1
|
||||
{destination_extra_set}
|
||||
ON MATCH SET m.mentions = coalesce(m.mentions, 0) + 1
|
||||
WITH n, m
|
||||
MERGE (n)-[rel:{relationship}]->(m)
|
||||
ON CREATE SET rel.created = timestamp(), rel.mentions = 1
|
||||
ON MATCH SET rel.mentions = coalesce(rel.mentions, 0) + 1
|
||||
RETURN n.name AS source, type(rel) AS relationship, m.name AS target
|
||||
"""
|
||||
params = {
|
||||
"source_id": source_id,
|
||||
"dest_id": destination_id,
|
||||
"source_name": source,
|
||||
"dest_name": destination,
|
||||
"source_embedding": source_embedding,
|
||||
"dest_embedding": dest_embedding,
|
||||
"user_id": user_id,
|
||||
}
|
||||
logger.debug(
|
||||
f"_add_new_entities_cypher:\n query={cypher}"
|
||||
)
|
||||
return cypher, params
|
||||
|
||||
def _search_source_node_cypher(self, source_embedding, user_id, threshold):
|
||||
"""
|
||||
Returns the OpenCypher query and parameters to search for source nodes
|
||||
|
||||
:param source_embedding: source vector
|
||||
:param user_id: user_id to use
|
||||
:param threshold: the threshold for similarity
|
||||
:return: str, dict
|
||||
"""
|
||||
|
||||
source_nodes = self.vector_store.search(
|
||||
query="",
|
||||
vectors=source_embedding,
|
||||
top_k=self.vector_store_limit,
|
||||
filters={"user_id": user_id},
|
||||
)
|
||||
|
||||
ids = [n.id for n in filter(lambda s: s.score > threshold, source_nodes)]
|
||||
|
||||
cypher = f"""
|
||||
MATCH (source_candidate {self.node_label})
|
||||
WHERE source_candidate.user_id = $user_id AND id(source_candidate) IN $ids
|
||||
RETURN id(source_candidate)
|
||||
"""
|
||||
|
||||
params = {
|
||||
"ids": ids,
|
||||
"source_embedding": source_embedding,
|
||||
"user_id": user_id,
|
||||
"threshold": threshold,
|
||||
}
|
||||
logger.debug(f"_search_source_node\n query={cypher}")
|
||||
return cypher, params
|
||||
|
||||
def _search_destination_node_cypher(self, destination_embedding, user_id, threshold):
|
||||
"""
|
||||
Returns the OpenCypher query and parameters to search for destination nodes
|
||||
|
||||
:param source_embedding: source vector
|
||||
:param user_id: user_id to use
|
||||
:param threshold: the threshold for similarity
|
||||
:return: str, dict
|
||||
"""
|
||||
destination_nodes = self.vector_store.search(
|
||||
query="",
|
||||
vectors=destination_embedding,
|
||||
top_k=self.vector_store_limit,
|
||||
filters={"user_id": user_id},
|
||||
)
|
||||
|
||||
ids = [n.id for n in filter(lambda d: d.score > threshold, destination_nodes)]
|
||||
|
||||
cypher = f"""
|
||||
MATCH (destination_candidate {self.node_label})
|
||||
WHERE destination_candidate.user_id = $user_id AND id(destination_candidate) IN $ids
|
||||
RETURN id(destination_candidate)
|
||||
"""
|
||||
|
||||
params = {
|
||||
"ids": ids,
|
||||
"destination_embedding": destination_embedding,
|
||||
"user_id": user_id,
|
||||
}
|
||||
|
||||
logger.debug(f"_search_destination_node\n query={cypher}")
|
||||
return cypher, params
|
||||
|
||||
def _delete_all_cypher(self, filters):
|
||||
"""
|
||||
Returns the OpenCypher query and parameters to delete all edges/nodes in the memory store
|
||||
|
||||
:param filters: search filters
|
||||
:return: str, dict
|
||||
"""
|
||||
|
||||
# remove the vector store index
|
||||
self.vector_store.reset()
|
||||
|
||||
# create a query that: deletes the nodes of the graph_store
|
||||
cypher = f"""
|
||||
MATCH (n {self.node_label} {{user_id: $user_id}})
|
||||
DETACH DELETE n
|
||||
"""
|
||||
params = {"user_id": filters["user_id"]}
|
||||
|
||||
logger.debug(f"delete_all query={cypher}")
|
||||
return cypher, params
|
||||
|
||||
def _get_all_cypher(self, filters, top_k):
|
||||
"""
|
||||
Returns the OpenCypher query and parameters to get all edges/nodes in the memory store
|
||||
|
||||
:param filters: search filters
|
||||
:param top_k: return limit
|
||||
:return: str, dict
|
||||
"""
|
||||
|
||||
cypher = f"""
|
||||
MATCH (n {self.node_label} {{user_id: $user_id}})-[r]->(m {self.node_label} {{user_id: $user_id}})
|
||||
RETURN n.name AS source, type(r) AS relationship, m.name AS target
|
||||
LIMIT $limit
|
||||
"""
|
||||
params = {"user_id": filters["user_id"], "limit": top_k}
|
||||
return cypher, params
|
||||
|
||||
def _search_graph_db_cypher(self, n_embedding, filters, top_k):
|
||||
"""
|
||||
Returns the OpenCypher query and parameters to search for similar nodes in the memory store
|
||||
|
||||
:param n_embedding: node vector
|
||||
:param filters: search filters
|
||||
:param top_k: return limit
|
||||
:return: str, dict
|
||||
"""
|
||||
|
||||
# search vector store for applicable nodes using cosine similarity
|
||||
search_nodes = self.vector_store.search(
|
||||
query="",
|
||||
vectors=n_embedding,
|
||||
top_k=self.vector_store_limit,
|
||||
filters=filters,
|
||||
)
|
||||
|
||||
ids = [n.id for n in search_nodes]
|
||||
|
||||
cypher_query = f"""
|
||||
MATCH (n {self.node_label})-[r]->(m)
|
||||
WHERE n.user_id = $user_id AND id(n) IN $n_ids
|
||||
RETURN n.name AS source, id(n) AS source_id, type(r) AS relationship, id(r) AS relation_id, m.name AS destination, id(m) AS destination_id
|
||||
UNION
|
||||
MATCH (m)-[r]->(n {self.node_label})
|
||||
RETURN m.name AS source, id(m) AS source_id, type(r) AS relationship, id(r) AS relation_id, n.name AS destination, id(n) AS destination_id
|
||||
LIMIT $limit
|
||||
"""
|
||||
params = {
|
||||
"n_ids": ids,
|
||||
"user_id": filters["user_id"],
|
||||
"limit": top_k,
|
||||
}
|
||||
logger.debug(f"_search_graph_db\n query={cypher_query}")
|
||||
|
||||
return cypher_query, params
|
||||
@@ -1,475 +0,0 @@
|
||||
import logging
|
||||
|
||||
from .base import NeptuneBase
|
||||
|
||||
try:
|
||||
from botocore.config import Config
|
||||
from langchain_aws import NeptuneAnalyticsGraph
|
||||
except ImportError:
|
||||
raise ImportError("langchain_aws is not installed. Please install it using 'make install_all'.")
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class MemoryGraph(NeptuneBase):
|
||||
def __init__(self, config):
|
||||
self.config = config
|
||||
|
||||
self.graph = None
|
||||
endpoint = self.config.graph_store.config.endpoint
|
||||
app_id = self.config.graph_store.config.app_id
|
||||
if endpoint and endpoint.startswith("neptune-graph://"):
|
||||
graph_identifier = endpoint.replace("neptune-graph://", "")
|
||||
self.graph = NeptuneAnalyticsGraph(graph_identifier = graph_identifier,
|
||||
config = Config(user_agent_appid=app_id))
|
||||
|
||||
if not self.graph:
|
||||
raise ValueError("Unable to create a Neptune client: missing 'endpoint' in config")
|
||||
|
||||
self.node_label = ":`__Entity__`" if self.config.graph_store.config.base_label else ""
|
||||
|
||||
self.embedding_model = NeptuneBase._create_embedding_model(self.config)
|
||||
|
||||
# Default to openai if no specific provider is configured
|
||||
self.llm_provider = "openai"
|
||||
if self.config.llm.provider:
|
||||
self.llm_provider = self.config.llm.provider
|
||||
if self.config.graph_store.llm:
|
||||
self.llm_provider = self.config.graph_store.llm.provider
|
||||
|
||||
self.llm = NeptuneBase._create_llm(self.config, self.llm_provider)
|
||||
self.user_id = None
|
||||
# Use threshold from graph_store config, default to 0.7 for backward compatibility
|
||||
self.threshold = self.config.graph_store.threshold if hasattr(self.config.graph_store, 'threshold') else 0.7
|
||||
|
||||
def _delete_entities_cypher(self, source, destination, relationship, user_id):
|
||||
"""
|
||||
Returns the OpenCypher query and parameters for deleting entities in the graph DB
|
||||
|
||||
:param source: source node
|
||||
:param destination: destination node
|
||||
:param relationship: relationship label
|
||||
:param user_id: user_id to use
|
||||
:return: str, dict
|
||||
"""
|
||||
|
||||
cypher = f"""
|
||||
MATCH (n {self.node_label} {{name: $source_name, user_id: $user_id}})
|
||||
-[r:{relationship}]->
|
||||
(m {self.node_label} {{name: $dest_name, user_id: $user_id}})
|
||||
DELETE r
|
||||
RETURN
|
||||
n.name AS source,
|
||||
m.name AS target,
|
||||
type(r) AS relationship
|
||||
"""
|
||||
params = {
|
||||
"source_name": source,
|
||||
"dest_name": destination,
|
||||
"user_id": user_id,
|
||||
}
|
||||
logger.debug(f"_delete_entities\n query={cypher}")
|
||||
return cypher, params
|
||||
|
||||
def _add_entities_by_source_cypher(
|
||||
self,
|
||||
source_node_list,
|
||||
destination,
|
||||
dest_embedding,
|
||||
destination_type,
|
||||
relationship,
|
||||
user_id,
|
||||
):
|
||||
"""
|
||||
Returns the OpenCypher query and parameters for adding entities in the graph DB
|
||||
|
||||
:param source_node_list: list of source nodes
|
||||
:param destination: destination name
|
||||
:param dest_embedding: destination embedding
|
||||
:param destination_type: destination node label
|
||||
:param relationship: relationship label
|
||||
:param user_id: user id to use
|
||||
:return: str, dict
|
||||
"""
|
||||
|
||||
destination_label = self.node_label if self.node_label else f":`{destination_type}`"
|
||||
destination_extra_set = f", destination:`{destination_type}`" if self.node_label else ""
|
||||
|
||||
cypher = f"""
|
||||
MATCH (source {{user_id: $user_id}})
|
||||
WHERE id(source) = $source_id
|
||||
SET source.mentions = coalesce(source.mentions, 0) + 1
|
||||
WITH source
|
||||
MERGE (destination {destination_label} {{name: $destination_name, user_id: $user_id}})
|
||||
ON CREATE SET
|
||||
destination.created = timestamp(),
|
||||
destination.updated = timestamp(),
|
||||
destination.mentions = 1
|
||||
{destination_extra_set}
|
||||
ON MATCH SET
|
||||
destination.mentions = coalesce(destination.mentions, 0) + 1,
|
||||
destination.updated = timestamp()
|
||||
WITH source, destination, $dest_embedding as dest_embedding
|
||||
CALL neptune.algo.vectors.upsert(destination, dest_embedding)
|
||||
WITH source, destination
|
||||
MERGE (source)-[r:{relationship}]->(destination)
|
||||
ON CREATE SET
|
||||
r.created = timestamp(),
|
||||
r.updated = timestamp(),
|
||||
r.mentions = 1
|
||||
ON MATCH SET
|
||||
r.mentions = coalesce(r.mentions, 0) + 1,
|
||||
r.updated = timestamp()
|
||||
RETURN source.name AS source, type(r) AS relationship, destination.name AS target
|
||||
"""
|
||||
|
||||
params = {
|
||||
"source_id": source_node_list[0]["id(source_candidate)"],
|
||||
"destination_name": destination,
|
||||
"dest_embedding": dest_embedding,
|
||||
"user_id": user_id,
|
||||
}
|
||||
logger.debug(
|
||||
f"_add_entities:\n source_node_search_result={source_node_list[0]}\n query={cypher}"
|
||||
)
|
||||
return cypher, params
|
||||
|
||||
def _add_entities_by_destination_cypher(
|
||||
self,
|
||||
source,
|
||||
source_embedding,
|
||||
source_type,
|
||||
destination_node_list,
|
||||
relationship,
|
||||
user_id,
|
||||
):
|
||||
"""
|
||||
Returns the OpenCypher query and parameters for adding entities in the graph DB
|
||||
|
||||
:param source: source node name
|
||||
:param source_embedding: source node embedding
|
||||
:param source_type: source node label
|
||||
:param destination_node_list: list of dest nodes
|
||||
:param relationship: relationship label
|
||||
:param user_id: user id to use
|
||||
:return: str, dict
|
||||
"""
|
||||
|
||||
source_label = self.node_label if self.node_label else f":`{source_type}`"
|
||||
source_extra_set = f", source:`{source_type}`" if self.node_label else ""
|
||||
|
||||
cypher = f"""
|
||||
MATCH (destination {{user_id: $user_id}})
|
||||
WHERE id(destination) = $destination_id
|
||||
SET
|
||||
destination.mentions = coalesce(destination.mentions, 0) + 1,
|
||||
destination.updated = timestamp()
|
||||
WITH destination
|
||||
MERGE (source {source_label} {{name: $source_name, user_id: $user_id}})
|
||||
ON CREATE SET
|
||||
source.created = timestamp(),
|
||||
source.updated = timestamp(),
|
||||
source.mentions = 1
|
||||
{source_extra_set}
|
||||
ON MATCH SET
|
||||
source.mentions = coalesce(source.mentions, 0) + 1,
|
||||
source.updated = timestamp()
|
||||
WITH source, destination, $source_embedding as source_embedding
|
||||
CALL neptune.algo.vectors.upsert(source, source_embedding)
|
||||
WITH source, destination
|
||||
MERGE (source)-[r:{relationship}]->(destination)
|
||||
ON CREATE SET
|
||||
r.created = timestamp(),
|
||||
r.updated = timestamp(),
|
||||
r.mentions = 1
|
||||
ON MATCH SET
|
||||
r.mentions = coalesce(r.mentions, 0) + 1,
|
||||
r.updated = timestamp()
|
||||
RETURN source.name AS source, type(r) AS relationship, destination.name AS target
|
||||
"""
|
||||
|
||||
params = {
|
||||
"destination_id": destination_node_list[0]["id(destination_candidate)"],
|
||||
"source_name": source,
|
||||
"source_embedding": source_embedding,
|
||||
"user_id": user_id,
|
||||
}
|
||||
logger.debug(
|
||||
f"_add_entities:\n destination_node_search_result={destination_node_list[0]}\n query={cypher}"
|
||||
)
|
||||
return cypher, params
|
||||
|
||||
def _add_relationship_entities_cypher(
|
||||
self,
|
||||
source_node_list,
|
||||
destination_node_list,
|
||||
relationship,
|
||||
user_id,
|
||||
):
|
||||
"""
|
||||
Returns the OpenCypher query and parameters for adding entities in the graph DB
|
||||
|
||||
:param source_node_list: list of source node ids
|
||||
:param destination_node_list: list of dest node ids
|
||||
:param relationship: relationship label
|
||||
:param user_id: user id to use
|
||||
:return: str, dict
|
||||
"""
|
||||
|
||||
cypher = f"""
|
||||
MATCH (source {{user_id: $user_id}})
|
||||
WHERE id(source) = $source_id
|
||||
SET
|
||||
source.mentions = coalesce(source.mentions, 0) + 1,
|
||||
source.updated = timestamp()
|
||||
WITH source
|
||||
MATCH (destination {{user_id: $user_id}})
|
||||
WHERE id(destination) = $destination_id
|
||||
SET
|
||||
destination.mentions = coalesce(destination.mentions) + 1,
|
||||
destination.updated = timestamp()
|
||||
MERGE (source)-[r:{relationship}]->(destination)
|
||||
ON CREATE SET
|
||||
r.created_at = timestamp(),
|
||||
r.updated_at = timestamp(),
|
||||
r.mentions = 1
|
||||
ON MATCH SET r.mentions = coalesce(r.mentions, 0) + 1
|
||||
RETURN source.name AS source, type(r) AS relationship, destination.name AS target
|
||||
"""
|
||||
params = {
|
||||
"source_id": source_node_list[0]["id(source_candidate)"],
|
||||
"destination_id": destination_node_list[0]["id(destination_candidate)"],
|
||||
"user_id": user_id,
|
||||
}
|
||||
logger.debug(
|
||||
f"_add_entities:\n destination_node_search_result={destination_node_list[0]}\n source_node_search_result={source_node_list[0]}\n query={cypher}"
|
||||
)
|
||||
return cypher, params
|
||||
|
||||
def _add_new_entities_cypher(
|
||||
self,
|
||||
source,
|
||||
source_embedding,
|
||||
source_type,
|
||||
destination,
|
||||
dest_embedding,
|
||||
destination_type,
|
||||
relationship,
|
||||
user_id,
|
||||
):
|
||||
"""
|
||||
Returns the OpenCypher query and parameters for adding entities in the graph DB
|
||||
|
||||
:param source: source node name
|
||||
:param source_embedding: source node embedding
|
||||
:param source_type: source node label
|
||||
:param destination: destination name
|
||||
:param dest_embedding: destination embedding
|
||||
:param destination_type: destination node label
|
||||
:param relationship: relationship label
|
||||
:param user_id: user id to use
|
||||
:return: str, dict
|
||||
"""
|
||||
|
||||
source_label = self.node_label if self.node_label else f":`{source_type}`"
|
||||
source_extra_set = f", source:`{source_type}`" if self.node_label else ""
|
||||
destination_label = self.node_label if self.node_label else f":`{destination_type}`"
|
||||
destination_extra_set = f", destination:`{destination_type}`" if self.node_label else ""
|
||||
|
||||
cypher = f"""
|
||||
MERGE (n {source_label} {{name: $source_name, user_id: $user_id}})
|
||||
ON CREATE SET n.created = timestamp(),
|
||||
n.updated = timestamp(),
|
||||
n.mentions = 1
|
||||
{source_extra_set}
|
||||
ON MATCH SET
|
||||
n.mentions = coalesce(n.mentions, 0) + 1,
|
||||
n.updated = timestamp()
|
||||
WITH n, $source_embedding as source_embedding
|
||||
CALL neptune.algo.vectors.upsert(n, source_embedding)
|
||||
WITH n
|
||||
MERGE (m {destination_label} {{name: $dest_name, user_id: $user_id}})
|
||||
ON CREATE SET
|
||||
m.created = timestamp(),
|
||||
m.updated = timestamp(),
|
||||
m.mentions = 1
|
||||
{destination_extra_set}
|
||||
ON MATCH SET
|
||||
m.updated = timestamp(),
|
||||
m.mentions = coalesce(m.mentions, 0) + 1
|
||||
WITH n, m, $dest_embedding as dest_embedding
|
||||
CALL neptune.algo.vectors.upsert(m, dest_embedding)
|
||||
WITH n, m
|
||||
MERGE (n)-[rel:{relationship}]->(m)
|
||||
ON CREATE SET
|
||||
rel.created = timestamp(),
|
||||
rel.updated = timestamp(),
|
||||
rel.mentions = 1
|
||||
ON MATCH SET
|
||||
rel.updated = timestamp(),
|
||||
rel.mentions = coalesce(rel.mentions, 0) + 1
|
||||
RETURN n.name AS source, type(rel) AS relationship, m.name AS target
|
||||
"""
|
||||
params = {
|
||||
"source_name": source,
|
||||
"dest_name": destination,
|
||||
"source_embedding": source_embedding,
|
||||
"dest_embedding": dest_embedding,
|
||||
"user_id": user_id,
|
||||
}
|
||||
logger.debug(
|
||||
f"_add_new_entities_cypher:\n query={cypher}"
|
||||
)
|
||||
return cypher, params
|
||||
|
||||
def _search_source_node_cypher(self, source_embedding, user_id, threshold):
|
||||
"""
|
||||
Returns the OpenCypher query and parameters to search for source nodes
|
||||
|
||||
:param source_embedding: source vector
|
||||
:param user_id: user_id to use
|
||||
:param threshold: the threshold for similarity
|
||||
:return: str, dict
|
||||
"""
|
||||
cypher = f"""
|
||||
MATCH (source_candidate {self.node_label})
|
||||
WHERE source_candidate.user_id = $user_id
|
||||
|
||||
WITH source_candidate, $source_embedding as v_embedding
|
||||
CALL neptune.algo.vectors.distanceByEmbedding(
|
||||
v_embedding,
|
||||
source_candidate,
|
||||
{{metric:"CosineSimilarity"}}
|
||||
) YIELD distance
|
||||
WITH source_candidate, distance AS cosine_similarity
|
||||
WHERE cosine_similarity >= $threshold
|
||||
|
||||
WITH source_candidate, cosine_similarity
|
||||
ORDER BY cosine_similarity DESC
|
||||
LIMIT 1
|
||||
|
||||
RETURN id(source_candidate), cosine_similarity
|
||||
"""
|
||||
|
||||
params = {
|
||||
"source_embedding": source_embedding,
|
||||
"user_id": user_id,
|
||||
"threshold": threshold,
|
||||
}
|
||||
logger.debug(f"_search_source_node\n query={cypher}")
|
||||
return cypher, params
|
||||
|
||||
def _search_destination_node_cypher(self, destination_embedding, user_id, threshold):
|
||||
"""
|
||||
Returns the OpenCypher query and parameters to search for destination nodes
|
||||
|
||||
:param source_embedding: source vector
|
||||
:param user_id: user_id to use
|
||||
:param threshold: the threshold for similarity
|
||||
:return: str, dict
|
||||
"""
|
||||
cypher = f"""
|
||||
MATCH (destination_candidate {self.node_label})
|
||||
WHERE destination_candidate.user_id = $user_id
|
||||
|
||||
WITH destination_candidate, $destination_embedding as v_embedding
|
||||
CALL neptune.algo.vectors.distanceByEmbedding(
|
||||
v_embedding,
|
||||
destination_candidate,
|
||||
{{metric:"CosineSimilarity"}}
|
||||
) YIELD distance
|
||||
WITH destination_candidate, distance AS cosine_similarity
|
||||
WHERE cosine_similarity >= $threshold
|
||||
|
||||
WITH destination_candidate, cosine_similarity
|
||||
ORDER BY cosine_similarity DESC
|
||||
LIMIT 1
|
||||
|
||||
RETURN id(destination_candidate), cosine_similarity
|
||||
"""
|
||||
params = {
|
||||
"destination_embedding": destination_embedding,
|
||||
"user_id": user_id,
|
||||
"threshold": threshold,
|
||||
}
|
||||
|
||||
logger.debug(f"_search_destination_node\n query={cypher}")
|
||||
return cypher, params
|
||||
|
||||
def _delete_all_cypher(self, filters):
|
||||
"""
|
||||
Returns the OpenCypher query and parameters to delete all edges/nodes in the memory store
|
||||
|
||||
:param filters: search filters
|
||||
:return: str, dict
|
||||
"""
|
||||
cypher = f"""
|
||||
MATCH (n {self.node_label} {{user_id: $user_id}})
|
||||
DETACH DELETE n
|
||||
"""
|
||||
params = {"user_id": filters["user_id"]}
|
||||
|
||||
logger.debug(f"delete_all query={cypher}")
|
||||
return cypher, params
|
||||
|
||||
def _get_all_cypher(self, filters, top_k):
|
||||
"""
|
||||
Returns the OpenCypher query and parameters to get all edges/nodes in the memory store
|
||||
|
||||
:param filters: search filters
|
||||
:param top_k: return limit
|
||||
:return: str, dict
|
||||
"""
|
||||
|
||||
cypher = f"""
|
||||
MATCH (n {self.node_label} {{user_id: $user_id}})-[r]->(m {self.node_label} {{user_id: $user_id}})
|
||||
RETURN n.name AS source, type(r) AS relationship, m.name AS target
|
||||
LIMIT $limit
|
||||
"""
|
||||
params = {"user_id": filters["user_id"], "limit": top_k}
|
||||
return cypher, params
|
||||
|
||||
def _search_graph_db_cypher(self, n_embedding, filters, top_k):
|
||||
"""
|
||||
Returns the OpenCypher query and parameters to search for similar nodes in the memory store
|
||||
|
||||
:param n_embedding: node vector
|
||||
:param filters: search filters
|
||||
:param top_k: return limit
|
||||
:return: str, dict
|
||||
"""
|
||||
|
||||
cypher_query = f"""
|
||||
MATCH (n {self.node_label})
|
||||
WHERE n.user_id = $user_id
|
||||
WITH n, $n_embedding as n_embedding
|
||||
CALL neptune.algo.vectors.distanceByEmbedding(
|
||||
n_embedding,
|
||||
n,
|
||||
{{metric:"CosineSimilarity"}}
|
||||
) YIELD distance
|
||||
WITH n, distance as similarity
|
||||
WHERE similarity >= $threshold
|
||||
CALL {{
|
||||
WITH n
|
||||
MATCH (n)-[r]->(m)
|
||||
RETURN n.name AS source, id(n) AS source_id, type(r) AS relationship, id(r) AS relation_id, m.name AS destination, id(m) AS destination_id
|
||||
UNION ALL
|
||||
WITH n
|
||||
MATCH (m)-[r]->(n)
|
||||
RETURN m.name AS source, id(m) AS source_id, type(r) AS relationship, id(r) AS relation_id, n.name AS destination, id(n) AS destination_id
|
||||
}}
|
||||
WITH distinct source, source_id, relationship, relation_id, destination, destination_id, similarity
|
||||
RETURN source, source_id, relationship, relation_id, destination, destination_id, similarity
|
||||
ORDER BY similarity DESC
|
||||
LIMIT $limit
|
||||
"""
|
||||
params = {
|
||||
"n_embedding": n_embedding,
|
||||
"threshold": self.threshold,
|
||||
"user_id": filters["user_id"],
|
||||
"limit": top_k,
|
||||
}
|
||||
logger.debug(f"_search_graph_db\n query={cypher_query}")
|
||||
|
||||
return cypher_query, params
|
||||
@@ -1,371 +0,0 @@
|
||||
UPDATE_MEMORY_TOOL_GRAPH = {
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "update_graph_memory",
|
||||
"description": "Update the relationship key of an existing graph memory based on new information. This function should be called when there's a need to modify an existing relationship in the knowledge graph. The update should only be performed if the new information is more recent, more accurate, or provides additional context compared to the existing information. The source and destination nodes of the relationship must remain the same as in the existing graph memory; only the relationship itself can be updated.",
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"source": {
|
||||
"type": "string",
|
||||
"description": "The identifier of the source node in the relationship to be updated. This should match an existing node in the graph.",
|
||||
},
|
||||
"destination": {
|
||||
"type": "string",
|
||||
"description": "The identifier of the destination node in the relationship to be updated. This should match an existing node in the graph.",
|
||||
},
|
||||
"relationship": {
|
||||
"type": "string",
|
||||
"description": "The new or updated relationship between the source and destination nodes. This should be a concise, clear description of how the two nodes are connected.",
|
||||
},
|
||||
},
|
||||
"required": ["source", "destination", "relationship"],
|
||||
"additionalProperties": False,
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
ADD_MEMORY_TOOL_GRAPH = {
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "add_graph_memory",
|
||||
"description": "Add a new graph memory to the knowledge graph. This function creates a new relationship between two nodes, potentially creating new nodes if they don't exist.",
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"source": {
|
||||
"type": "string",
|
||||
"description": "The identifier of the source node in the new relationship. This can be an existing node or a new node to be created.",
|
||||
},
|
||||
"destination": {
|
||||
"type": "string",
|
||||
"description": "The identifier of the destination node in the new relationship. This can be an existing node or a new node to be created.",
|
||||
},
|
||||
"relationship": {
|
||||
"type": "string",
|
||||
"description": "The type of relationship between the source and destination nodes. This should be a concise, clear description of how the two nodes are connected.",
|
||||
},
|
||||
"source_type": {
|
||||
"type": "string",
|
||||
"description": "The type or category of the source node. This helps in classifying and organizing nodes in the graph.",
|
||||
},
|
||||
"destination_type": {
|
||||
"type": "string",
|
||||
"description": "The type or category of the destination node. This helps in classifying and organizing nodes in the graph.",
|
||||
},
|
||||
},
|
||||
"required": [
|
||||
"source",
|
||||
"destination",
|
||||
"relationship",
|
||||
"source_type",
|
||||
"destination_type",
|
||||
],
|
||||
"additionalProperties": False,
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
NOOP_TOOL = {
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "noop",
|
||||
"description": "No operation should be performed to the graph entities. This function is called when the system determines that no changes or additions are necessary based on the current input or context. It serves as a placeholder action when no other actions are required, ensuring that the system can explicitly acknowledge situations where no modifications to the graph are needed.",
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
"properties": {},
|
||||
"required": [],
|
||||
"additionalProperties": False,
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
RELATIONS_TOOL = {
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "establish_relationships",
|
||||
"description": "Establish relationships among the entities based on the provided text.",
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"entities": {
|
||||
"type": "array",
|
||||
"items": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"source": {"type": "string", "description": "The source entity of the relationship."},
|
||||
"relationship": {
|
||||
"type": "string",
|
||||
"description": "The relationship between the source and destination entities.",
|
||||
},
|
||||
"destination": {
|
||||
"type": "string",
|
||||
"description": "The destination entity of the relationship.",
|
||||
},
|
||||
},
|
||||
"required": [
|
||||
"source",
|
||||
"relationship",
|
||||
"destination",
|
||||
],
|
||||
"additionalProperties": False,
|
||||
},
|
||||
}
|
||||
},
|
||||
"required": ["entities"],
|
||||
"additionalProperties": False,
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
EXTRACT_ENTITIES_TOOL = {
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "extract_entities",
|
||||
"description": "Extract entities and their types from the text.",
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"entities": {
|
||||
"type": "array",
|
||||
"items": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"entity": {"type": "string", "description": "The name or identifier of the entity."},
|
||||
"entity_type": {"type": "string", "description": "The type or category of the entity."},
|
||||
},
|
||||
"required": ["entity", "entity_type"],
|
||||
"additionalProperties": False,
|
||||
},
|
||||
"description": "An array of entities with their types.",
|
||||
}
|
||||
},
|
||||
"required": ["entities"],
|
||||
"additionalProperties": False,
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
UPDATE_MEMORY_STRUCT_TOOL_GRAPH = {
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "update_graph_memory",
|
||||
"description": "Update the relationship key of an existing graph memory based on new information. This function should be called when there's a need to modify an existing relationship in the knowledge graph. The update should only be performed if the new information is more recent, more accurate, or provides additional context compared to the existing information. The source and destination nodes of the relationship must remain the same as in the existing graph memory; only the relationship itself can be updated.",
|
||||
"strict": True,
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"source": {
|
||||
"type": "string",
|
||||
"description": "The identifier of the source node in the relationship to be updated. This should match an existing node in the graph.",
|
||||
},
|
||||
"destination": {
|
||||
"type": "string",
|
||||
"description": "The identifier of the destination node in the relationship to be updated. This should match an existing node in the graph.",
|
||||
},
|
||||
"relationship": {
|
||||
"type": "string",
|
||||
"description": "The new or updated relationship between the source and destination nodes. This should be a concise, clear description of how the two nodes are connected.",
|
||||
},
|
||||
},
|
||||
"required": ["source", "destination", "relationship"],
|
||||
"additionalProperties": False,
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
ADD_MEMORY_STRUCT_TOOL_GRAPH = {
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "add_graph_memory",
|
||||
"description": "Add a new graph memory to the knowledge graph. This function creates a new relationship between two nodes, potentially creating new nodes if they don't exist.",
|
||||
"strict": True,
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"source": {
|
||||
"type": "string",
|
||||
"description": "The identifier of the source node in the new relationship. This can be an existing node or a new node to be created.",
|
||||
},
|
||||
"destination": {
|
||||
"type": "string",
|
||||
"description": "The identifier of the destination node in the new relationship. This can be an existing node or a new node to be created.",
|
||||
},
|
||||
"relationship": {
|
||||
"type": "string",
|
||||
"description": "The type of relationship between the source and destination nodes. This should be a concise, clear description of how the two nodes are connected.",
|
||||
},
|
||||
"source_type": {
|
||||
"type": "string",
|
||||
"description": "The type or category of the source node. This helps in classifying and organizing nodes in the graph.",
|
||||
},
|
||||
"destination_type": {
|
||||
"type": "string",
|
||||
"description": "The type or category of the destination node. This helps in classifying and organizing nodes in the graph.",
|
||||
},
|
||||
},
|
||||
"required": [
|
||||
"source",
|
||||
"destination",
|
||||
"relationship",
|
||||
"source_type",
|
||||
"destination_type",
|
||||
],
|
||||
"additionalProperties": False,
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
NOOP_STRUCT_TOOL = {
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "noop",
|
||||
"description": "No operation should be performed to the graph entities. This function is called when the system determines that no changes or additions are necessary based on the current input or context. It serves as a placeholder action when no other actions are required, ensuring that the system can explicitly acknowledge situations where no modifications to the graph are needed.",
|
||||
"strict": True,
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
"properties": {},
|
||||
"required": [],
|
||||
"additionalProperties": False,
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
RELATIONS_STRUCT_TOOL = {
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "establish_relations",
|
||||
"description": "Establish relationships among the entities based on the provided text.",
|
||||
"strict": True,
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"entities": {
|
||||
"type": "array",
|
||||
"items": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"source": {
|
||||
"type": "string",
|
||||
"description": "The source entity of the relationship.",
|
||||
},
|
||||
"relationship": {
|
||||
"type": "string",
|
||||
"description": "The relationship between the source and destination entities.",
|
||||
},
|
||||
"destination": {
|
||||
"type": "string",
|
||||
"description": "The destination entity of the relationship.",
|
||||
},
|
||||
},
|
||||
"required": [
|
||||
"source",
|
||||
"relationship",
|
||||
"destination",
|
||||
],
|
||||
"additionalProperties": False,
|
||||
},
|
||||
}
|
||||
},
|
||||
"required": ["entities"],
|
||||
"additionalProperties": False,
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
EXTRACT_ENTITIES_STRUCT_TOOL = {
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "extract_entities",
|
||||
"description": "Extract entities and their types from the text.",
|
||||
"strict": True,
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"entities": {
|
||||
"type": "array",
|
||||
"items": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"entity": {"type": "string", "description": "The name or identifier of the entity."},
|
||||
"entity_type": {"type": "string", "description": "The type or category of the entity."},
|
||||
},
|
||||
"required": ["entity", "entity_type"],
|
||||
"additionalProperties": False,
|
||||
},
|
||||
"description": "An array of entities with their types.",
|
||||
}
|
||||
},
|
||||
"required": ["entities"],
|
||||
"additionalProperties": False,
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
DELETE_MEMORY_STRUCT_TOOL_GRAPH = {
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "delete_graph_memory",
|
||||
"description": "Delete the relationship between two nodes. This function deletes the existing relationship.",
|
||||
"strict": True,
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"source": {
|
||||
"type": "string",
|
||||
"description": "The identifier of the source node in the relationship.",
|
||||
},
|
||||
"relationship": {
|
||||
"type": "string",
|
||||
"description": "The existing relationship between the source and destination nodes that needs to be deleted.",
|
||||
},
|
||||
"destination": {
|
||||
"type": "string",
|
||||
"description": "The identifier of the destination node in the relationship.",
|
||||
},
|
||||
},
|
||||
"required": [
|
||||
"source",
|
||||
"relationship",
|
||||
"destination",
|
||||
],
|
||||
"additionalProperties": False,
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
DELETE_MEMORY_TOOL_GRAPH = {
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "delete_graph_memory",
|
||||
"description": "Delete the relationship between two nodes. This function deletes the existing relationship.",
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"source": {
|
||||
"type": "string",
|
||||
"description": "The identifier of the source node in the relationship.",
|
||||
},
|
||||
"relationship": {
|
||||
"type": "string",
|
||||
"description": "The existing relationship between the source and destination nodes that needs to be deleted.",
|
||||
},
|
||||
"destination": {
|
||||
"type": "string",
|
||||
"description": "The identifier of the destination node in the relationship.",
|
||||
},
|
||||
},
|
||||
"required": [
|
||||
"source",
|
||||
"relationship",
|
||||
"destination",
|
||||
],
|
||||
"additionalProperties": False,
|
||||
},
|
||||
},
|
||||
}
|
||||
@@ -1,97 +0,0 @@
|
||||
UPDATE_GRAPH_PROMPT = """
|
||||
You are an AI expert specializing in graph memory management and optimization. Your task is to analyze existing graph memories alongside new information, and update the relationships in the memory list to ensure the most accurate, current, and coherent representation of knowledge.
|
||||
|
||||
Input:
|
||||
1. Existing Graph Memories: A list of current graph memories, each containing source, target, and relationship information.
|
||||
2. New Graph Memory: Fresh information to be integrated into the existing graph structure.
|
||||
|
||||
Guidelines:
|
||||
1. Identification: Use the source and target as primary identifiers when matching existing memories with new information.
|
||||
2. Conflict Resolution:
|
||||
- If new information contradicts an existing memory:
|
||||
a) For matching source and target but differing content, update the relationship of the existing memory.
|
||||
b) If the new memory provides more recent or accurate information, update the existing memory accordingly.
|
||||
3. Comprehensive Review: Thoroughly examine each existing graph memory against the new information, updating relationships as necessary. Multiple updates may be required.
|
||||
4. Consistency: Maintain a uniform and clear style across all memories. Each entry should be concise yet comprehensive.
|
||||
5. Semantic Coherence: Ensure that updates maintain or improve the overall semantic structure of the graph.
|
||||
6. Temporal Awareness: If timestamps are available, consider the recency of information when making updates.
|
||||
7. Relationship Refinement: Look for opportunities to refine relationship descriptions for greater precision or clarity.
|
||||
8. Redundancy Elimination: Identify and merge any redundant or highly similar relationships that may result from the update.
|
||||
|
||||
Memory Format:
|
||||
source -- RELATIONSHIP -- destination
|
||||
|
||||
Task Details:
|
||||
======= Existing Graph Memories:=======
|
||||
{existing_memories}
|
||||
|
||||
======= New Graph Memory:=======
|
||||
{new_memories}
|
||||
|
||||
Output:
|
||||
Provide a list of update instructions, each specifying the source, target, and the new relationship to be set. Only include memories that require updates.
|
||||
"""
|
||||
|
||||
EXTRACT_RELATIONS_PROMPT = """
|
||||
|
||||
You are an advanced algorithm designed to extract structured information from text to construct knowledge graphs. Your goal is to capture comprehensive and accurate information. Follow these key principles:
|
||||
|
||||
1. Extract only explicitly stated information from the text.
|
||||
2. Establish relationships among the entities provided.
|
||||
3. Use "USER_ID" as the source entity for any self-references (e.g., "I," "me," "my," etc.) in user messages.
|
||||
CUSTOM_PROMPT
|
||||
|
||||
Relationships:
|
||||
- Use consistent, general, and timeless relationship types.
|
||||
- Example: Prefer "professor" over "became_professor."
|
||||
- Relationships should only be established among the entities explicitly mentioned in the user message.
|
||||
|
||||
Entity Consistency:
|
||||
- Ensure that relationships are coherent and logically align with the context of the message.
|
||||
- Maintain consistent naming for entities across the extracted data.
|
||||
|
||||
Strive to construct a coherent and easily understandable knowledge graph by establishing all the relationships among the entities and adherence to the user’s context.
|
||||
|
||||
Adhere strictly to these guidelines to ensure high-quality knowledge graph extraction."""
|
||||
|
||||
DELETE_RELATIONS_SYSTEM_PROMPT = """
|
||||
You are a graph memory manager specializing in identifying, managing, and optimizing relationships within graph-based memories. Your primary task is to analyze a list of existing relationships and determine which ones should be deleted based on the new information provided.
|
||||
Input:
|
||||
1. Existing Graph Memories: A list of current graph memories, each containing source, relationship, and destination information.
|
||||
2. New Text: The new information to be integrated into the existing graph structure.
|
||||
3. Use "USER_ID" as node for any self-references (e.g., "I," "me," "my," etc.) in user messages.
|
||||
|
||||
Guidelines:
|
||||
1. Identification: Use the new information to evaluate existing relationships in the memory graph.
|
||||
2. Deletion Criteria: Delete a relationship only if it meets at least one of these conditions:
|
||||
- Outdated or Inaccurate: The new information is more recent or accurate.
|
||||
- Contradictory: The new information conflicts with or negates the existing information.
|
||||
3. DO NOT DELETE if their is a possibility of same type of relationship but different destination nodes.
|
||||
4. Comprehensive Analysis:
|
||||
- Thoroughly examine each existing relationship against the new information and delete as necessary.
|
||||
- Multiple deletions may be required based on the new information.
|
||||
5. Semantic Integrity:
|
||||
- Ensure that deletions maintain or improve the overall semantic structure of the graph.
|
||||
- Avoid deleting relationships that are NOT contradictory/outdated to the new information.
|
||||
6. Temporal Awareness: Prioritize recency when timestamps are available.
|
||||
7. Necessity Principle: Only DELETE relationships that must be deleted and are contradictory/outdated to the new information to maintain an accurate and coherent memory graph.
|
||||
|
||||
Note: DO NOT DELETE if their is a possibility of same type of relationship but different destination nodes.
|
||||
|
||||
For example:
|
||||
Existing Memory: alice -- loves_to_eat -- pizza
|
||||
New Information: Alice also loves to eat burger.
|
||||
|
||||
Do not delete in the above example because there is a possibility that Alice loves to eat both pizza and burger.
|
||||
|
||||
Memory Format:
|
||||
source -- relationship -- destination
|
||||
|
||||
Provide a list of deletion instructions, each specifying the relationship to be deleted.
|
||||
"""
|
||||
|
||||
|
||||
def get_delete_messages(existing_memories_string, data, user_id):
|
||||
return DELETE_RELATIONS_SYSTEM_PROMPT.replace(
|
||||
"USER_ID", user_id
|
||||
), f"Here are the existing memories: {existing_memories_string} \n\n New Information: {data}"
|
||||
@@ -1,595 +0,0 @@
|
||||
import json
|
||||
import logging
|
||||
import time
|
||||
|
||||
from mem0.memory.utils import format_entities, sanitize_relationship_for_cypher
|
||||
|
||||
try:
|
||||
import age
|
||||
except ImportError:
|
||||
raise ImportError("apache-age-python is not installed. Please install it using pip install apache-age-python")
|
||||
|
||||
try:
|
||||
from rank_bm25 import BM25Okapi
|
||||
except ImportError:
|
||||
raise ImportError("rank_bm25 is not installed. Please install it using pip install rank-bm25")
|
||||
|
||||
from mem0.graphs.tools import (
|
||||
DELETE_MEMORY_STRUCT_TOOL_GRAPH,
|
||||
DELETE_MEMORY_TOOL_GRAPH,
|
||||
EXTRACT_ENTITIES_STRUCT_TOOL,
|
||||
EXTRACT_ENTITIES_TOOL,
|
||||
RELATIONS_STRUCT_TOOL,
|
||||
RELATIONS_TOOL,
|
||||
)
|
||||
from mem0.graphs.utils import EXTRACT_RELATIONS_PROMPT, get_delete_messages
|
||||
from mem0.utils.factory import EmbedderFactory, LlmFactory
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def _cosine_similarity(vec1, vec2):
|
||||
"""Compute cosine similarity between two vectors without numpy."""
|
||||
dot = sum(a * b for a, b in zip(vec1, vec2))
|
||||
norm1 = sum(a * a for a in vec1) ** 0.5
|
||||
norm2 = sum(b * b for b in vec2) ** 0.5
|
||||
if norm1 == 0 or norm2 == 0:
|
||||
return 0.0
|
||||
return dot / (norm1 * norm2)
|
||||
|
||||
|
||||
def _get_similar_nodes(nodes, query_embedding, filters, threshold):
|
||||
"""Find nodes above the similarity threshold from a fetched node list.
|
||||
|
||||
Shared by ``_find_similar_node`` and ``_search_graph_db`` to avoid
|
||||
duplicating the client-side cosine similarity logic.
|
||||
"""
|
||||
matches = []
|
||||
for node in nodes:
|
||||
props = node if isinstance(node, dict) else {}
|
||||
stored_emb = props.get("embedding")
|
||||
if not stored_emb:
|
||||
continue
|
||||
if isinstance(stored_emb, str):
|
||||
stored_emb = json.loads(stored_emb)
|
||||
|
||||
if filters.get("agent_id") and props.get("agent_id") != filters["agent_id"]:
|
||||
continue
|
||||
if filters.get("run_id") and props.get("run_id") != filters["run_id"]:
|
||||
continue
|
||||
|
||||
sim = _cosine_similarity(query_embedding, stored_emb)
|
||||
if sim >= threshold:
|
||||
matches.append({"name": props.get("name"), "similarity": sim, "props": props})
|
||||
|
||||
matches.sort(key=lambda x: x["similarity"], reverse=True)
|
||||
return matches
|
||||
|
||||
|
||||
class MemoryGraph:
|
||||
def __init__(self, config):
|
||||
self.config = config
|
||||
|
||||
graph_cfg = self.config.graph_store.config
|
||||
self.graph_name = graph_cfg.graph_name
|
||||
|
||||
# Connect using the Apache AGE Python driver (psycopg2-based)
|
||||
self.ag = age.connect(
|
||||
graph=self.graph_name,
|
||||
host=graph_cfg.host,
|
||||
port=graph_cfg.port,
|
||||
dbname=graph_cfg.database,
|
||||
user=graph_cfg.username,
|
||||
password=graph_cfg.password,
|
||||
)
|
||||
|
||||
self.embedding_model = EmbedderFactory.create(
|
||||
self.config.embedder.provider, self.config.embedder.config, self.config.vector_store.config
|
||||
)
|
||||
|
||||
# Default to openai if no specific provider is configured
|
||||
self.llm_provider = "openai"
|
||||
if self.config.llm and self.config.llm.provider:
|
||||
self.llm_provider = self.config.llm.provider
|
||||
if self.config.graph_store and self.config.graph_store.llm and self.config.graph_store.llm.provider:
|
||||
self.llm_provider = self.config.graph_store.llm.provider
|
||||
|
||||
# Get LLM config with proper null checks
|
||||
llm_config = None
|
||||
if self.config.graph_store and self.config.graph_store.llm and hasattr(self.config.graph_store.llm, "config"):
|
||||
llm_config = self.config.graph_store.llm.config
|
||||
elif hasattr(self.config.llm, "config"):
|
||||
llm_config = self.config.llm.config
|
||||
self.llm = LlmFactory.create(self.llm_provider, llm_config)
|
||||
self.user_id = None
|
||||
self.threshold = self.config.graph_store.threshold if hasattr(self.config.graph_store, "threshold") else 0.7
|
||||
|
||||
# -- helpers ---------------------------------------------------------------
|
||||
|
||||
def _exec_cypher(self, cypher_stmt, cols=None, params=None):
|
||||
"""Execute a Cypher query via the AGE driver and return results.
|
||||
|
||||
Uses ``ag.execCypher`` which delegates to psycopg2's safe parameter
|
||||
substitution (``%s`` placeholders). The *cols* argument specifies the
|
||||
column names in the ``AS (…)`` clause — when ``None`` the driver
|
||||
defaults to a single ``v agtype`` column.
|
||||
|
||||
When *cols* are provided, returns a list of dicts keyed by column name.
|
||||
When *cols* is ``None``, returns vertex/edge property dicts or raw values.
|
||||
"""
|
||||
cursor = self.ag.execCypher(cypher_stmt, cols=cols, params=params)
|
||||
rows = cursor.fetchall()
|
||||
if not rows:
|
||||
return []
|
||||
|
||||
col_names = [desc[0] for desc in cursor.description] if cursor.description else None
|
||||
results = []
|
||||
for row in rows:
|
||||
if col_names and len(col_names) > 1:
|
||||
record = {}
|
||||
for i, col_name in enumerate(col_names):
|
||||
val = row[i]
|
||||
if hasattr(val, "properties"):
|
||||
record[col_name] = val.properties
|
||||
else:
|
||||
record[col_name] = val
|
||||
results.append(record)
|
||||
else:
|
||||
val = row[0] if len(row) == 1 else row
|
||||
if hasattr(val, "properties"):
|
||||
results.append(val.properties)
|
||||
else:
|
||||
results.append(val)
|
||||
return results
|
||||
|
||||
def _fetch_user_nodes_with_embeddings(self, user_id):
|
||||
"""Fetch all nodes with embeddings for a given user_id."""
|
||||
return self._exec_cypher(
|
||||
"MATCH (n {user_id: %s}) WHERE n.embedding IS NOT NULL RETURN n",
|
||||
params=(user_id,),
|
||||
)
|
||||
|
||||
def _find_similar_node(self, embedding, filters, threshold=0.9):
|
||||
"""Find the most similar existing node by cosine similarity.
|
||||
|
||||
Apache AGE does not have a built-in vector index, so we fetch all
|
||||
node embeddings matching the filters and compute cosine similarity
|
||||
on the client side. This is adequate for moderate graph sizes; for
|
||||
very large graphs consider pairing AGE with pgvector.
|
||||
"""
|
||||
nodes = self._fetch_user_nodes_with_embeddings(filters["user_id"])
|
||||
matches = _get_similar_nodes(nodes, embedding, filters, threshold)
|
||||
return matches[0]["props"] if matches else None
|
||||
|
||||
def _merge_node(self, user_id, name, embedding, agent_id=None, run_id=None):
|
||||
"""Create a node if it doesn't exist, or update mentions if it does.
|
||||
|
||||
Apache AGE does not support ``ON CREATE SET`` / ``ON MATCH SET``, so
|
||||
we use ``MERGE … SET`` which always applies the SET clause. Embeddings
|
||||
and optional filter properties are set in a single query.
|
||||
"""
|
||||
set_parts = [
|
||||
"n.embedding = %s",
|
||||
"n.mentions = coalesce(n.mentions, 0) + 1",
|
||||
"n.created = coalesce(n.created, %s)",
|
||||
]
|
||||
params = [user_id, name, json.dumps(embedding), int(time.time() * 1000)]
|
||||
|
||||
if agent_id:
|
||||
set_parts.append("n.agent_id = %s")
|
||||
params.append(agent_id)
|
||||
if run_id:
|
||||
set_parts.append("n.run_id = %s")
|
||||
params.append(run_id)
|
||||
|
||||
set_clause = ", ".join(set_parts)
|
||||
self._exec_cypher(
|
||||
f"MERGE (n {{user_id: %s, name: %s}}) SET {set_clause}",
|
||||
params=tuple(params),
|
||||
)
|
||||
|
||||
def close(self):
|
||||
"""Close the underlying database connection."""
|
||||
if self.ag:
|
||||
self.ag.close()
|
||||
|
||||
# -- public API ------------------------------------------------------------
|
||||
|
||||
def add(self, data, filters):
|
||||
"""
|
||||
Adds data to the graph.
|
||||
|
||||
Args:
|
||||
data (str): The data to add to the graph.
|
||||
filters (dict): A dictionary containing filters to be applied during the addition.
|
||||
"""
|
||||
entity_type_map = self._retrieve_nodes_from_data(data, filters)
|
||||
to_be_added = self._establish_nodes_relations_from_data(data, filters, entity_type_map)
|
||||
search_output = self._search_graph_db(node_list=list(entity_type_map.keys()), filters=filters)
|
||||
to_be_deleted = self._get_delete_entities_from_search_output(search_output, data, filters)
|
||||
|
||||
deleted_entities = self._delete_entities(to_be_deleted, filters)
|
||||
added_entities = self._add_entities(to_be_added, filters, entity_type_map)
|
||||
|
||||
return {"deleted_entities": deleted_entities, "added_entities": added_entities}
|
||||
|
||||
def search(self, query, filters, top_k=100):
|
||||
"""
|
||||
Search for memories and related graph data.
|
||||
|
||||
Args:
|
||||
query (str): Query to search for.
|
||||
filters (dict): A dictionary containing filters to be applied during the search.
|
||||
top_k (int): The maximum number of nodes and relationships to retrieve. Defaults to 100.
|
||||
|
||||
Returns:
|
||||
list: A list of dicts with keys "source", "relationship", "destination".
|
||||
"""
|
||||
entity_type_map = self._retrieve_nodes_from_data(query, filters)
|
||||
search_output = self._search_graph_db(node_list=list(entity_type_map.keys()), filters=filters)
|
||||
|
||||
if not search_output:
|
||||
return []
|
||||
|
||||
search_outputs_sequence = [
|
||||
[item["source"], item["relationship"], item["destination"]] for item in search_output
|
||||
]
|
||||
bm25 = BM25Okapi(search_outputs_sequence)
|
||||
|
||||
tokenized_query = query.split(" ")
|
||||
reranked_results = bm25.get_top_n(tokenized_query, search_outputs_sequence, n=5)
|
||||
|
||||
search_results = []
|
||||
for item in reranked_results:
|
||||
search_results.append({"source": item[0], "relationship": item[1], "destination": item[2]})
|
||||
|
||||
logger.info(f"Returned {len(search_results)} search results")
|
||||
return search_results
|
||||
|
||||
def delete(self, data, filters):
|
||||
"""
|
||||
Delete graph entities associated with the given memory text.
|
||||
|
||||
Extracts entities and relationships from the memory text using the same
|
||||
pipeline as add(), then deletes the matching relationships in the graph.
|
||||
|
||||
Args:
|
||||
data (str): The memory text whose graph entities should be removed.
|
||||
filters (dict): Scope filters (user_id, agent_id, run_id).
|
||||
"""
|
||||
try:
|
||||
entity_type_map = self._retrieve_nodes_from_data(data, filters)
|
||||
if not entity_type_map:
|
||||
logger.debug("No entities found in memory text, skipping graph cleanup")
|
||||
return
|
||||
to_be_deleted = self._establish_nodes_relations_from_data(data, filters, entity_type_map)
|
||||
if to_be_deleted:
|
||||
self._delete_entities(to_be_deleted, filters)
|
||||
except Exception as e:
|
||||
logger.error(f"Error during graph cleanup for memory delete: {e}")
|
||||
|
||||
def delete_all(self, filters):
|
||||
"""Delete all nodes and relationships for a user or specific agent."""
|
||||
where_parts = ["n.user_id = %s"]
|
||||
params = [filters["user_id"]]
|
||||
if filters.get("agent_id"):
|
||||
where_parts.append("n.agent_id = %s")
|
||||
params.append(filters["agent_id"])
|
||||
if filters.get("run_id"):
|
||||
where_parts.append("n.run_id = %s")
|
||||
params.append(filters["run_id"])
|
||||
where_clause = " AND ".join(where_parts)
|
||||
|
||||
self._exec_cypher(
|
||||
f"MATCH (n) WHERE {where_clause} DETACH DELETE n",
|
||||
params=tuple(params),
|
||||
)
|
||||
self.ag.commit()
|
||||
|
||||
def get_all(self, filters, top_k=100):
|
||||
"""
|
||||
Retrieves all nodes and relationships from the graph database based on optional filtering criteria.
|
||||
|
||||
Args:
|
||||
filters (dict): A dictionary containing filters to be applied during the retrieval.
|
||||
top_k (int): The maximum number of nodes and relationships to retrieve. Defaults to 100.
|
||||
Returns:
|
||||
list: A list of dictionaries, each containing:
|
||||
- 'source': The source node name.
|
||||
- 'relationship': The relationship type.
|
||||
- 'target': The target node name.
|
||||
"""
|
||||
where_parts = ["n.user_id = %s", "m.user_id = %s"]
|
||||
params = [filters["user_id"], filters["user_id"]]
|
||||
if filters.get("agent_id"):
|
||||
where_parts.extend(["n.agent_id = %s", "m.agent_id = %s"])
|
||||
params.extend([filters["agent_id"], filters["agent_id"]])
|
||||
if filters.get("run_id"):
|
||||
where_parts.extend(["n.run_id = %s", "m.run_id = %s"])
|
||||
params.extend([filters["run_id"], filters["run_id"]])
|
||||
where_clause = " AND ".join(where_parts)
|
||||
params.append(top_k)
|
||||
|
||||
results = self._exec_cypher(
|
||||
f"MATCH (n)-[r]->(m) WHERE {where_clause} "
|
||||
f"RETURN n.name, type(r), m.name LIMIT %s",
|
||||
cols=["source", "relationship", "target"],
|
||||
params=tuple(params),
|
||||
)
|
||||
|
||||
final_results = []
|
||||
for result in results:
|
||||
final_results.append(
|
||||
{
|
||||
"source": result["source"],
|
||||
"relationship": result["relationship"],
|
||||
"target": result["target"],
|
||||
}
|
||||
)
|
||||
|
||||
logger.info(f"Retrieved {len(final_results)} relationships")
|
||||
return final_results
|
||||
|
||||
# -- LLM-driven extraction -------------------------------------------------
|
||||
|
||||
def _retrieve_nodes_from_data(self, data, filters):
|
||||
"""Extracts all the entities mentioned in the query."""
|
||||
_tools = [EXTRACT_ENTITIES_TOOL]
|
||||
if self.llm_provider in ["azure_openai_structured", "openai_structured"]:
|
||||
_tools = [EXTRACT_ENTITIES_STRUCT_TOOL]
|
||||
search_results = self.llm.generate_response(
|
||||
messages=[
|
||||
{
|
||||
"role": "system",
|
||||
"content": f"You are a smart assistant who understands entities and their types in a given text. If user message contains self reference such as 'I', 'me', 'my' etc. then use {filters['user_id']} as the source entity. Extract all the entities from the text. ***DO NOT*** answer the question itself if the given text is a question.",
|
||||
},
|
||||
{"role": "user", "content": data},
|
||||
],
|
||||
tools=_tools,
|
||||
)
|
||||
|
||||
entity_type_map = {}
|
||||
|
||||
try:
|
||||
for tool_call in search_results["tool_calls"]:
|
||||
if tool_call["name"] != "extract_entities":
|
||||
continue
|
||||
for item in tool_call.get("arguments", {}).get("entities", []):
|
||||
if "entity" in item and "entity_type" in item:
|
||||
entity_type_map[item["entity"]] = item["entity_type"]
|
||||
except Exception as e:
|
||||
logger.exception(
|
||||
f"Error in search tool: {e}, llm_provider={self.llm_provider}, search_results={search_results}"
|
||||
)
|
||||
|
||||
entity_type_map = {k.lower().replace(" ", "_"): v.lower().replace(" ", "_") for k, v in entity_type_map.items()}
|
||||
logger.debug(f"Entity type map: {entity_type_map}\n search_results={search_results}")
|
||||
return entity_type_map
|
||||
|
||||
def _establish_nodes_relations_from_data(self, data, filters, entity_type_map):
|
||||
"""Establish relations among the extracted nodes."""
|
||||
|
||||
user_identity = f"user_id: {filters['user_id']}"
|
||||
if filters.get("agent_id"):
|
||||
user_identity += f", agent_id: {filters['agent_id']}"
|
||||
if filters.get("run_id"):
|
||||
user_identity += f", run_id: {filters['run_id']}"
|
||||
|
||||
if self.config.graph_store.custom_prompt:
|
||||
system_content = EXTRACT_RELATIONS_PROMPT.replace("USER_ID", user_identity)
|
||||
system_content = system_content.replace("CUSTOM_PROMPT", f"4. {self.config.graph_store.custom_prompt}")
|
||||
messages = [
|
||||
{"role": "system", "content": system_content},
|
||||
{"role": "user", "content": data},
|
||||
]
|
||||
else:
|
||||
system_content = EXTRACT_RELATIONS_PROMPT.replace("USER_ID", user_identity)
|
||||
messages = [
|
||||
{"role": "system", "content": system_content},
|
||||
{"role": "user", "content": f"List of entities: {list(entity_type_map.keys())}. \n\nText: {data}"},
|
||||
]
|
||||
|
||||
_tools = [RELATIONS_TOOL]
|
||||
if self.llm_provider in ["azure_openai_structured", "openai_structured"]:
|
||||
_tools = [RELATIONS_STRUCT_TOOL]
|
||||
|
||||
extracted_entities = self.llm.generate_response(
|
||||
messages=messages,
|
||||
tools=_tools,
|
||||
)
|
||||
|
||||
entities = []
|
||||
if extracted_entities and extracted_entities.get("tool_calls"):
|
||||
entities = extracted_entities["tool_calls"][0].get("arguments", {}).get("entities", [])
|
||||
|
||||
entities = self._remove_spaces_from_entities(entities)
|
||||
logger.debug(f"Extracted entities: {entities}")
|
||||
return entities
|
||||
|
||||
# -- graph DB operations ---------------------------------------------------
|
||||
|
||||
def _search_graph_db(self, node_list, filters, top_k=100):
|
||||
"""Search similar nodes and their respective incoming and outgoing relations."""
|
||||
result_relations = []
|
||||
|
||||
for node in node_list:
|
||||
n_embedding = self.embedding_model.embed(node)
|
||||
|
||||
nodes = self._fetch_user_nodes_with_embeddings(filters["user_id"])
|
||||
similar_nodes = _get_similar_nodes(nodes, n_embedding, filters, self.threshold)
|
||||
|
||||
# Build WHERE clause for relationship target filtering
|
||||
rel_where_parts = ["m.user_id = %s"]
|
||||
rel_params_suffix = [filters["user_id"]]
|
||||
if filters.get("agent_id"):
|
||||
rel_where_parts.append("m.agent_id = %s")
|
||||
rel_params_suffix.append(filters["agent_id"])
|
||||
if filters.get("run_id"):
|
||||
rel_where_parts.append("m.run_id = %s")
|
||||
rel_params_suffix.append(filters["run_id"])
|
||||
rel_where = " AND ".join(rel_where_parts)
|
||||
|
||||
# For each similar node, fetch its relationships
|
||||
for sn in similar_nodes[:top_k]:
|
||||
node_name = sn["name"]
|
||||
similarity = sn["similarity"]
|
||||
|
||||
out_params = (filters["user_id"], node_name) + tuple(rel_params_suffix)
|
||||
out_results = self._exec_cypher(
|
||||
f"MATCH (n {{user_id: %s, name: %s}})-[r]->(m) "
|
||||
f"WHERE {rel_where} "
|
||||
f"RETURN n.name, type(r), m.name",
|
||||
cols=["source", "relationship", "destination"],
|
||||
params=out_params,
|
||||
)
|
||||
|
||||
in_params = (filters["user_id"], node_name) + tuple(rel_params_suffix)
|
||||
in_results = self._exec_cypher(
|
||||
f"MATCH (n {{user_id: %s, name: %s}})<-[r]-(m) "
|
||||
f"WHERE {rel_where} "
|
||||
f"RETURN m.name, type(r), n.name",
|
||||
cols=["source", "relationship", "destination"],
|
||||
params=in_params,
|
||||
)
|
||||
|
||||
for rel in out_results + in_results:
|
||||
rel["similarity"] = similarity
|
||||
result_relations.append(rel)
|
||||
|
||||
return result_relations
|
||||
|
||||
def _get_delete_entities_from_search_output(self, search_output, data, filters):
|
||||
"""Get the entities to be deleted from the search output."""
|
||||
search_output_string = format_entities(search_output)
|
||||
|
||||
user_identity = f"user_id: {filters['user_id']}"
|
||||
if filters.get("agent_id"):
|
||||
user_identity += f", agent_id: {filters['agent_id']}"
|
||||
if filters.get("run_id"):
|
||||
user_identity += f", run_id: {filters['run_id']}"
|
||||
|
||||
system_prompt, user_prompt = get_delete_messages(search_output_string, data, user_identity)
|
||||
|
||||
_tools = [DELETE_MEMORY_TOOL_GRAPH]
|
||||
if self.llm_provider in ["azure_openai_structured", "openai_structured"]:
|
||||
_tools = [DELETE_MEMORY_STRUCT_TOOL_GRAPH]
|
||||
|
||||
memory_updates = self.llm.generate_response(
|
||||
messages=[
|
||||
{"role": "system", "content": system_prompt},
|
||||
{"role": "user", "content": user_prompt},
|
||||
],
|
||||
tools=_tools,
|
||||
)
|
||||
|
||||
to_be_deleted = []
|
||||
for item in memory_updates.get("tool_calls", []):
|
||||
if item.get("name") == "delete_graph_memory":
|
||||
to_be_deleted.append(item.get("arguments"))
|
||||
to_be_deleted = self._remove_spaces_from_entities(to_be_deleted)
|
||||
logger.debug(f"Deleted relationships: {to_be_deleted}")
|
||||
return to_be_deleted
|
||||
|
||||
def _delete_entities(self, to_be_deleted, filters):
|
||||
"""Delete the entities from the graph."""
|
||||
user_id = filters["user_id"]
|
||||
agent_id = filters.get("agent_id")
|
||||
run_id = filters.get("run_id")
|
||||
results = []
|
||||
|
||||
try:
|
||||
for item in to_be_deleted:
|
||||
source = item["source"]
|
||||
destination = item["destination"]
|
||||
relationship = item["relationship"]
|
||||
|
||||
where_parts = [
|
||||
"n.user_id = %s", "n.name = %s",
|
||||
"m.user_id = %s", "m.name = %s",
|
||||
]
|
||||
params = [user_id, source, user_id, destination]
|
||||
if agent_id:
|
||||
where_parts.extend(["n.agent_id = %s", "m.agent_id = %s"])
|
||||
params.extend([agent_id, agent_id])
|
||||
if run_id:
|
||||
where_parts.extend(["n.run_id = %s", "m.run_id = %s"])
|
||||
params.extend([run_id, run_id])
|
||||
where_clause = " AND ".join(where_parts)
|
||||
|
||||
result = self._exec_cypher(
|
||||
f"MATCH (n)-[r:{relationship}]->(m) "
|
||||
f"WHERE {where_clause} "
|
||||
f"DELETE r "
|
||||
f"RETURN n.name, type(r), m.name",
|
||||
cols=["source", "relationship", "target"],
|
||||
params=tuple(params),
|
||||
)
|
||||
results.append(result)
|
||||
|
||||
self.ag.commit()
|
||||
except Exception:
|
||||
self.ag.rollback()
|
||||
raise
|
||||
|
||||
return results
|
||||
|
||||
def _add_entities(self, to_be_added, filters, entity_type_map):
|
||||
"""Add new entities to the graph. Merge nodes if they already exist.
|
||||
|
||||
Apache AGE does not support ``ON CREATE SET`` / ``ON MATCH SET``, so we
|
||||
use ``MERGE … SET`` which always applies. The ``coalesce`` pattern
|
||||
ensures ``created`` is only set on the first merge.
|
||||
"""
|
||||
user_id = filters["user_id"]
|
||||
agent_id = filters.get("agent_id")
|
||||
run_id = filters.get("run_id")
|
||||
results = []
|
||||
|
||||
try:
|
||||
for item in to_be_added:
|
||||
source = item["source"]
|
||||
destination = item["destination"]
|
||||
relationship = item["relationship"]
|
||||
|
||||
source_embedding = self.embedding_model.embed(source)
|
||||
dest_embedding = self.embedding_model.embed(destination)
|
||||
|
||||
source_match = self._find_similar_node(source_embedding, filters, threshold=self.threshold)
|
||||
dest_match = self._find_similar_node(dest_embedding, filters, threshold=self.threshold)
|
||||
|
||||
effective_source = source_match["name"] if source_match else source
|
||||
effective_dest = dest_match["name"] if dest_match else destination
|
||||
|
||||
# Merge source and destination nodes
|
||||
self._merge_node(user_id, effective_source, source_embedding, agent_id, run_id)
|
||||
self._merge_node(user_id, effective_dest, dest_embedding, agent_id, run_id)
|
||||
|
||||
# Merge relationship
|
||||
result = self._exec_cypher(
|
||||
f"MATCH (s {{user_id: %s, name: %s}}), (d {{user_id: %s, name: %s}}) "
|
||||
f"MERGE (s)-[r:{relationship}]->(d) "
|
||||
f"RETURN s.name, type(r), d.name",
|
||||
cols=["source", "relationship", "target"],
|
||||
params=(user_id, effective_source, user_id, effective_dest),
|
||||
)
|
||||
results.append(result)
|
||||
|
||||
self.ag.commit()
|
||||
except Exception:
|
||||
self.ag.rollback()
|
||||
raise
|
||||
|
||||
return results
|
||||
|
||||
def _remove_spaces_from_entities(self, entity_list):
|
||||
for item in entity_list:
|
||||
item["source"] = item["source"].lower().replace(" ", "_")
|
||||
item["relationship"] = sanitize_relationship_for_cypher(item["relationship"].lower().replace(" ", "_"))
|
||||
item["destination"] = item["destination"].lower().replace(" ", "_")
|
||||
return entity_list
|
||||
|
||||
def reset(self):
|
||||
"""Reset the graph by clearing all nodes and relationships."""
|
||||
logger.warning("Clearing graph...")
|
||||
self._exec_cypher("MATCH (n) DETACH DELETE n")
|
||||
self.ag.commit()
|
||||
@@ -1,744 +0,0 @@
|
||||
import logging
|
||||
|
||||
from mem0.memory.utils import format_entities, remove_spaces_from_entities
|
||||
|
||||
try:
|
||||
from langchain_neo4j import Neo4jGraph
|
||||
except ImportError:
|
||||
raise ImportError("langchain_neo4j is not installed. Please install it using pip install langchain-neo4j")
|
||||
|
||||
try:
|
||||
from rank_bm25 import BM25Okapi
|
||||
except ImportError:
|
||||
raise ImportError("rank_bm25 is not installed. Please install it using pip install rank-bm25")
|
||||
|
||||
from mem0.graphs.tools import (
|
||||
DELETE_MEMORY_STRUCT_TOOL_GRAPH,
|
||||
DELETE_MEMORY_TOOL_GRAPH,
|
||||
EXTRACT_ENTITIES_STRUCT_TOOL,
|
||||
EXTRACT_ENTITIES_TOOL,
|
||||
RELATIONS_STRUCT_TOOL,
|
||||
RELATIONS_TOOL,
|
||||
)
|
||||
from mem0.graphs.utils import EXTRACT_RELATIONS_PROMPT, get_delete_messages
|
||||
from mem0.utils.factory import EmbedderFactory, LlmFactory
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
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,
|
||||
refresh_schema=False,
|
||||
driver_config={"notifications_min_severity": "OFF"},
|
||||
)
|
||||
self.embedding_model = EmbedderFactory.create(
|
||||
self.config.embedder.provider, self.config.embedder.config, self.config.vector_store.config
|
||||
)
|
||||
self.node_label = ":`__Entity__`" if self.config.graph_store.config.base_label else ""
|
||||
|
||||
if self.config.graph_store.config.base_label:
|
||||
# Safely add user_id index
|
||||
try:
|
||||
self.graph.query(f"CREATE INDEX entity_single IF NOT EXISTS FOR (n {self.node_label}) ON (n.user_id)")
|
||||
except Exception:
|
||||
pass
|
||||
try: # Safely try to add composite index (Enterprise only)
|
||||
self.graph.query(
|
||||
f"CREATE INDEX entity_composite IF NOT EXISTS FOR (n {self.node_label}) ON (n.name, n.user_id)"
|
||||
)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
# Default to openai if no specific provider is configured
|
||||
self.llm_provider = "openai"
|
||||
if self.config.llm and self.config.llm.provider:
|
||||
self.llm_provider = self.config.llm.provider
|
||||
if self.config.graph_store and self.config.graph_store.llm and self.config.graph_store.llm.provider:
|
||||
self.llm_provider = self.config.graph_store.llm.provider
|
||||
|
||||
# Get LLM config with proper null checks
|
||||
llm_config = None
|
||||
if self.config.graph_store and self.config.graph_store.llm and hasattr(self.config.graph_store.llm, "config"):
|
||||
llm_config = self.config.graph_store.llm.config
|
||||
elif hasattr(self.config.llm, "config"):
|
||||
llm_config = self.config.llm.config
|
||||
self.llm = LlmFactory.create(self.llm_provider, llm_config)
|
||||
self.user_id = None
|
||||
# Use threshold from graph_store config, default to 0.7 for backward compatibility
|
||||
self.threshold = self.config.graph_store.threshold if hasattr(self.config.graph_store, 'threshold') else 0.7
|
||||
|
||||
def add(self, data, filters):
|
||||
"""
|
||||
Adds data to the graph.
|
||||
|
||||
Args:
|
||||
data (str): The data to add to the graph.
|
||||
filters (dict): A dictionary containing filters to be applied during the addition.
|
||||
"""
|
||||
entity_type_map = self._retrieve_nodes_from_data(data, filters)
|
||||
to_be_added = self._establish_nodes_relations_from_data(data, filters, entity_type_map)
|
||||
search_output = self._search_graph_db(node_list=list(entity_type_map.keys()), filters=filters)
|
||||
to_be_deleted = self._get_delete_entities_from_search_output(search_output, data, filters)
|
||||
|
||||
# TODO: Batch queries with APOC plugin
|
||||
# TODO: Add more filter support
|
||||
deleted_entities = self._delete_entities(to_be_deleted, filters)
|
||||
added_entities = self._add_entities(to_be_added, filters, entity_type_map)
|
||||
|
||||
return {"deleted_entities": deleted_entities, "added_entities": added_entities}
|
||||
|
||||
def search(self, query, filters, top_k=100):
|
||||
"""
|
||||
Search for memories and related graph data.
|
||||
|
||||
Args:
|
||||
query (str): Query to search for.
|
||||
filters (dict): A dictionary containing filters to be applied during the search.
|
||||
top_k (int): The maximum number of nodes and relationships to retrieve. Defaults to 100.
|
||||
|
||||
Returns:
|
||||
dict: A dictionary containing:
|
||||
- "contexts": List of search results from the base data store.
|
||||
- "entities": List of related graph data based on the query.
|
||||
"""
|
||||
entity_type_map = self._retrieve_nodes_from_data(query, filters)
|
||||
search_output = self._search_graph_db(node_list=list(entity_type_map.keys()), filters=filters)
|
||||
|
||||
if not search_output:
|
||||
return []
|
||||
|
||||
search_outputs_sequence = [
|
||||
[item["source"], item["relationship"], item["destination"]] for item in search_output
|
||||
]
|
||||
bm25 = BM25Okapi(search_outputs_sequence)
|
||||
|
||||
tokenized_query = query.split(" ")
|
||||
reranked_results = bm25.get_top_n(tokenized_query, search_outputs_sequence, n=5)
|
||||
|
||||
search_results = []
|
||||
for item in reranked_results:
|
||||
search_results.append({"source": item[0], "relationship": item[1], "destination": item[2]})
|
||||
|
||||
logger.info(f"Returned {len(search_results)} search results")
|
||||
|
||||
return search_results
|
||||
|
||||
def delete(self, data, filters):
|
||||
"""
|
||||
Delete graph entities associated with the given memory text.
|
||||
|
||||
Extracts entities and relationships from the memory text using the same
|
||||
pipeline as add(), then soft-deletes the matching relationships in the graph.
|
||||
|
||||
Args:
|
||||
data (str): The memory text whose graph entities should be removed.
|
||||
filters (dict): Scope filters (user_id, agent_id, run_id).
|
||||
"""
|
||||
try:
|
||||
entity_type_map = self._retrieve_nodes_from_data(data, filters)
|
||||
if not entity_type_map:
|
||||
logger.debug("No entities found in memory text, skipping graph cleanup")
|
||||
return
|
||||
to_be_deleted = self._establish_nodes_relations_from_data(data, filters, entity_type_map)
|
||||
if to_be_deleted:
|
||||
self._delete_entities(to_be_deleted, filters)
|
||||
except Exception as e:
|
||||
logger.error(f"Error during graph cleanup for memory delete: {e}")
|
||||
|
||||
def delete_all(self, filters):
|
||||
# Build node properties for filtering
|
||||
node_props = ["user_id: $user_id"]
|
||||
if filters.get("agent_id"):
|
||||
node_props.append("agent_id: $agent_id")
|
||||
if filters.get("run_id"):
|
||||
node_props.append("run_id: $run_id")
|
||||
node_props_str = ", ".join(node_props)
|
||||
|
||||
cypher = f"""
|
||||
MATCH (n {self.node_label} {{{node_props_str}}})
|
||||
DETACH DELETE n
|
||||
"""
|
||||
params = {"user_id": filters["user_id"]}
|
||||
if filters.get("agent_id"):
|
||||
params["agent_id"] = filters["agent_id"]
|
||||
if filters.get("run_id"):
|
||||
params["run_id"] = filters["run_id"]
|
||||
self.graph.query(cypher, params=params)
|
||||
|
||||
def get_all(self, filters, top_k=100):
|
||||
"""
|
||||
Retrieves all nodes and relationships from the graph database based on optional filtering criteria.
|
||||
Args:
|
||||
filters (dict): A dictionary containing filters to be applied during the retrieval.
|
||||
top_k (int): The maximum number of nodes and relationships to retrieve. Defaults to 100.
|
||||
Returns:
|
||||
list: A list of dictionaries, each containing:
|
||||
- 'contexts': The base data store response for each memory.
|
||||
- 'entities': A list of strings representing the nodes and relationships
|
||||
"""
|
||||
params = {"user_id": filters["user_id"], "limit": top_k}
|
||||
|
||||
# Build node properties based on filters
|
||||
node_props = ["user_id: $user_id"]
|
||||
if filters.get("agent_id"):
|
||||
node_props.append("agent_id: $agent_id")
|
||||
params["agent_id"] = filters["agent_id"]
|
||||
if filters.get("run_id"):
|
||||
node_props.append("run_id: $run_id")
|
||||
params["run_id"] = filters["run_id"]
|
||||
node_props_str = ", ".join(node_props)
|
||||
|
||||
query = f"""
|
||||
MATCH (n {self.node_label} {{{node_props_str}}})-[r]->(m {self.node_label} {{{node_props_str}}})
|
||||
WHERE r.valid IS NULL OR r.valid = true
|
||||
RETURN n.name AS source, type(r) AS relationship, m.name AS target
|
||||
LIMIT $limit
|
||||
"""
|
||||
results = self.graph.query(query, params=params)
|
||||
|
||||
final_results = []
|
||||
for result in results:
|
||||
final_results.append(
|
||||
{
|
||||
"source": result["source"],
|
||||
"relationship": result["relationship"],
|
||||
"target": result["target"],
|
||||
}
|
||||
)
|
||||
|
||||
logger.info(f"Retrieved {len(final_results)} relationships")
|
||||
|
||||
return final_results
|
||||
|
||||
def _retrieve_nodes_from_data(self, data, filters):
|
||||
"""Extracts all the entities mentioned in the query."""
|
||||
_tools = [EXTRACT_ENTITIES_TOOL]
|
||||
if self.llm_provider in ["azure_openai_structured", "openai_structured"]:
|
||||
_tools = [EXTRACT_ENTITIES_STRUCT_TOOL]
|
||||
search_results = self.llm.generate_response(
|
||||
messages=[
|
||||
{
|
||||
"role": "system",
|
||||
"content": f"You are a smart assistant who understands entities and their types in a given text. If user message contains self reference such as 'I', 'me', 'my' etc. then use {filters['user_id']} as the source entity. Extract all the entities from the text. ***DO NOT*** answer the question itself if the given text is a question.",
|
||||
},
|
||||
{"role": "user", "content": data},
|
||||
],
|
||||
tools=_tools,
|
||||
)
|
||||
|
||||
entity_type_map = {}
|
||||
|
||||
try:
|
||||
for tool_call in search_results["tool_calls"]:
|
||||
if tool_call["name"] != "extract_entities":
|
||||
continue
|
||||
for item in tool_call.get("arguments", {}).get("entities", []):
|
||||
entity_type_map[item["entity"]] = item["entity_type"]
|
||||
except Exception as e:
|
||||
logger.exception(
|
||||
f"Error in search tool: {e}, llm_provider={self.llm_provider}, search_results={search_results}"
|
||||
)
|
||||
|
||||
entity_type_map = {k.lower().replace(" ", "_"): v.lower().replace(" ", "_") for k, v in entity_type_map.items()}
|
||||
logger.debug(f"Entity type map: {entity_type_map}\n search_results={search_results}")
|
||||
return entity_type_map
|
||||
|
||||
def _establish_nodes_relations_from_data(self, data, filters, entity_type_map):
|
||||
"""Establish relations among the extracted nodes."""
|
||||
|
||||
# Compose user identification string for prompt
|
||||
user_identity = f"user_id: {filters['user_id']}"
|
||||
if filters.get("agent_id"):
|
||||
user_identity += f", agent_id: {filters['agent_id']}"
|
||||
if filters.get("run_id"):
|
||||
user_identity += f", run_id: {filters['run_id']}"
|
||||
|
||||
if self.config.graph_store.custom_prompt:
|
||||
system_content = EXTRACT_RELATIONS_PROMPT.replace("USER_ID", user_identity)
|
||||
# Add the custom prompt line if configured
|
||||
system_content = system_content.replace("CUSTOM_PROMPT", f"4. {self.config.graph_store.custom_prompt}")
|
||||
messages = [
|
||||
{"role": "system", "content": system_content},
|
||||
{"role": "user", "content": data},
|
||||
]
|
||||
else:
|
||||
system_content = EXTRACT_RELATIONS_PROMPT.replace("USER_ID", user_identity)
|
||||
messages = [
|
||||
{"role": "system", "content": system_content},
|
||||
{"role": "user", "content": f"List of entities: {list(entity_type_map.keys())}. \n\nText: {data}"},
|
||||
]
|
||||
|
||||
_tools = [RELATIONS_TOOL]
|
||||
if self.llm_provider in ["azure_openai_structured", "openai_structured"]:
|
||||
_tools = [RELATIONS_STRUCT_TOOL]
|
||||
|
||||
extracted_entities = self.llm.generate_response(
|
||||
messages=messages,
|
||||
tools=_tools,
|
||||
)
|
||||
|
||||
entities = []
|
||||
if extracted_entities.get("tool_calls"):
|
||||
entities = extracted_entities["tool_calls"][0].get("arguments", {}).get("entities", [])
|
||||
|
||||
entities = self._remove_spaces_from_entities(entities)
|
||||
logger.debug(f"Extracted entities: {entities}")
|
||||
return entities
|
||||
|
||||
def _search_graph_db(self, node_list, filters, top_k=100):
|
||||
"""Search similar nodes among and their respective incoming and outgoing relations."""
|
||||
result_relations = []
|
||||
|
||||
# Build node properties for filtering
|
||||
node_props = ["user_id: $user_id"]
|
||||
if filters.get("agent_id"):
|
||||
node_props.append("agent_id: $agent_id")
|
||||
if filters.get("run_id"):
|
||||
node_props.append("run_id: $run_id")
|
||||
node_props_str = ", ".join(node_props)
|
||||
|
||||
for node in node_list:
|
||||
n_embedding = self.embedding_model.embed(node)
|
||||
|
||||
cypher_query = f"""
|
||||
MATCH (n {self.node_label} {{{node_props_str}}})
|
||||
WHERE n.embedding IS NOT NULL
|
||||
WITH n, round(2 * vector.similarity.cosine(n.embedding, $n_embedding) - 1, 4) AS similarity // denormalize for backward compatibility
|
||||
WHERE similarity >= $threshold
|
||||
CALL {{
|
||||
WITH n
|
||||
MATCH (n)-[r]->(m {self.node_label} {{{node_props_str}}})
|
||||
WHERE r.valid IS NULL OR r.valid = true
|
||||
RETURN n.name AS source, elementId(n) AS source_id, type(r) AS relationship, elementId(r) AS relation_id, m.name AS destination, elementId(m) AS destination_id
|
||||
UNION
|
||||
WITH n
|
||||
MATCH (n)<-[r]-(m {self.node_label} {{{node_props_str}}})
|
||||
WHERE r.valid IS NULL OR r.valid = true
|
||||
RETURN m.name AS source, elementId(m) AS source_id, type(r) AS relationship, elementId(r) AS relation_id, n.name AS destination, elementId(n) AS destination_id
|
||||
}}
|
||||
WITH distinct source, source_id, relationship, relation_id, destination, destination_id, similarity
|
||||
RETURN source, source_id, relationship, relation_id, destination, destination_id, similarity
|
||||
ORDER BY similarity DESC
|
||||
LIMIT $limit
|
||||
"""
|
||||
|
||||
params = {
|
||||
"n_embedding": n_embedding,
|
||||
"threshold": self.threshold,
|
||||
"user_id": filters["user_id"],
|
||||
"limit": top_k,
|
||||
}
|
||||
if filters.get("agent_id"):
|
||||
params["agent_id"] = filters["agent_id"]
|
||||
if filters.get("run_id"):
|
||||
params["run_id"] = filters["run_id"]
|
||||
|
||||
ans = self.graph.query(cypher_query, params=params)
|
||||
result_relations.extend(ans)
|
||||
|
||||
return result_relations
|
||||
|
||||
def _get_delete_entities_from_search_output(self, search_output, data, filters):
|
||||
"""Get the entities to be deleted from the search output."""
|
||||
search_output_string = format_entities(search_output)
|
||||
|
||||
# Compose user identification string for prompt
|
||||
user_identity = f"user_id: {filters['user_id']}"
|
||||
if filters.get("agent_id"):
|
||||
user_identity += f", agent_id: {filters['agent_id']}"
|
||||
if filters.get("run_id"):
|
||||
user_identity += f", run_id: {filters['run_id']}"
|
||||
|
||||
system_prompt, user_prompt = get_delete_messages(search_output_string, data, user_identity)
|
||||
|
||||
_tools = [DELETE_MEMORY_TOOL_GRAPH]
|
||||
if self.llm_provider in ["azure_openai_structured", "openai_structured"]:
|
||||
_tools = [
|
||||
DELETE_MEMORY_STRUCT_TOOL_GRAPH,
|
||||
]
|
||||
|
||||
memory_updates = self.llm.generate_response(
|
||||
messages=[
|
||||
{"role": "system", "content": system_prompt},
|
||||
{"role": "user", "content": user_prompt},
|
||||
],
|
||||
tools=_tools,
|
||||
)
|
||||
|
||||
to_be_deleted = []
|
||||
for item in memory_updates.get("tool_calls", []):
|
||||
if item.get("name") == "delete_graph_memory":
|
||||
to_be_deleted.append(item.get("arguments"))
|
||||
# Clean entities formatting
|
||||
to_be_deleted = self._remove_spaces_from_entities(to_be_deleted)
|
||||
logger.debug(f"Deleted relationships: {to_be_deleted}")
|
||||
return to_be_deleted
|
||||
|
||||
def _delete_entities(self, to_be_deleted, filters):
|
||||
"""Delete the entities from the graph."""
|
||||
user_id = filters["user_id"]
|
||||
agent_id = filters.get("agent_id", None)
|
||||
run_id = filters.get("run_id", None)
|
||||
results = []
|
||||
|
||||
for item in to_be_deleted:
|
||||
source = item["source"]
|
||||
destination = item["destination"]
|
||||
relationship = item["relationship"]
|
||||
|
||||
# Build the agent filter for the query
|
||||
|
||||
params = {
|
||||
"source_name": source,
|
||||
"dest_name": destination,
|
||||
"user_id": user_id,
|
||||
}
|
||||
|
||||
if agent_id:
|
||||
params["agent_id"] = agent_id
|
||||
if run_id:
|
||||
params["run_id"] = run_id
|
||||
|
||||
# Build node properties for filtering
|
||||
source_props = ["name: $source_name", "user_id: $user_id"]
|
||||
dest_props = ["name: $dest_name", "user_id: $user_id"]
|
||||
if agent_id:
|
||||
source_props.append("agent_id: $agent_id")
|
||||
dest_props.append("agent_id: $agent_id")
|
||||
if run_id:
|
||||
source_props.append("run_id: $run_id")
|
||||
dest_props.append("run_id: $run_id")
|
||||
source_props_str = ", ".join(source_props)
|
||||
dest_props_str = ", ".join(dest_props)
|
||||
|
||||
# Soft-delete: mark relationship as invalid instead of removing it,
|
||||
# enabling temporal reasoning over historical graph state.
|
||||
# See: https://github.com/mem0ai/mem0/issues/4187
|
||||
cypher = f"""
|
||||
MATCH (n {self.node_label} {{{source_props_str}}})
|
||||
-[r:{relationship}]->
|
||||
(m {self.node_label} {{{dest_props_str}}})
|
||||
WHERE r.valid IS NULL OR r.valid = true
|
||||
SET r.valid = false, r.invalidated_at = datetime()
|
||||
RETURN
|
||||
n.name AS source,
|
||||
m.name AS target,
|
||||
type(r) AS relationship
|
||||
"""
|
||||
|
||||
result = self.graph.query(cypher, params=params)
|
||||
results.append(result)
|
||||
|
||||
return results
|
||||
|
||||
def _add_entities(self, to_be_added, filters, entity_type_map):
|
||||
"""Add the new entities to the graph. Merge the nodes if they already exist."""
|
||||
user_id = filters["user_id"]
|
||||
agent_id = filters.get("agent_id", None)
|
||||
run_id = filters.get("run_id", None)
|
||||
results = []
|
||||
for item in to_be_added:
|
||||
# entities
|
||||
source = item["source"]
|
||||
destination = item["destination"]
|
||||
relationship = item["relationship"]
|
||||
|
||||
# types
|
||||
source_type = entity_type_map.get(source, "__User__")
|
||||
source_label = self.node_label if self.node_label else f":`{source_type}`"
|
||||
source_extra_set = f", source:`{source_type}`" if self.node_label else ""
|
||||
destination_type = entity_type_map.get(destination, "__User__")
|
||||
destination_label = self.node_label if self.node_label else f":`{destination_type}`"
|
||||
destination_extra_set = f", destination:`{destination_type}`" if self.node_label else ""
|
||||
|
||||
# embeddings
|
||||
source_embedding = self.embedding_model.embed(source)
|
||||
dest_embedding = self.embedding_model.embed(destination)
|
||||
|
||||
# search for the nodes with the closest embeddings
|
||||
source_node_search_result = self._search_source_node(source_embedding, filters, threshold=self.threshold)
|
||||
destination_node_search_result = self._search_destination_node(dest_embedding, filters, threshold=self.threshold)
|
||||
|
||||
# TODO: Create a cypher query and common params for all the cases
|
||||
if not destination_node_search_result and source_node_search_result:
|
||||
# Build destination MERGE properties
|
||||
merge_props = ["name: $destination_name", "user_id: $user_id"]
|
||||
if agent_id:
|
||||
merge_props.append("agent_id: $agent_id")
|
||||
if run_id:
|
||||
merge_props.append("run_id: $run_id")
|
||||
merge_props_str = ", ".join(merge_props)
|
||||
|
||||
cypher = f"""
|
||||
MATCH (source)
|
||||
WHERE elementId(source) = $source_id
|
||||
SET source.mentions = coalesce(source.mentions, 0) + 1
|
||||
WITH source
|
||||
MERGE (destination {destination_label} {{{merge_props_str}}})
|
||||
ON CREATE SET
|
||||
destination.created = timestamp(),
|
||||
destination.mentions = 1
|
||||
{destination_extra_set}
|
||||
ON MATCH SET
|
||||
destination.mentions = coalesce(destination.mentions, 0) + 1
|
||||
WITH source, destination
|
||||
CALL db.create.setNodeVectorProperty(destination, 'embedding', $destination_embedding)
|
||||
WITH source, destination
|
||||
MERGE (source)-[r:{relationship}]->(destination)
|
||||
ON CREATE SET
|
||||
r.created_at = timestamp(),
|
||||
r.updated_at = timestamp(),
|
||||
r.mentions = 1,
|
||||
r.valid = true
|
||||
ON MATCH SET
|
||||
r.mentions = coalesce(r.mentions, 0) + 1,
|
||||
r.valid = true,
|
||||
r.updated_at = timestamp(),
|
||||
r.invalidated_at = null
|
||||
RETURN source.name AS source, type(r) AS relationship, destination.name AS target
|
||||
"""
|
||||
|
||||
params = {
|
||||
"source_id": source_node_search_result[0]["elementId(source_candidate)"],
|
||||
"destination_name": destination,
|
||||
"destination_embedding": dest_embedding,
|
||||
"user_id": user_id,
|
||||
}
|
||||
if agent_id:
|
||||
params["agent_id"] = agent_id
|
||||
if run_id:
|
||||
params["run_id"] = run_id
|
||||
|
||||
elif destination_node_search_result and not source_node_search_result:
|
||||
# Build source MERGE properties
|
||||
merge_props = ["name: $source_name", "user_id: $user_id"]
|
||||
if agent_id:
|
||||
merge_props.append("agent_id: $agent_id")
|
||||
if run_id:
|
||||
merge_props.append("run_id: $run_id")
|
||||
merge_props_str = ", ".join(merge_props)
|
||||
|
||||
cypher = f"""
|
||||
MATCH (destination)
|
||||
WHERE elementId(destination) = $destination_id
|
||||
SET destination.mentions = coalesce(destination.mentions, 0) + 1
|
||||
WITH destination
|
||||
MERGE (source {source_label} {{{merge_props_str}}})
|
||||
ON CREATE SET
|
||||
source.created = timestamp(),
|
||||
source.mentions = 1
|
||||
{source_extra_set}
|
||||
ON MATCH SET
|
||||
source.mentions = coalesce(source.mentions, 0) + 1
|
||||
WITH source, destination
|
||||
CALL db.create.setNodeVectorProperty(source, 'embedding', $source_embedding)
|
||||
WITH source, destination
|
||||
MERGE (source)-[r:{relationship}]->(destination)
|
||||
ON CREATE SET
|
||||
r.created_at = timestamp(),
|
||||
r.updated_at = timestamp(),
|
||||
r.mentions = 1,
|
||||
r.valid = true
|
||||
ON MATCH SET
|
||||
r.mentions = coalesce(r.mentions, 0) + 1,
|
||||
r.valid = true,
|
||||
r.updated_at = timestamp(),
|
||||
r.invalidated_at = null
|
||||
RETURN source.name AS source, type(r) AS relationship, destination.name AS target
|
||||
"""
|
||||
|
||||
params = {
|
||||
"destination_id": destination_node_search_result[0]["elementId(destination_candidate)"],
|
||||
"source_name": source,
|
||||
"source_embedding": source_embedding,
|
||||
"user_id": user_id,
|
||||
}
|
||||
if agent_id:
|
||||
params["agent_id"] = agent_id
|
||||
if run_id:
|
||||
params["run_id"] = run_id
|
||||
|
||||
elif source_node_search_result and destination_node_search_result:
|
||||
cypher = f"""
|
||||
MATCH (source)
|
||||
WHERE elementId(source) = $source_id
|
||||
SET source.mentions = coalesce(source.mentions, 0) + 1
|
||||
WITH source
|
||||
MATCH (destination)
|
||||
WHERE elementId(destination) = $destination_id
|
||||
SET destination.mentions = coalesce(destination.mentions, 0) + 1
|
||||
MERGE (source)-[r:{relationship}]->(destination)
|
||||
ON CREATE SET
|
||||
r.created_at = timestamp(),
|
||||
r.updated_at = timestamp(),
|
||||
r.mentions = 1,
|
||||
r.valid = true
|
||||
ON MATCH SET
|
||||
r.mentions = coalesce(r.mentions, 0) + 1,
|
||||
r.valid = true,
|
||||
r.updated_at = timestamp(),
|
||||
r.invalidated_at = null
|
||||
RETURN source.name AS source, type(r) AS relationship, destination.name AS target
|
||||
"""
|
||||
|
||||
params = {
|
||||
"source_id": source_node_search_result[0]["elementId(source_candidate)"],
|
||||
"destination_id": destination_node_search_result[0]["elementId(destination_candidate)"],
|
||||
"user_id": user_id,
|
||||
}
|
||||
if agent_id:
|
||||
params["agent_id"] = agent_id
|
||||
if run_id:
|
||||
params["run_id"] = run_id
|
||||
|
||||
else:
|
||||
# Build dynamic MERGE props for both source and destination
|
||||
source_props = ["name: $source_name", "user_id: $user_id"]
|
||||
dest_props = ["name: $dest_name", "user_id: $user_id"]
|
||||
if agent_id:
|
||||
source_props.append("agent_id: $agent_id")
|
||||
dest_props.append("agent_id: $agent_id")
|
||||
if run_id:
|
||||
source_props.append("run_id: $run_id")
|
||||
dest_props.append("run_id: $run_id")
|
||||
source_props_str = ", ".join(source_props)
|
||||
dest_props_str = ", ".join(dest_props)
|
||||
|
||||
cypher = f"""
|
||||
MERGE (source {source_label} {{{source_props_str}}})
|
||||
ON CREATE SET source.created = timestamp(),
|
||||
source.mentions = 1
|
||||
{source_extra_set}
|
||||
ON MATCH SET source.mentions = coalesce(source.mentions, 0) + 1
|
||||
WITH source
|
||||
CALL db.create.setNodeVectorProperty(source, 'embedding', $source_embedding)
|
||||
WITH source
|
||||
MERGE (destination {destination_label} {{{dest_props_str}}})
|
||||
ON CREATE SET destination.created = timestamp(),
|
||||
destination.mentions = 1
|
||||
{destination_extra_set}
|
||||
ON MATCH SET destination.mentions = coalesce(destination.mentions, 0) + 1
|
||||
WITH source, destination
|
||||
CALL db.create.setNodeVectorProperty(destination, 'embedding', $dest_embedding)
|
||||
WITH source, destination
|
||||
MERGE (source)-[r:{relationship}]->(destination)
|
||||
ON CREATE SET
|
||||
r.created_at = timestamp(),
|
||||
r.updated_at = timestamp(),
|
||||
r.mentions = 1,
|
||||
r.valid = true
|
||||
ON MATCH SET
|
||||
r.mentions = coalesce(r.mentions, 0) + 1,
|
||||
r.valid = true,
|
||||
r.updated_at = timestamp(),
|
||||
r.invalidated_at = null
|
||||
RETURN source.name AS source, type(r) AS relationship, destination.name AS target
|
||||
"""
|
||||
|
||||
params = {
|
||||
"source_name": source,
|
||||
"dest_name": destination,
|
||||
"source_embedding": source_embedding,
|
||||
"dest_embedding": dest_embedding,
|
||||
"user_id": user_id,
|
||||
}
|
||||
if agent_id:
|
||||
params["agent_id"] = agent_id
|
||||
if run_id:
|
||||
params["run_id"] = run_id
|
||||
result = self.graph.query(cypher, params=params)
|
||||
results.append(result)
|
||||
return results
|
||||
|
||||
def _remove_spaces_from_entities(self, entity_list):
|
||||
return remove_spaces_from_entities(entity_list, sanitize_relationship=True)
|
||||
|
||||
def _search_source_node(self, source_embedding, filters, threshold=0.9):
|
||||
# Build WHERE conditions
|
||||
where_conditions = ["source_candidate.embedding IS NOT NULL", "source_candidate.user_id = $user_id"]
|
||||
if filters.get("agent_id"):
|
||||
where_conditions.append("source_candidate.agent_id = $agent_id")
|
||||
if filters.get("run_id"):
|
||||
where_conditions.append("source_candidate.run_id = $run_id")
|
||||
where_clause = " AND ".join(where_conditions)
|
||||
|
||||
cypher = f"""
|
||||
MATCH (source_candidate {self.node_label})
|
||||
WHERE {where_clause}
|
||||
|
||||
WITH source_candidate,
|
||||
round(2 * vector.similarity.cosine(source_candidate.embedding, $source_embedding) - 1, 4) AS source_similarity // denormalize for backward compatibility
|
||||
WHERE source_similarity >= $threshold
|
||||
|
||||
WITH source_candidate, source_similarity
|
||||
ORDER BY source_similarity DESC
|
||||
LIMIT 1
|
||||
|
||||
RETURN elementId(source_candidate)
|
||||
"""
|
||||
|
||||
params = {
|
||||
"source_embedding": source_embedding,
|
||||
"user_id": filters["user_id"],
|
||||
"threshold": threshold,
|
||||
}
|
||||
if filters.get("agent_id"):
|
||||
params["agent_id"] = filters["agent_id"]
|
||||
if filters.get("run_id"):
|
||||
params["run_id"] = filters["run_id"]
|
||||
|
||||
result = self.graph.query(cypher, params=params)
|
||||
return result
|
||||
|
||||
def _search_destination_node(self, destination_embedding, filters, threshold=0.9):
|
||||
# Build WHERE conditions
|
||||
where_conditions = ["destination_candidate.embedding IS NOT NULL", "destination_candidate.user_id = $user_id"]
|
||||
if filters.get("agent_id"):
|
||||
where_conditions.append("destination_candidate.agent_id = $agent_id")
|
||||
if filters.get("run_id"):
|
||||
where_conditions.append("destination_candidate.run_id = $run_id")
|
||||
where_clause = " AND ".join(where_conditions)
|
||||
|
||||
cypher = f"""
|
||||
MATCH (destination_candidate {self.node_label})
|
||||
WHERE {where_clause}
|
||||
|
||||
WITH destination_candidate,
|
||||
round(2 * vector.similarity.cosine(destination_candidate.embedding, $destination_embedding) - 1, 4) AS destination_similarity // denormalize for backward compatibility
|
||||
|
||||
WHERE destination_similarity >= $threshold
|
||||
|
||||
WITH destination_candidate, destination_similarity
|
||||
ORDER BY destination_similarity DESC
|
||||
LIMIT 1
|
||||
|
||||
RETURN elementId(destination_candidate)
|
||||
"""
|
||||
|
||||
params = {
|
||||
"destination_embedding": destination_embedding,
|
||||
"user_id": filters["user_id"],
|
||||
"threshold": threshold,
|
||||
}
|
||||
if filters.get("agent_id"):
|
||||
params["agent_id"] = filters["agent_id"]
|
||||
if filters.get("run_id"):
|
||||
params["run_id"] = filters["run_id"]
|
||||
|
||||
result = self.graph.query(cypher, params=params)
|
||||
return result
|
||||
|
||||
# Reset is not defined in base.py
|
||||
def reset(self):
|
||||
"""Reset the graph by clearing all nodes and relationships."""
|
||||
logger.warning("Clearing graph...")
|
||||
cypher_query = """
|
||||
MATCH (n) DETACH DELETE n
|
||||
"""
|
||||
return self.graph.query(cypher_query)
|
||||
@@ -1,732 +0,0 @@
|
||||
import logging
|
||||
|
||||
from mem0.memory.utils import format_entities, remove_spaces_from_entities
|
||||
|
||||
try:
|
||||
import kuzu
|
||||
except ImportError:
|
||||
raise ImportError("kuzu is not installed. Please install it using pip install kuzu")
|
||||
|
||||
try:
|
||||
from rank_bm25 import BM25Okapi
|
||||
except ImportError:
|
||||
raise ImportError("rank_bm25 is not installed. Please install it using pip install rank-bm25")
|
||||
|
||||
from mem0.graphs.tools import (
|
||||
DELETE_MEMORY_STRUCT_TOOL_GRAPH,
|
||||
DELETE_MEMORY_TOOL_GRAPH,
|
||||
EXTRACT_ENTITIES_STRUCT_TOOL,
|
||||
EXTRACT_ENTITIES_TOOL,
|
||||
RELATIONS_STRUCT_TOOL,
|
||||
RELATIONS_TOOL,
|
||||
)
|
||||
from mem0.graphs.utils import EXTRACT_RELATIONS_PROMPT, get_delete_messages
|
||||
from mem0.utils.factory import EmbedderFactory, LlmFactory
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class MemoryGraph:
|
||||
def __init__(self, config):
|
||||
self.config = config
|
||||
|
||||
self.embedding_model = EmbedderFactory.create(
|
||||
self.config.embedder.provider,
|
||||
self.config.embedder.config,
|
||||
self.config.vector_store.config,
|
||||
)
|
||||
self.embedding_dims = self.embedding_model.config.embedding_dims
|
||||
|
||||
if self.embedding_dims is None or self.embedding_dims <= 0:
|
||||
raise ValueError(f"embedding_dims must be a positive integer. Given: {self.embedding_dims}")
|
||||
|
||||
self.db = kuzu.Database(self.config.graph_store.config.db)
|
||||
self.graph = kuzu.Connection(self.db)
|
||||
|
||||
self.node_label = ":Entity"
|
||||
self.rel_label = ":CONNECTED_TO"
|
||||
self.kuzu_create_schema()
|
||||
|
||||
# Default to openai if no specific provider is configured
|
||||
self.llm_provider = "openai"
|
||||
if self.config.llm and self.config.llm.provider:
|
||||
self.llm_provider = self.config.llm.provider
|
||||
if self.config.graph_store and self.config.graph_store.llm and self.config.graph_store.llm.provider:
|
||||
self.llm_provider = self.config.graph_store.llm.provider
|
||||
# Get LLM config with proper null checks
|
||||
llm_config = None
|
||||
if self.config.graph_store and self.config.graph_store.llm and hasattr(self.config.graph_store.llm, "config"):
|
||||
llm_config = self.config.graph_store.llm.config
|
||||
elif hasattr(self.config.llm, "config"):
|
||||
llm_config = self.config.llm.config
|
||||
self.llm = LlmFactory.create(self.llm_provider, llm_config)
|
||||
|
||||
self.user_id = None
|
||||
# Use threshold from graph_store config, default to 0.7 for backward compatibility
|
||||
self.threshold = self.config.graph_store.threshold if hasattr(self.config.graph_store, 'threshold') else 0.7
|
||||
|
||||
def kuzu_create_schema(self):
|
||||
self.kuzu_execute(
|
||||
"""
|
||||
CREATE NODE TABLE IF NOT EXISTS Entity(
|
||||
id SERIAL PRIMARY KEY,
|
||||
user_id STRING,
|
||||
agent_id STRING,
|
||||
run_id STRING,
|
||||
name STRING,
|
||||
mentions INT64,
|
||||
created TIMESTAMP,
|
||||
embedding FLOAT[]);
|
||||
"""
|
||||
)
|
||||
self.kuzu_execute(
|
||||
"""
|
||||
CREATE REL TABLE IF NOT EXISTS CONNECTED_TO(
|
||||
FROM Entity TO Entity,
|
||||
name STRING,
|
||||
mentions INT64,
|
||||
created TIMESTAMP,
|
||||
updated TIMESTAMP
|
||||
);
|
||||
"""
|
||||
)
|
||||
|
||||
def kuzu_execute(self, query, parameters=None):
|
||||
results = self.graph.execute(query, parameters)
|
||||
return list(results.rows_as_dict())
|
||||
|
||||
def add(self, data, filters):
|
||||
"""
|
||||
Adds data to the graph.
|
||||
|
||||
Args:
|
||||
data (str): The data to add to the graph.
|
||||
filters (dict): A dictionary containing filters to be applied during the addition.
|
||||
"""
|
||||
entity_type_map = self._retrieve_nodes_from_data(data, filters)
|
||||
to_be_added = self._establish_nodes_relations_from_data(data, filters, entity_type_map)
|
||||
search_output = self._search_graph_db(node_list=list(entity_type_map.keys()), filters=filters)
|
||||
to_be_deleted = self._get_delete_entities_from_search_output(search_output, data, filters)
|
||||
|
||||
deleted_entities = self._delete_entities(to_be_deleted, filters)
|
||||
added_entities = self._add_entities(to_be_added, filters, entity_type_map)
|
||||
|
||||
return {"deleted_entities": deleted_entities, "added_entities": added_entities}
|
||||
|
||||
def search(self, query, filters, top_k=5):
|
||||
"""
|
||||
Search for memories and related graph data.
|
||||
|
||||
Args:
|
||||
query (str): Query to search for.
|
||||
filters (dict): A dictionary containing filters to be applied during the search.
|
||||
top_k (int): The maximum number of nodes and relationships to retrieve. Defaults to 100.
|
||||
|
||||
Returns:
|
||||
dict: A dictionary containing:
|
||||
- "contexts": List of search results from the base data store.
|
||||
- "entities": List of related graph data based on the query.
|
||||
"""
|
||||
entity_type_map = self._retrieve_nodes_from_data(query, filters)
|
||||
search_output = self._search_graph_db(node_list=list(entity_type_map.keys()), filters=filters)
|
||||
|
||||
if not search_output:
|
||||
return []
|
||||
|
||||
search_outputs_sequence = [
|
||||
[item["source"], item["relationship"], item["destination"]] for item in search_output
|
||||
]
|
||||
bm25 = BM25Okapi(search_outputs_sequence)
|
||||
|
||||
tokenized_query = query.split(" ")
|
||||
reranked_results = bm25.get_top_n(tokenized_query, search_outputs_sequence, n=top_k)
|
||||
|
||||
search_results = []
|
||||
for item in reranked_results:
|
||||
search_results.append({"source": item[0], "relationship": item[1], "destination": item[2]})
|
||||
|
||||
logger.info(f"Returned {len(search_results)} search results")
|
||||
|
||||
return search_results
|
||||
|
||||
def delete(self, data, filters):
|
||||
"""
|
||||
Delete graph entities associated with the given memory text.
|
||||
|
||||
Extracts entities and relationships from the memory text using the same
|
||||
pipeline as add(), then deletes the matching relationships in the graph.
|
||||
|
||||
Args:
|
||||
data (str): The memory text whose graph entities should be removed.
|
||||
filters (dict): Scope filters (user_id, agent_id, run_id).
|
||||
"""
|
||||
try:
|
||||
entity_type_map = self._retrieve_nodes_from_data(data, filters)
|
||||
if not entity_type_map:
|
||||
logger.debug("No entities found in memory text, skipping graph cleanup")
|
||||
return
|
||||
to_be_deleted = self._establish_nodes_relations_from_data(data, filters, entity_type_map)
|
||||
if to_be_deleted:
|
||||
self._delete_entities(to_be_deleted, filters)
|
||||
except Exception as e:
|
||||
logger.error(f"Error during graph cleanup for memory delete: {e}")
|
||||
|
||||
def delete_all(self, filters):
|
||||
# Build node properties for filtering
|
||||
node_props = ["user_id: $user_id"]
|
||||
if filters.get("agent_id"):
|
||||
node_props.append("agent_id: $agent_id")
|
||||
if filters.get("run_id"):
|
||||
node_props.append("run_id: $run_id")
|
||||
node_props_str = ", ".join(node_props)
|
||||
|
||||
cypher = f"""
|
||||
MATCH (n {self.node_label} {{{node_props_str}}})
|
||||
DETACH DELETE n
|
||||
"""
|
||||
params = {"user_id": filters["user_id"]}
|
||||
if filters.get("agent_id"):
|
||||
params["agent_id"] = filters["agent_id"]
|
||||
if filters.get("run_id"):
|
||||
params["run_id"] = filters["run_id"]
|
||||
self.kuzu_execute(cypher, parameters=params)
|
||||
|
||||
def get_all(self, filters, top_k=100):
|
||||
"""
|
||||
Retrieves all nodes and relationships from the graph database based on optional filtering criteria.
|
||||
Args:
|
||||
filters (dict): A dictionary containing filters to be applied during the retrieval.
|
||||
top_k (int): The maximum number of nodes and relationships to retrieve. Defaults to 100.
|
||||
Returns:
|
||||
list: A list of dictionaries, each containing:
|
||||
- 'contexts': The base data store response for each memory.
|
||||
- 'entities': A list of strings representing the nodes and relationships
|
||||
"""
|
||||
|
||||
params = {
|
||||
"user_id": filters["user_id"],
|
||||
"limit": top_k,
|
||||
}
|
||||
# Build node properties based on filters
|
||||
node_props = ["user_id: $user_id"]
|
||||
if filters.get("agent_id"):
|
||||
node_props.append("agent_id: $agent_id")
|
||||
params["agent_id"] = filters["agent_id"]
|
||||
if filters.get("run_id"):
|
||||
node_props.append("run_id: $run_id")
|
||||
params["run_id"] = filters["run_id"]
|
||||
node_props_str = ", ".join(node_props)
|
||||
|
||||
query = f"""
|
||||
MATCH (n {self.node_label} {{{node_props_str}}})-[r]->(m {self.node_label} {{{node_props_str}}})
|
||||
RETURN
|
||||
n.name AS source,
|
||||
r.name AS relationship,
|
||||
m.name AS target
|
||||
LIMIT $limit
|
||||
"""
|
||||
results = self.kuzu_execute(query, parameters=params)
|
||||
|
||||
final_results = []
|
||||
for result in results:
|
||||
final_results.append(
|
||||
{
|
||||
"source": result["source"],
|
||||
"relationship": result["relationship"],
|
||||
"target": result["target"],
|
||||
}
|
||||
)
|
||||
|
||||
logger.info(f"Retrieved {len(final_results)} relationships")
|
||||
|
||||
return final_results
|
||||
|
||||
def _retrieve_nodes_from_data(self, data, filters):
|
||||
"""Extracts all the entities mentioned in the query."""
|
||||
_tools = [EXTRACT_ENTITIES_TOOL]
|
||||
if self.llm_provider in ["azure_openai_structured", "openai_structured"]:
|
||||
_tools = [EXTRACT_ENTITIES_STRUCT_TOOL]
|
||||
search_results = self.llm.generate_response(
|
||||
messages=[
|
||||
{
|
||||
"role": "system",
|
||||
"content": f"You are a smart assistant who understands entities and their types in a given text. If user message contains self reference such as 'I', 'me', 'my' etc. then use {filters['user_id']} as the source entity. Extract all the entities from the text. ***DO NOT*** answer the question itself if the given text is a question.",
|
||||
},
|
||||
{"role": "user", "content": data},
|
||||
],
|
||||
tools=_tools,
|
||||
)
|
||||
|
||||
entity_type_map = {}
|
||||
|
||||
try:
|
||||
for tool_call in search_results["tool_calls"]:
|
||||
if tool_call["name"] != "extract_entities":
|
||||
continue
|
||||
for item in tool_call.get("arguments", {}).get("entities", []):
|
||||
entity_type_map[item["entity"]] = item["entity_type"]
|
||||
except Exception as e:
|
||||
logger.exception(
|
||||
f"Error in search tool: {e}, llm_provider={self.llm_provider}, search_results={search_results}"
|
||||
)
|
||||
|
||||
entity_type_map = {k.lower().replace(" ", "_"): v.lower().replace(" ", "_") for k, v in entity_type_map.items()}
|
||||
logger.debug(f"Entity type map: {entity_type_map}\n search_results={search_results}")
|
||||
return entity_type_map
|
||||
|
||||
def _establish_nodes_relations_from_data(self, data, filters, entity_type_map):
|
||||
"""Establish relations among the extracted nodes."""
|
||||
|
||||
# Compose user identification string for prompt
|
||||
user_identity = f"user_id: {filters['user_id']}"
|
||||
if filters.get("agent_id"):
|
||||
user_identity += f", agent_id: {filters['agent_id']}"
|
||||
if filters.get("run_id"):
|
||||
user_identity += f", run_id: {filters['run_id']}"
|
||||
|
||||
if self.config.graph_store.custom_prompt:
|
||||
system_content = EXTRACT_RELATIONS_PROMPT.replace("USER_ID", user_identity)
|
||||
# Add the custom prompt line if configured
|
||||
system_content = system_content.replace("CUSTOM_PROMPT", f"4. {self.config.graph_store.custom_prompt}")
|
||||
messages = [
|
||||
{"role": "system", "content": system_content},
|
||||
{"role": "user", "content": data},
|
||||
]
|
||||
else:
|
||||
system_content = EXTRACT_RELATIONS_PROMPT.replace("USER_ID", user_identity)
|
||||
messages = [
|
||||
{"role": "system", "content": system_content},
|
||||
{"role": "user", "content": f"List of entities: {list(entity_type_map.keys())}. \n\nText: {data}"},
|
||||
]
|
||||
|
||||
_tools = [RELATIONS_TOOL]
|
||||
if self.llm_provider in ["azure_openai_structured", "openai_structured"]:
|
||||
_tools = [RELATIONS_STRUCT_TOOL]
|
||||
|
||||
extracted_entities = self.llm.generate_response(
|
||||
messages=messages,
|
||||
tools=_tools,
|
||||
)
|
||||
|
||||
entities = []
|
||||
if extracted_entities.get("tool_calls"):
|
||||
entities = extracted_entities["tool_calls"][0].get("arguments", {}).get("entities", [])
|
||||
|
||||
entities = self._remove_spaces_from_entities(entities)
|
||||
logger.debug(f"Extracted entities: {entities}")
|
||||
return entities
|
||||
|
||||
def _search_graph_db(self, node_list, filters, top_k=100, threshold=None):
|
||||
"""Search similar nodes among and their respective incoming and outgoing relations."""
|
||||
result_relations = []
|
||||
|
||||
params = {
|
||||
"threshold": threshold if threshold else self.threshold,
|
||||
"user_id": filters["user_id"],
|
||||
"limit": top_k,
|
||||
}
|
||||
# Build node properties for filtering
|
||||
node_props = ["user_id: $user_id"]
|
||||
if filters.get("agent_id"):
|
||||
node_props.append("agent_id: $agent_id")
|
||||
params["agent_id"] = filters["agent_id"]
|
||||
if filters.get("run_id"):
|
||||
node_props.append("run_id: $run_id")
|
||||
params["run_id"] = filters["run_id"]
|
||||
node_props_str = ", ".join(node_props)
|
||||
|
||||
for node in node_list:
|
||||
n_embedding = self.embedding_model.embed(node)
|
||||
params["n_embedding"] = n_embedding
|
||||
|
||||
results = []
|
||||
for match_fragment in [
|
||||
f"(n)-[r]->(m {self.node_label} {{{node_props_str}}}) WITH n as src, r, m as dst, similarity",
|
||||
f"(m {self.node_label} {{{node_props_str}}})-[r]->(n) WITH m as src, r, n as dst, similarity"
|
||||
]:
|
||||
results.extend(self.kuzu_execute(
|
||||
f"""
|
||||
MATCH (n {self.node_label} {{{node_props_str}}})
|
||||
WHERE n.embedding IS NOT NULL
|
||||
WITH n, array_cosine_similarity(n.embedding, CAST($n_embedding,'FLOAT[{self.embedding_dims}]')) AS similarity
|
||||
WHERE similarity >= CAST($threshold, 'DOUBLE')
|
||||
MATCH {match_fragment}
|
||||
RETURN
|
||||
src.name AS source,
|
||||
id(src) AS source_id,
|
||||
r.name AS relationship,
|
||||
id(r) AS relation_id,
|
||||
dst.name AS destination,
|
||||
id(dst) AS destination_id,
|
||||
similarity
|
||||
LIMIT $limit
|
||||
""",
|
||||
parameters=params))
|
||||
|
||||
# Kuzu does not support sort/limit over unions. Do it manually for now.
|
||||
result_relations.extend(sorted(results, key=lambda x: x["similarity"], reverse=True)[:top_k])
|
||||
|
||||
return result_relations
|
||||
|
||||
def _get_delete_entities_from_search_output(self, search_output, data, filters):
|
||||
"""Get the entities to be deleted from the search output."""
|
||||
search_output_string = format_entities(search_output)
|
||||
|
||||
# Compose user identification string for prompt
|
||||
user_identity = f"user_id: {filters['user_id']}"
|
||||
if filters.get("agent_id"):
|
||||
user_identity += f", agent_id: {filters['agent_id']}"
|
||||
if filters.get("run_id"):
|
||||
user_identity += f", run_id: {filters['run_id']}"
|
||||
|
||||
system_prompt, user_prompt = get_delete_messages(search_output_string, data, user_identity)
|
||||
|
||||
_tools = [DELETE_MEMORY_TOOL_GRAPH]
|
||||
if self.llm_provider in ["azure_openai_structured", "openai_structured"]:
|
||||
_tools = [
|
||||
DELETE_MEMORY_STRUCT_TOOL_GRAPH,
|
||||
]
|
||||
|
||||
memory_updates = self.llm.generate_response(
|
||||
messages=[
|
||||
{"role": "system", "content": system_prompt},
|
||||
{"role": "user", "content": user_prompt},
|
||||
],
|
||||
tools=_tools,
|
||||
)
|
||||
|
||||
to_be_deleted = []
|
||||
for item in memory_updates.get("tool_calls", []):
|
||||
if item.get("name") == "delete_graph_memory":
|
||||
to_be_deleted.append(item.get("arguments"))
|
||||
# Clean entities formatting
|
||||
to_be_deleted = self._remove_spaces_from_entities(to_be_deleted)
|
||||
logger.debug(f"Deleted relationships: {to_be_deleted}")
|
||||
return to_be_deleted
|
||||
|
||||
def _delete_entities(self, to_be_deleted, filters):
|
||||
"""Delete the entities from the graph."""
|
||||
user_id = filters["user_id"]
|
||||
agent_id = filters.get("agent_id", None)
|
||||
run_id = filters.get("run_id", None)
|
||||
results = []
|
||||
|
||||
for item in to_be_deleted:
|
||||
source = item["source"]
|
||||
destination = item["destination"]
|
||||
relationship = item["relationship"]
|
||||
|
||||
params = {
|
||||
"source_name": source,
|
||||
"dest_name": destination,
|
||||
"user_id": user_id,
|
||||
"relationship_name": relationship,
|
||||
}
|
||||
# Build node properties for filtering
|
||||
source_props = ["name: $source_name", "user_id: $user_id"]
|
||||
dest_props = ["name: $dest_name", "user_id: $user_id"]
|
||||
if agent_id:
|
||||
source_props.append("agent_id: $agent_id")
|
||||
dest_props.append("agent_id: $agent_id")
|
||||
params["agent_id"] = agent_id
|
||||
if run_id:
|
||||
source_props.append("run_id: $run_id")
|
||||
dest_props.append("run_id: $run_id")
|
||||
params["run_id"] = run_id
|
||||
source_props_str = ", ".join(source_props)
|
||||
dest_props_str = ", ".join(dest_props)
|
||||
|
||||
# Delete the specific relationship between nodes
|
||||
cypher = f"""
|
||||
MATCH (n {self.node_label} {{{source_props_str}}})
|
||||
-[r {self.rel_label} {{name: $relationship_name}}]->
|
||||
(m {self.node_label} {{{dest_props_str}}})
|
||||
DELETE r
|
||||
RETURN
|
||||
n.name AS source,
|
||||
r.name AS relationship,
|
||||
m.name AS target
|
||||
"""
|
||||
|
||||
result = self.kuzu_execute(cypher, parameters=params)
|
||||
results.append(result)
|
||||
|
||||
return results
|
||||
|
||||
def _add_entities(self, to_be_added, filters, entity_type_map):
|
||||
"""Add the new entities to the graph. Merge the nodes if they already exist."""
|
||||
user_id = filters["user_id"]
|
||||
agent_id = filters.get("agent_id", None)
|
||||
run_id = filters.get("run_id", None)
|
||||
results = []
|
||||
for item in to_be_added:
|
||||
# entities
|
||||
source = item["source"]
|
||||
source_label = self.node_label
|
||||
|
||||
destination = item["destination"]
|
||||
destination_label = self.node_label
|
||||
|
||||
relationship = item["relationship"]
|
||||
relationship_label = self.rel_label
|
||||
|
||||
# embeddings
|
||||
source_embedding = self.embedding_model.embed(source)
|
||||
dest_embedding = self.embedding_model.embed(destination)
|
||||
|
||||
# search for the nodes with the closest embeddings
|
||||
source_node_search_result = self._search_source_node(source_embedding, filters, threshold=self.threshold)
|
||||
destination_node_search_result = self._search_destination_node(dest_embedding, filters, threshold=self.threshold)
|
||||
|
||||
if not destination_node_search_result and source_node_search_result:
|
||||
params = {
|
||||
"table_id": source_node_search_result[0]["id"]["table"],
|
||||
"offset_id": source_node_search_result[0]["id"]["offset"],
|
||||
"destination_name": destination,
|
||||
"destination_embedding": dest_embedding,
|
||||
"relationship_name": relationship,
|
||||
"user_id": user_id,
|
||||
}
|
||||
# Build source MERGE properties
|
||||
merge_props = ["name: $destination_name", "user_id: $user_id"]
|
||||
if agent_id:
|
||||
merge_props.append("agent_id: $agent_id")
|
||||
params["agent_id"] = agent_id
|
||||
if run_id:
|
||||
merge_props.append("run_id: $run_id")
|
||||
params["run_id"] = run_id
|
||||
merge_props_str = ", ".join(merge_props)
|
||||
|
||||
cypher = f"""
|
||||
MATCH (source)
|
||||
WHERE id(source) = internal_id($table_id, $offset_id)
|
||||
SET source.mentions = coalesce(source.mentions, 0) + 1
|
||||
WITH source
|
||||
MERGE (destination {destination_label} {{{merge_props_str}}})
|
||||
ON CREATE SET
|
||||
destination.created = current_timestamp(),
|
||||
destination.mentions = 1,
|
||||
destination.embedding = CAST($destination_embedding,'FLOAT[{self.embedding_dims}]')
|
||||
ON MATCH SET
|
||||
destination.mentions = coalesce(destination.mentions, 0) + 1,
|
||||
destination.embedding = CAST($destination_embedding,'FLOAT[{self.embedding_dims}]')
|
||||
WITH source, destination
|
||||
MERGE (source)-[r {relationship_label} {{name: $relationship_name}}]->(destination)
|
||||
ON CREATE SET
|
||||
r.created = current_timestamp(),
|
||||
r.mentions = 1
|
||||
ON MATCH SET
|
||||
r.mentions = coalesce(r.mentions, 0) + 1
|
||||
RETURN
|
||||
source.name AS source,
|
||||
r.name AS relationship,
|
||||
destination.name AS target
|
||||
"""
|
||||
elif destination_node_search_result and not source_node_search_result:
|
||||
params = {
|
||||
"table_id": destination_node_search_result[0]["id"]["table"],
|
||||
"offset_id": destination_node_search_result[0]["id"]["offset"],
|
||||
"source_name": source,
|
||||
"source_embedding": source_embedding,
|
||||
"user_id": user_id,
|
||||
"relationship_name": relationship,
|
||||
}
|
||||
# Build source MERGE properties
|
||||
merge_props = ["name: $source_name", "user_id: $user_id"]
|
||||
if agent_id:
|
||||
merge_props.append("agent_id: $agent_id")
|
||||
params["agent_id"] = agent_id
|
||||
if run_id:
|
||||
merge_props.append("run_id: $run_id")
|
||||
params["run_id"] = run_id
|
||||
merge_props_str = ", ".join(merge_props)
|
||||
|
||||
cypher = f"""
|
||||
MATCH (destination)
|
||||
WHERE id(destination) = internal_id($table_id, $offset_id)
|
||||
SET destination.mentions = coalesce(destination.mentions, 0) + 1
|
||||
WITH destination
|
||||
MERGE (source {source_label} {{{merge_props_str}}})
|
||||
ON CREATE SET
|
||||
source.created = current_timestamp(),
|
||||
source.mentions = 1,
|
||||
source.embedding = CAST($source_embedding,'FLOAT[{self.embedding_dims}]')
|
||||
ON MATCH SET
|
||||
source.mentions = coalesce(source.mentions, 0) + 1,
|
||||
source.embedding = CAST($source_embedding,'FLOAT[{self.embedding_dims}]')
|
||||
WITH source, destination
|
||||
MERGE (source)-[r {relationship_label} {{name: $relationship_name}}]->(destination)
|
||||
ON CREATE SET
|
||||
r.created = current_timestamp(),
|
||||
r.mentions = 1
|
||||
ON MATCH SET
|
||||
r.mentions = coalesce(r.mentions, 0) + 1
|
||||
RETURN
|
||||
source.name AS source,
|
||||
r.name AS relationship,
|
||||
destination.name AS target
|
||||
"""
|
||||
elif source_node_search_result and destination_node_search_result:
|
||||
cypher = f"""
|
||||
MATCH (source)
|
||||
WHERE id(source) = internal_id($src_table, $src_offset)
|
||||
SET source.mentions = coalesce(source.mentions, 0) + 1
|
||||
WITH source
|
||||
MATCH (destination)
|
||||
WHERE id(destination) = internal_id($dst_table, $dst_offset)
|
||||
SET destination.mentions = coalesce(destination.mentions, 0) + 1
|
||||
MERGE (source)-[r {relationship_label} {{name: $relationship_name}}]->(destination)
|
||||
ON CREATE SET
|
||||
r.created = current_timestamp(),
|
||||
r.updated = current_timestamp(),
|
||||
r.mentions = 1
|
||||
ON MATCH SET r.mentions = coalesce(r.mentions, 0) + 1
|
||||
RETURN
|
||||
source.name AS source,
|
||||
r.name AS relationship,
|
||||
destination.name AS target
|
||||
"""
|
||||
|
||||
params = {
|
||||
"src_table": source_node_search_result[0]["id"]["table"],
|
||||
"src_offset": source_node_search_result[0]["id"]["offset"],
|
||||
"dst_table": destination_node_search_result[0]["id"]["table"],
|
||||
"dst_offset": destination_node_search_result[0]["id"]["offset"],
|
||||
"relationship_name": relationship,
|
||||
}
|
||||
else:
|
||||
params = {
|
||||
"source_name": source,
|
||||
"dest_name": destination,
|
||||
"relationship_name": relationship,
|
||||
"source_embedding": source_embedding,
|
||||
"dest_embedding": dest_embedding,
|
||||
"user_id": user_id,
|
||||
}
|
||||
# Build dynamic MERGE props for both source and destination
|
||||
source_props = ["name: $source_name", "user_id: $user_id"]
|
||||
dest_props = ["name: $dest_name", "user_id: $user_id"]
|
||||
if agent_id:
|
||||
source_props.append("agent_id: $agent_id")
|
||||
dest_props.append("agent_id: $agent_id")
|
||||
params["agent_id"] = agent_id
|
||||
if run_id:
|
||||
source_props.append("run_id: $run_id")
|
||||
dest_props.append("run_id: $run_id")
|
||||
params["run_id"] = run_id
|
||||
source_props_str = ", ".join(source_props)
|
||||
dest_props_str = ", ".join(dest_props)
|
||||
|
||||
cypher = f"""
|
||||
MERGE (source {source_label} {{{source_props_str}}})
|
||||
ON CREATE SET
|
||||
source.created = current_timestamp(),
|
||||
source.mentions = 1,
|
||||
source.embedding = CAST($source_embedding,'FLOAT[{self.embedding_dims}]')
|
||||
ON MATCH SET
|
||||
source.mentions = coalesce(source.mentions, 0) + 1,
|
||||
source.embedding = CAST($source_embedding,'FLOAT[{self.embedding_dims}]')
|
||||
WITH source
|
||||
MERGE (destination {destination_label} {{{dest_props_str}}})
|
||||
ON CREATE SET
|
||||
destination.created = current_timestamp(),
|
||||
destination.mentions = 1,
|
||||
destination.embedding = CAST($dest_embedding,'FLOAT[{self.embedding_dims}]')
|
||||
ON MATCH SET
|
||||
destination.mentions = coalesce(destination.mentions, 0) + 1,
|
||||
destination.embedding = CAST($dest_embedding,'FLOAT[{self.embedding_dims}]')
|
||||
WITH source, destination
|
||||
MERGE (source)-[rel {relationship_label} {{name: $relationship_name}}]->(destination)
|
||||
ON CREATE SET
|
||||
rel.created = current_timestamp(),
|
||||
rel.mentions = 1
|
||||
ON MATCH SET
|
||||
rel.mentions = coalesce(rel.mentions, 0) + 1
|
||||
RETURN
|
||||
source.name AS source,
|
||||
rel.name AS relationship,
|
||||
destination.name AS target
|
||||
"""
|
||||
|
||||
result = self.kuzu_execute(cypher, parameters=params)
|
||||
results.append(result)
|
||||
|
||||
return results
|
||||
|
||||
def _remove_spaces_from_entities(self, entity_list):
|
||||
return remove_spaces_from_entities(entity_list, sanitize_relationship=False)
|
||||
|
||||
def _search_source_node(self, source_embedding, filters, threshold=0.9):
|
||||
params = {
|
||||
"source_embedding": source_embedding,
|
||||
"user_id": filters["user_id"],
|
||||
"threshold": threshold,
|
||||
}
|
||||
where_conditions = ["source_candidate.embedding IS NOT NULL", "source_candidate.user_id = $user_id"]
|
||||
if filters.get("agent_id"):
|
||||
where_conditions.append("source_candidate.agent_id = $agent_id")
|
||||
params["agent_id"] = filters["agent_id"]
|
||||
if filters.get("run_id"):
|
||||
where_conditions.append("source_candidate.run_id = $run_id")
|
||||
params["run_id"] = filters["run_id"]
|
||||
where_clause = " AND ".join(where_conditions)
|
||||
|
||||
cypher = f"""
|
||||
MATCH (source_candidate {self.node_label})
|
||||
WHERE {where_clause}
|
||||
|
||||
WITH source_candidate,
|
||||
array_cosine_similarity(source_candidate.embedding, CAST($source_embedding,'FLOAT[{self.embedding_dims}]')) AS source_similarity
|
||||
|
||||
WHERE source_similarity >= $threshold
|
||||
|
||||
WITH source_candidate, source_similarity
|
||||
ORDER BY source_similarity DESC
|
||||
LIMIT 2
|
||||
|
||||
RETURN id(source_candidate) as id, source_similarity
|
||||
"""
|
||||
|
||||
return self.kuzu_execute(cypher, parameters=params)
|
||||
|
||||
def _search_destination_node(self, destination_embedding, filters, threshold=0.9):
|
||||
params = {
|
||||
"destination_embedding": destination_embedding,
|
||||
"user_id": filters["user_id"],
|
||||
"threshold": threshold,
|
||||
}
|
||||
where_conditions = ["destination_candidate.embedding IS NOT NULL", "destination_candidate.user_id = $user_id"]
|
||||
if filters.get("agent_id"):
|
||||
where_conditions.append("destination_candidate.agent_id = $agent_id")
|
||||
params["agent_id"] = filters["agent_id"]
|
||||
if filters.get("run_id"):
|
||||
where_conditions.append("destination_candidate.run_id = $run_id")
|
||||
params["run_id"] = filters["run_id"]
|
||||
where_clause = " AND ".join(where_conditions)
|
||||
|
||||
cypher = f"""
|
||||
MATCH (destination_candidate {self.node_label})
|
||||
WHERE {where_clause}
|
||||
|
||||
WITH destination_candidate,
|
||||
array_cosine_similarity(destination_candidate.embedding, CAST($destination_embedding,'FLOAT[{self.embedding_dims}]')) AS destination_similarity
|
||||
|
||||
WHERE destination_similarity >= $threshold
|
||||
|
||||
WITH destination_candidate, destination_similarity
|
||||
ORDER BY destination_similarity DESC
|
||||
LIMIT 2
|
||||
|
||||
RETURN id(destination_candidate) as id, destination_similarity
|
||||
"""
|
||||
|
||||
return self.kuzu_execute(cypher, parameters=params)
|
||||
|
||||
# Reset is not defined in base.py
|
||||
def reset(self):
|
||||
"""Reset the graph by clearing all nodes and relationships."""
|
||||
logger.warning("Clearing graph...")
|
||||
cypher_query = """
|
||||
MATCH (n) DETACH DELETE n
|
||||
"""
|
||||
return self.kuzu_execute(cypher_query)
|
||||
+1090
-832
File diff suppressed because it is too large
Load Diff
@@ -1,708 +0,0 @@
|
||||
import logging
|
||||
|
||||
from mem0.memory.utils import format_entities, remove_spaces_from_entities
|
||||
|
||||
try:
|
||||
from langchain_memgraph.graphs.memgraph import Memgraph
|
||||
except ImportError:
|
||||
raise ImportError("langchain_memgraph is not installed. Please install it using pip install langchain-memgraph")
|
||||
|
||||
try:
|
||||
from rank_bm25 import BM25Okapi
|
||||
except ImportError:
|
||||
raise ImportError("rank_bm25 is not installed. Please install it using pip install rank-bm25")
|
||||
|
||||
from mem0.graphs.tools import (
|
||||
DELETE_MEMORY_STRUCT_TOOL_GRAPH,
|
||||
DELETE_MEMORY_TOOL_GRAPH,
|
||||
EXTRACT_ENTITIES_STRUCT_TOOL,
|
||||
EXTRACT_ENTITIES_TOOL,
|
||||
RELATIONS_STRUCT_TOOL,
|
||||
RELATIONS_TOOL,
|
||||
)
|
||||
from mem0.graphs.utils import EXTRACT_RELATIONS_PROMPT, get_delete_messages
|
||||
from mem0.utils.factory import EmbedderFactory, LlmFactory
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class MemoryGraph:
|
||||
def __init__(self, config):
|
||||
self.config = config
|
||||
self.graph = Memgraph(
|
||||
self.config.graph_store.config.url,
|
||||
self.config.graph_store.config.username,
|
||||
self.config.graph_store.config.password,
|
||||
)
|
||||
self.embedding_model = EmbedderFactory.create(
|
||||
self.config.embedder.provider,
|
||||
self.config.embedder.config,
|
||||
{"enable_embeddings": True},
|
||||
)
|
||||
|
||||
# Default to openai if no specific provider is configured
|
||||
self.llm_provider = "openai"
|
||||
if self.config.llm and self.config.llm.provider:
|
||||
self.llm_provider = self.config.llm.provider
|
||||
if self.config.graph_store and self.config.graph_store.llm and self.config.graph_store.llm.provider:
|
||||
self.llm_provider = self.config.graph_store.llm.provider
|
||||
|
||||
# Get LLM config with proper null checks
|
||||
llm_config = None
|
||||
if self.config.graph_store and self.config.graph_store.llm and hasattr(self.config.graph_store.llm, "config"):
|
||||
llm_config = self.config.graph_store.llm.config
|
||||
elif hasattr(self.config.llm, "config"):
|
||||
llm_config = self.config.llm.config
|
||||
self.llm = LlmFactory.create(self.llm_provider, llm_config)
|
||||
self.user_id = None
|
||||
# Use threshold from graph_store config, default to 0.7 for backward compatibility
|
||||
self.threshold = self.config.graph_store.threshold if hasattr(self.config.graph_store, 'threshold') else 0.7
|
||||
|
||||
# Setup Memgraph:
|
||||
# 1. Create vector index (created Entity label on all nodes)
|
||||
# 2. Create label property index for performance optimizations
|
||||
embedding_dims = self.config.embedder.config["embedding_dims"]
|
||||
index_info = self._fetch_existing_indexes()
|
||||
|
||||
# Create vector index if not exists
|
||||
if not self._vector_index_exists(index_info, "memzero"):
|
||||
self.graph.query(
|
||||
f"CREATE VECTOR INDEX memzero ON :Entity(embedding) WITH CONFIG {{'dimension': {embedding_dims}, 'capacity': 1000, 'metric': 'cos'}};"
|
||||
)
|
||||
|
||||
# Create label+property index if not exists
|
||||
if not self._label_property_index_exists(index_info, "Entity", "user_id"):
|
||||
self.graph.query("CREATE INDEX ON :Entity(user_id);")
|
||||
|
||||
# Create label index if not exists
|
||||
if not self._label_index_exists(index_info, "Entity"):
|
||||
self.graph.query("CREATE INDEX ON :Entity;")
|
||||
|
||||
def add(self, data, filters):
|
||||
"""
|
||||
Adds data to the graph.
|
||||
|
||||
Args:
|
||||
data (str): The data to add to the graph.
|
||||
filters (dict): A dictionary containing filters to be applied during the addition.
|
||||
"""
|
||||
entity_type_map = self._retrieve_nodes_from_data(data, filters)
|
||||
to_be_added = self._establish_nodes_relations_from_data(data, filters, entity_type_map)
|
||||
search_output = self._search_graph_db(node_list=list(entity_type_map.keys()), filters=filters)
|
||||
to_be_deleted = self._get_delete_entities_from_search_output(search_output, data, filters)
|
||||
|
||||
# TODO: Batch queries with APOC plugin
|
||||
# TODO: Add more filter support
|
||||
deleted_entities = self._delete_entities(to_be_deleted, filters)
|
||||
added_entities = self._add_entities(to_be_added, filters, entity_type_map)
|
||||
|
||||
return {"deleted_entities": deleted_entities, "added_entities": added_entities}
|
||||
|
||||
def search(self, query, filters, top_k=100):
|
||||
"""
|
||||
Search for memories and related graph data.
|
||||
|
||||
Args:
|
||||
query (str): Query to search for.
|
||||
filters (dict): A dictionary containing filters to be applied during the search.
|
||||
top_k (int): The maximum number of nodes and relationships to retrieve. Defaults to 100.
|
||||
|
||||
Returns:
|
||||
dict: A dictionary containing:
|
||||
- "contexts": List of search results from the base data store.
|
||||
- "entities": List of related graph data based on the query.
|
||||
"""
|
||||
entity_type_map = self._retrieve_nodes_from_data(query, filters)
|
||||
search_output = self._search_graph_db(node_list=list(entity_type_map.keys()), filters=filters)
|
||||
|
||||
if not search_output:
|
||||
return []
|
||||
|
||||
search_outputs_sequence = [
|
||||
[item["source"], item["relationship"], item["destination"]] for item in search_output
|
||||
]
|
||||
bm25 = BM25Okapi(search_outputs_sequence)
|
||||
|
||||
tokenized_query = query.split(" ")
|
||||
reranked_results = bm25.get_top_n(tokenized_query, search_outputs_sequence, n=5)
|
||||
|
||||
search_results = []
|
||||
for item in reranked_results:
|
||||
search_results.append({"source": item[0], "relationship": item[1], "destination": item[2]})
|
||||
|
||||
logger.info(f"Returned {len(search_results)} search results")
|
||||
|
||||
return search_results
|
||||
|
||||
def delete(self, data, filters):
|
||||
"""
|
||||
Delete graph entities associated with the given memory text.
|
||||
|
||||
Extracts entities and relationships from the memory text using the same
|
||||
pipeline as add(), then deletes the matching relationships in the graph.
|
||||
|
||||
Args:
|
||||
data (str): The memory text whose graph entities should be removed.
|
||||
filters (dict): Scope filters (user_id, agent_id).
|
||||
"""
|
||||
try:
|
||||
entity_type_map = self._retrieve_nodes_from_data(data, filters)
|
||||
if not entity_type_map:
|
||||
logger.debug("No entities found in memory text, skipping graph cleanup")
|
||||
return
|
||||
to_be_deleted = self._establish_nodes_relations_from_data(data, filters, entity_type_map)
|
||||
if to_be_deleted:
|
||||
self._delete_entities(to_be_deleted, filters)
|
||||
except Exception as e:
|
||||
logger.error(f"Error during graph cleanup for memory delete: {e}")
|
||||
|
||||
def delete_all(self, filters):
|
||||
"""Delete all nodes and relationships for a user or specific agent."""
|
||||
if filters.get("agent_id"):
|
||||
cypher = """
|
||||
MATCH (n:Entity {user_id: $user_id, agent_id: $agent_id})
|
||||
DETACH DELETE n
|
||||
"""
|
||||
params = {"user_id": filters["user_id"], "agent_id": filters["agent_id"]}
|
||||
else:
|
||||
cypher = """
|
||||
MATCH (n:Entity {user_id: $user_id})
|
||||
DETACH DELETE n
|
||||
"""
|
||||
params = {"user_id": filters["user_id"]}
|
||||
self.graph.query(cypher, params=params)
|
||||
|
||||
def get_all(self, filters, top_k=100):
|
||||
"""
|
||||
Retrieves all nodes and relationships from the graph database based on optional filtering criteria.
|
||||
|
||||
Args:
|
||||
filters (dict): A dictionary containing filters to be applied during the retrieval.
|
||||
Supports 'user_id' (required) and 'agent_id' (optional).
|
||||
top_k (int): The maximum number of nodes and relationships to retrieve. Defaults to 100.
|
||||
Returns:
|
||||
list: A list of dictionaries, each containing:
|
||||
- 'source': The source node name.
|
||||
- 'relationship': The relationship type.
|
||||
- 'target': The target node name.
|
||||
"""
|
||||
# Build query based on whether agent_id is provided
|
||||
if filters.get("agent_id"):
|
||||
query = """
|
||||
MATCH (n:Entity {user_id: $user_id, agent_id: $agent_id})-[r]->(m:Entity {user_id: $user_id, agent_id: $agent_id})
|
||||
RETURN n.name AS source, type(r) AS relationship, m.name AS target
|
||||
LIMIT $limit
|
||||
"""
|
||||
params = {"user_id": filters["user_id"], "agent_id": filters["agent_id"], "limit": top_k}
|
||||
else:
|
||||
query = """
|
||||
MATCH (n:Entity {user_id: $user_id})-[r]->(m:Entity {user_id: $user_id})
|
||||
RETURN n.name AS source, type(r) AS relationship, m.name AS target
|
||||
LIMIT $limit
|
||||
"""
|
||||
params = {"user_id": filters["user_id"], "limit": top_k}
|
||||
|
||||
results = self.graph.query(query, params=params)
|
||||
|
||||
final_results = []
|
||||
for result in results:
|
||||
final_results.append(
|
||||
{
|
||||
"source": result["source"],
|
||||
"relationship": result["relationship"],
|
||||
"target": result["target"],
|
||||
}
|
||||
)
|
||||
|
||||
logger.info(f"Retrieved {len(final_results)} relationships")
|
||||
|
||||
return final_results
|
||||
|
||||
def _retrieve_nodes_from_data(self, data, filters):
|
||||
"""Extracts all the entities mentioned in the query."""
|
||||
_tools = [EXTRACT_ENTITIES_TOOL]
|
||||
if self.llm_provider in ["azure_openai_structured", "openai_structured"]:
|
||||
_tools = [EXTRACT_ENTITIES_STRUCT_TOOL]
|
||||
search_results = self.llm.generate_response(
|
||||
messages=[
|
||||
{
|
||||
"role": "system",
|
||||
"content": f"You are a smart assistant who understands entities and their types in a given text. If user message contains self reference such as 'I', 'me', 'my' etc. then use {filters['user_id']} as the source entity. Extract all the entities from the text. ***DO NOT*** answer the question itself if the given text is a question.",
|
||||
},
|
||||
{"role": "user", "content": data},
|
||||
],
|
||||
tools=_tools,
|
||||
)
|
||||
|
||||
entity_type_map = {}
|
||||
|
||||
try:
|
||||
for tool_call in search_results["tool_calls"]:
|
||||
if tool_call["name"] != "extract_entities":
|
||||
continue
|
||||
for item in tool_call.get("arguments", {}).get("entities", []):
|
||||
if "entity" in item and "entity_type" in item:
|
||||
entity_type_map[item["entity"]] = item["entity_type"]
|
||||
except Exception as e:
|
||||
logger.exception(
|
||||
f"Error in search tool: {e}, llm_provider={self.llm_provider}, search_results={search_results}"
|
||||
)
|
||||
|
||||
entity_type_map = {k.lower().replace(" ", "_"): v.lower().replace(" ", "_") for k, v in entity_type_map.items()}
|
||||
logger.debug(f"Entity type map: {entity_type_map}\n search_results={search_results}")
|
||||
return entity_type_map
|
||||
|
||||
def _establish_nodes_relations_from_data(self, data, filters, entity_type_map):
|
||||
"""Eshtablish relations among the extracted nodes."""
|
||||
if self.config.graph_store.custom_prompt:
|
||||
messages = [
|
||||
{
|
||||
"role": "system",
|
||||
"content": EXTRACT_RELATIONS_PROMPT.replace("USER_ID", filters["user_id"]).replace(
|
||||
"CUSTOM_PROMPT", f"4. {self.config.graph_store.custom_prompt}"
|
||||
),
|
||||
},
|
||||
{"role": "user", "content": data},
|
||||
]
|
||||
else:
|
||||
messages = [
|
||||
{
|
||||
"role": "system",
|
||||
"content": EXTRACT_RELATIONS_PROMPT.replace("USER_ID", filters["user_id"]),
|
||||
},
|
||||
{
|
||||
"role": "user",
|
||||
"content": f"List of entities: {list(entity_type_map.keys())}. \n\nText: {data}",
|
||||
},
|
||||
]
|
||||
|
||||
_tools = [RELATIONS_TOOL]
|
||||
if self.llm_provider in ["azure_openai_structured", "openai_structured"]:
|
||||
_tools = [RELATIONS_STRUCT_TOOL]
|
||||
|
||||
extracted_entities = self.llm.generate_response(
|
||||
messages=messages,
|
||||
tools=_tools,
|
||||
)
|
||||
|
||||
entities = []
|
||||
if extracted_entities and extracted_entities.get("tool_calls"):
|
||||
entities = extracted_entities["tool_calls"][0].get("arguments", {}).get("entities", [])
|
||||
|
||||
entities = self._remove_spaces_from_entities(entities)
|
||||
logger.debug(f"Extracted entities: {entities}")
|
||||
return entities
|
||||
|
||||
def _search_graph_db(self, node_list, filters, top_k=100):
|
||||
"""Search similar nodes among and their respective incoming and outgoing relations."""
|
||||
result_relations = []
|
||||
|
||||
for node in node_list:
|
||||
n_embedding = self.embedding_model.embed(node)
|
||||
|
||||
# Build query based on whether agent_id is provided
|
||||
if filters.get("agent_id"):
|
||||
cypher_query = """
|
||||
CALL vector_search.search("memzero", $limit, $n_embedding)
|
||||
YIELD distance, node, similarity
|
||||
WITH node AS n, similarity
|
||||
WHERE n:Entity AND n.user_id = $user_id AND n.agent_id = $agent_id AND n.embedding IS NOT NULL AND similarity >= $threshold
|
||||
MATCH (n)-[r]->(m:Entity)
|
||||
RETURN n.name AS source, id(n) AS source_id, type(r) AS relationship, id(r) AS relation_id, m.name AS destination, id(m) AS destination_id, similarity
|
||||
UNION
|
||||
CALL vector_search.search("memzero", $limit, $n_embedding)
|
||||
YIELD distance, node, similarity
|
||||
WITH node AS n, similarity
|
||||
WHERE n:Entity AND n.user_id = $user_id AND n.agent_id = $agent_id AND n.embedding IS NOT NULL AND similarity >= $threshold
|
||||
MATCH (m:Entity)-[r]->(n)
|
||||
RETURN m.name AS source, id(m) AS source_id, type(r) AS relationship, id(r) AS relation_id, n.name AS destination, id(n) AS destination_id, similarity
|
||||
ORDER BY similarity DESC
|
||||
LIMIT $limit;
|
||||
"""
|
||||
params = {
|
||||
"n_embedding": n_embedding,
|
||||
"threshold": self.threshold,
|
||||
"user_id": filters["user_id"],
|
||||
"agent_id": filters["agent_id"],
|
||||
"limit": top_k,
|
||||
}
|
||||
else:
|
||||
cypher_query = """
|
||||
CALL vector_search.search("memzero", $limit, $n_embedding)
|
||||
YIELD distance, node, similarity
|
||||
WITH node AS n, similarity
|
||||
WHERE n:Entity AND n.user_id = $user_id AND n.embedding IS NOT NULL AND similarity >= $threshold
|
||||
MATCH (n)-[r]->(m:Entity)
|
||||
RETURN n.name AS source, id(n) AS source_id, type(r) AS relationship, id(r) AS relation_id, m.name AS destination, id(m) AS destination_id, similarity
|
||||
UNION
|
||||
CALL vector_search.search("memzero", $limit, $n_embedding)
|
||||
YIELD distance, node, similarity
|
||||
WITH node AS n, similarity
|
||||
WHERE n:Entity AND n.user_id = $user_id AND n.embedding IS NOT NULL AND similarity >= $threshold
|
||||
MATCH (m:Entity)-[r]->(n)
|
||||
RETURN m.name AS source, id(m) AS source_id, type(r) AS relationship, id(r) AS relation_id, n.name AS destination, id(n) AS destination_id, similarity
|
||||
ORDER BY similarity DESC
|
||||
LIMIT $limit;
|
||||
"""
|
||||
params = {
|
||||
"n_embedding": n_embedding,
|
||||
"threshold": self.threshold,
|
||||
"user_id": filters["user_id"],
|
||||
"limit": top_k,
|
||||
}
|
||||
|
||||
ans = self.graph.query(cypher_query, params=params)
|
||||
result_relations.extend(ans)
|
||||
|
||||
return result_relations
|
||||
|
||||
def _get_delete_entities_from_search_output(self, search_output, data, filters):
|
||||
"""Get the entities to be deleted from the search output."""
|
||||
search_output_string = format_entities(search_output)
|
||||
system_prompt, user_prompt = get_delete_messages(search_output_string, data, filters["user_id"])
|
||||
|
||||
_tools = [DELETE_MEMORY_TOOL_GRAPH]
|
||||
if self.llm_provider in ["azure_openai_structured", "openai_structured"]:
|
||||
_tools = [
|
||||
DELETE_MEMORY_STRUCT_TOOL_GRAPH,
|
||||
]
|
||||
|
||||
memory_updates = self.llm.generate_response(
|
||||
messages=[
|
||||
{"role": "system", "content": system_prompt},
|
||||
{"role": "user", "content": user_prompt},
|
||||
],
|
||||
tools=_tools,
|
||||
)
|
||||
to_be_deleted = []
|
||||
for item in memory_updates["tool_calls"]:
|
||||
if item["name"] == "delete_graph_memory":
|
||||
to_be_deleted.append(item["arguments"])
|
||||
# in case if it is not in the correct format
|
||||
to_be_deleted = self._remove_spaces_from_entities(to_be_deleted)
|
||||
logger.debug(f"Deleted relationships: {to_be_deleted}")
|
||||
return to_be_deleted
|
||||
|
||||
def _delete_entities(self, to_be_deleted, filters):
|
||||
"""Delete the entities from the graph."""
|
||||
user_id = filters["user_id"]
|
||||
agent_id = filters.get("agent_id", None)
|
||||
results = []
|
||||
|
||||
for item in to_be_deleted:
|
||||
source = item["source"]
|
||||
destination = item["destination"]
|
||||
relationship = item["relationship"]
|
||||
|
||||
# Build the agent filter for the query
|
||||
agent_filter = ""
|
||||
params = {
|
||||
"source_name": source,
|
||||
"dest_name": destination,
|
||||
"user_id": user_id,
|
||||
}
|
||||
|
||||
if agent_id:
|
||||
agent_filter = "AND n.agent_id = $agent_id AND m.agent_id = $agent_id"
|
||||
params["agent_id"] = agent_id
|
||||
|
||||
# Delete the specific relationship between nodes
|
||||
cypher = f"""
|
||||
MATCH (n:Entity {{name: $source_name, user_id: $user_id}})
|
||||
-[r:{relationship}]->
|
||||
(m:Entity {{name: $dest_name, user_id: $user_id}})
|
||||
WHERE 1=1 {agent_filter}
|
||||
DELETE r
|
||||
RETURN
|
||||
n.name AS source,
|
||||
m.name AS target,
|
||||
type(r) AS relationship
|
||||
"""
|
||||
|
||||
result = self.graph.query(cypher, params=params)
|
||||
results.append(result)
|
||||
|
||||
return results
|
||||
|
||||
# added Entity label to all nodes for vector search to work
|
||||
def _add_entities(self, to_be_added, filters, entity_type_map):
|
||||
"""Add the new entities to the graph. Merge the nodes if they already exist."""
|
||||
user_id = filters["user_id"]
|
||||
agent_id = filters.get("agent_id", None)
|
||||
results = []
|
||||
|
||||
for item in to_be_added:
|
||||
# entities
|
||||
source = item["source"]
|
||||
destination = item["destination"]
|
||||
relationship = item["relationship"]
|
||||
|
||||
# types
|
||||
source_type = entity_type_map.get(source, "__User__")
|
||||
destination_type = entity_type_map.get(destination, "__User__")
|
||||
|
||||
# embeddings
|
||||
source_embedding = self.embedding_model.embed(source)
|
||||
dest_embedding = self.embedding_model.embed(destination)
|
||||
|
||||
# search for the nodes with the closest embeddings
|
||||
source_node_search_result = self._search_source_node(source_embedding, filters, threshold=self.threshold)
|
||||
destination_node_search_result = self._search_destination_node(dest_embedding, filters, threshold=self.threshold)
|
||||
|
||||
# Prepare agent_id for node creation
|
||||
agent_id_clause = ""
|
||||
if agent_id:
|
||||
agent_id_clause = ", agent_id: $agent_id"
|
||||
|
||||
# TODO: Create a cypher query and common params for all the cases
|
||||
if not destination_node_search_result and source_node_search_result:
|
||||
cypher = f"""
|
||||
MATCH (source:Entity)
|
||||
WHERE id(source) = $source_id
|
||||
MERGE (destination:{destination_type}:Entity {{name: $destination_name, user_id: $user_id{agent_id_clause}}})
|
||||
ON CREATE SET
|
||||
destination.created = timestamp(),
|
||||
destination.embedding = $destination_embedding,
|
||||
destination:Entity
|
||||
MERGE (source)-[r:{relationship}]->(destination)
|
||||
ON CREATE SET
|
||||
r.created = timestamp()
|
||||
RETURN source.name AS source, type(r) AS relationship, destination.name AS target
|
||||
"""
|
||||
|
||||
params = {
|
||||
"source_id": source_node_search_result[0]["id(source_candidate)"],
|
||||
"destination_name": destination,
|
||||
"destination_embedding": dest_embedding,
|
||||
"user_id": user_id,
|
||||
}
|
||||
if agent_id:
|
||||
params["agent_id"] = agent_id
|
||||
|
||||
elif destination_node_search_result and not source_node_search_result:
|
||||
cypher = f"""
|
||||
MATCH (destination:Entity)
|
||||
WHERE id(destination) = $destination_id
|
||||
MERGE (source:{source_type}:Entity {{name: $source_name, user_id: $user_id{agent_id_clause}}})
|
||||
ON CREATE SET
|
||||
source.created = timestamp(),
|
||||
source.embedding = $source_embedding,
|
||||
source:Entity
|
||||
MERGE (source)-[r:{relationship}]->(destination)
|
||||
ON CREATE SET
|
||||
r.created = timestamp()
|
||||
RETURN source.name AS source, type(r) AS relationship, destination.name AS target
|
||||
"""
|
||||
|
||||
params = {
|
||||
"destination_id": destination_node_search_result[0]["id(destination_candidate)"],
|
||||
"source_name": source,
|
||||
"source_embedding": source_embedding,
|
||||
"user_id": user_id,
|
||||
}
|
||||
if agent_id:
|
||||
params["agent_id"] = agent_id
|
||||
|
||||
elif source_node_search_result and destination_node_search_result:
|
||||
cypher = f"""
|
||||
MATCH (source:Entity)
|
||||
WHERE id(source) = $source_id
|
||||
MATCH (destination:Entity)
|
||||
WHERE id(destination) = $destination_id
|
||||
MERGE (source)-[r:{relationship}]->(destination)
|
||||
ON CREATE SET
|
||||
r.created_at = timestamp(),
|
||||
r.updated_at = timestamp()
|
||||
RETURN source.name AS source, type(r) AS relationship, destination.name AS target
|
||||
"""
|
||||
params = {
|
||||
"source_id": source_node_search_result[0]["id(source_candidate)"],
|
||||
"destination_id": destination_node_search_result[0]["id(destination_candidate)"],
|
||||
"user_id": user_id,
|
||||
}
|
||||
if agent_id:
|
||||
params["agent_id"] = agent_id
|
||||
|
||||
else:
|
||||
cypher = f"""
|
||||
MERGE (n:{source_type}:Entity {{name: $source_name, user_id: $user_id{agent_id_clause}}})
|
||||
ON CREATE SET n.created = timestamp(), n.embedding = $source_embedding, n:Entity
|
||||
ON MATCH SET n.embedding = $source_embedding
|
||||
MERGE (m:{destination_type}:Entity {{name: $dest_name, user_id: $user_id{agent_id_clause}}})
|
||||
ON CREATE SET m.created = timestamp(), m.embedding = $dest_embedding, m:Entity
|
||||
ON MATCH SET m.embedding = $dest_embedding
|
||||
MERGE (n)-[rel:{relationship}]->(m)
|
||||
ON CREATE SET rel.created = timestamp()
|
||||
RETURN n.name AS source, type(rel) AS relationship, m.name AS target
|
||||
"""
|
||||
params = {
|
||||
"source_name": source,
|
||||
"dest_name": destination,
|
||||
"source_embedding": source_embedding,
|
||||
"dest_embedding": dest_embedding,
|
||||
"user_id": user_id,
|
||||
}
|
||||
if agent_id:
|
||||
params["agent_id"] = agent_id
|
||||
|
||||
result = self.graph.query(cypher, params=params)
|
||||
results.append(result)
|
||||
return results
|
||||
|
||||
def _remove_spaces_from_entities(self, entity_list):
|
||||
return remove_spaces_from_entities(entity_list, sanitize_relationship=True)
|
||||
|
||||
def _search_source_node(self, source_embedding, filters, threshold=0.9):
|
||||
"""Search for source nodes with similar embeddings."""
|
||||
user_id = filters["user_id"]
|
||||
agent_id = filters.get("agent_id", None)
|
||||
|
||||
if agent_id:
|
||||
cypher = """
|
||||
CALL vector_search.search("memzero", 1, $source_embedding)
|
||||
YIELD distance, node, similarity
|
||||
WITH node AS source_candidate, similarity
|
||||
WHERE source_candidate.user_id = $user_id
|
||||
AND source_candidate.agent_id = $agent_id
|
||||
AND similarity >= $threshold
|
||||
RETURN id(source_candidate);
|
||||
"""
|
||||
params = {
|
||||
"source_embedding": source_embedding,
|
||||
"user_id": user_id,
|
||||
"agent_id": agent_id,
|
||||
"threshold": threshold,
|
||||
}
|
||||
else:
|
||||
cypher = """
|
||||
CALL vector_search.search("memzero", 1, $source_embedding)
|
||||
YIELD distance, node, similarity
|
||||
WITH node AS source_candidate, similarity
|
||||
WHERE source_candidate.user_id = $user_id
|
||||
AND similarity >= $threshold
|
||||
RETURN id(source_candidate);
|
||||
"""
|
||||
params = {
|
||||
"source_embedding": source_embedding,
|
||||
"user_id": user_id,
|
||||
"threshold": threshold,
|
||||
}
|
||||
|
||||
result = self.graph.query(cypher, params=params)
|
||||
return result
|
||||
|
||||
def _search_destination_node(self, destination_embedding, filters, threshold=0.9):
|
||||
"""Search for destination nodes with similar embeddings."""
|
||||
user_id = filters["user_id"]
|
||||
agent_id = filters.get("agent_id", None)
|
||||
|
||||
if agent_id:
|
||||
cypher = """
|
||||
CALL vector_search.search("memzero", 1, $destination_embedding)
|
||||
YIELD distance, node, similarity
|
||||
WITH node AS destination_candidate, similarity
|
||||
WHERE node.user_id = $user_id
|
||||
AND node.agent_id = $agent_id
|
||||
AND similarity >= $threshold
|
||||
RETURN id(destination_candidate);
|
||||
"""
|
||||
params = {
|
||||
"destination_embedding": destination_embedding,
|
||||
"user_id": user_id,
|
||||
"agent_id": agent_id,
|
||||
"threshold": threshold,
|
||||
}
|
||||
else:
|
||||
cypher = """
|
||||
CALL vector_search.search("memzero", 1, $destination_embedding)
|
||||
YIELD distance, node, similarity
|
||||
WITH node AS destination_candidate, similarity
|
||||
WHERE node.user_id = $user_id
|
||||
AND similarity >= $threshold
|
||||
RETURN id(destination_candidate);
|
||||
"""
|
||||
params = {
|
||||
"destination_embedding": destination_embedding,
|
||||
"user_id": user_id,
|
||||
"threshold": threshold,
|
||||
}
|
||||
|
||||
result = self.graph.query(cypher, params=params)
|
||||
return result
|
||||
|
||||
|
||||
def _vector_index_exists(self, index_info, index_name):
|
||||
"""
|
||||
Check if a vector index exists, compatible with both Memgraph versions.
|
||||
|
||||
Args:
|
||||
index_info (dict): Index information from _fetch_existing_indexes
|
||||
index_name (str): Name of the index to check
|
||||
|
||||
Returns:
|
||||
bool: True if index exists, False otherwise
|
||||
"""
|
||||
vector_indexes = index_info.get("vector_index_exists", [])
|
||||
|
||||
# Check for index by name regardless of version-specific format differences
|
||||
return any(
|
||||
idx.get("index_name") == index_name or
|
||||
idx.get("index name") == index_name or
|
||||
idx.get("name") == index_name
|
||||
for idx in vector_indexes
|
||||
)
|
||||
|
||||
def _label_property_index_exists(self, index_info, label, property_name):
|
||||
"""
|
||||
Check if a label+property index exists, compatible with both versions.
|
||||
|
||||
Args:
|
||||
index_info (dict): Index information from _fetch_existing_indexes
|
||||
label (str): Label name
|
||||
property_name (str): Property name
|
||||
|
||||
Returns:
|
||||
bool: True if index exists, False otherwise
|
||||
"""
|
||||
indexes = index_info.get("index_exists", [])
|
||||
|
||||
return any(
|
||||
(idx.get("index type") == "label+property" or idx.get("index_type") == "label+property") and
|
||||
(idx.get("label") == label) and
|
||||
(idx.get("property") == property_name or property_name in str(idx.get("properties", "")))
|
||||
for idx in indexes
|
||||
)
|
||||
|
||||
def _label_index_exists(self, index_info, label):
|
||||
"""
|
||||
Check if a label index exists, compatible with both versions.
|
||||
|
||||
Args:
|
||||
index_info (dict): Index information from _fetch_existing_indexes
|
||||
label (str): Label name
|
||||
|
||||
Returns:
|
||||
bool: True if index exists, False otherwise
|
||||
"""
|
||||
indexes = index_info.get("index_exists", [])
|
||||
|
||||
return any(
|
||||
(idx.get("index type") == "label" or idx.get("index_type") == "label") and
|
||||
(idx.get("label") == label)
|
||||
for idx in indexes
|
||||
)
|
||||
|
||||
def _fetch_existing_indexes(self):
|
||||
"""
|
||||
Retrieves information about existing indexes and vector indexes in the Memgraph database.
|
||||
|
||||
Returns:
|
||||
dict: A dictionary containing lists of existing indexes and vector indexes.
|
||||
"""
|
||||
try:
|
||||
index_exists = list(self.graph.query("SHOW INDEX INFO;"))
|
||||
vector_index_exists = list(self.graph.query("SHOW VECTOR INDEX INFO;"))
|
||||
return {"index_exists": index_exists, "vector_index_exists": vector_index_exists}
|
||||
except Exception as e:
|
||||
logger.warning(f"Error fetching indexes: {e}. Returning empty index info.")
|
||||
return {"index_exists": [], "vector_index_exists": []}
|
||||
+132
-3
@@ -2,6 +2,7 @@ import logging
|
||||
import sqlite3
|
||||
import threading
|
||||
import uuid
|
||||
from datetime import datetime, timezone
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -14,6 +15,7 @@ class SQLiteManager:
|
||||
self._lock = threading.Lock()
|
||||
self._migrate_history_table()
|
||||
self._create_history_table()
|
||||
self._create_messages_table()
|
||||
|
||||
def _migrate_history_table(self) -> None:
|
||||
"""
|
||||
@@ -123,6 +125,28 @@ class SQLiteManager:
|
||||
logger.error(f"Failed to create history table: {e}")
|
||||
raise
|
||||
|
||||
def _create_messages_table(self) -> None:
|
||||
with self._lock:
|
||||
try:
|
||||
self.connection.execute("BEGIN")
|
||||
self.connection.execute(
|
||||
"""
|
||||
CREATE TABLE IF NOT EXISTS messages (
|
||||
id TEXT PRIMARY KEY,
|
||||
session_scope TEXT,
|
||||
role TEXT,
|
||||
content TEXT,
|
||||
name TEXT,
|
||||
created_at DATETIME
|
||||
)
|
||||
"""
|
||||
)
|
||||
self.connection.execute("COMMIT")
|
||||
except Exception as e:
|
||||
self.connection.execute("ROLLBACK")
|
||||
logger.error(f"Failed to create messages table: {e}")
|
||||
raise
|
||||
|
||||
def add_history(
|
||||
self,
|
||||
memory_id: str,
|
||||
@@ -166,6 +190,40 @@ class SQLiteManager:
|
||||
logger.error(f"Failed to add history record: {e}")
|
||||
raise
|
||||
|
||||
def batch_add_history(self, records: List[Dict[str, Any]]) -> None:
|
||||
with self._lock:
|
||||
try:
|
||||
self.connection.execute("BEGIN")
|
||||
self.connection.executemany(
|
||||
"""
|
||||
INSERT INTO history (
|
||||
id, memory_id, old_memory, new_memory, event,
|
||||
created_at, updated_at, is_deleted, actor_id, role
|
||||
)
|
||||
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
""",
|
||||
[
|
||||
(
|
||||
str(uuid.uuid4()),
|
||||
record.get("memory_id"),
|
||||
record.get("old_memory"),
|
||||
record.get("new_memory"),
|
||||
record.get("event"),
|
||||
record.get("created_at"),
|
||||
record.get("updated_at"),
|
||||
record.get("is_deleted", 0),
|
||||
record.get("actor_id"),
|
||||
record.get("role"),
|
||||
)
|
||||
for record in records
|
||||
],
|
||||
)
|
||||
self.connection.execute("COMMIT")
|
||||
except Exception as e:
|
||||
self.connection.execute("ROLLBACK")
|
||||
logger.error(f"Failed to batch add history records: {e}")
|
||||
raise
|
||||
|
||||
def get_history(self, memory_id: str) -> List[Dict[str, Any]]:
|
||||
with self._lock:
|
||||
cur = self.connection.execute(
|
||||
@@ -196,18 +254,89 @@ class SQLiteManager:
|
||||
for r in rows
|
||||
]
|
||||
|
||||
def save_messages(self, messages: List[Dict[str, Any]], session_scope: str) -> None:
|
||||
if not messages:
|
||||
return
|
||||
with self._lock:
|
||||
try:
|
||||
self.connection.execute("BEGIN")
|
||||
now = datetime.now(timezone.utc).isoformat()
|
||||
for message in messages:
|
||||
self.connection.execute(
|
||||
"""
|
||||
INSERT INTO messages (id, session_scope, role, content, name, created_at)
|
||||
VALUES (?, ?, ?, ?, ?, ?)
|
||||
""",
|
||||
(
|
||||
str(uuid.uuid4()),
|
||||
session_scope,
|
||||
message.get("role"),
|
||||
message.get("content"),
|
||||
message.get("name"),
|
||||
now,
|
||||
),
|
||||
)
|
||||
# Evict old messages beyond the most recent 10 for this scope.
|
||||
# Wrapped in a derived table to force SQLite to materialize the
|
||||
# ORDER BY before the outer NOT IN evaluates it.
|
||||
self.connection.execute(
|
||||
"""
|
||||
DELETE FROM messages WHERE session_scope = ? AND id NOT IN (
|
||||
SELECT id FROM (
|
||||
SELECT id FROM messages WHERE session_scope = ? ORDER BY created_at DESC LIMIT 10
|
||||
)
|
||||
)
|
||||
""",
|
||||
(session_scope, session_scope),
|
||||
)
|
||||
self.connection.execute("COMMIT")
|
||||
except Exception as e:
|
||||
self.connection.execute("ROLLBACK")
|
||||
logger.error(f"Failed to save messages: {e}")
|
||||
raise
|
||||
|
||||
def get_last_messages(self, session_scope: str, limit: int = 10) -> List[Dict[str, Any]]:
|
||||
with self._lock:
|
||||
# Subquery picks the latest N rows (DESC + LIMIT), outer query
|
||||
# re-sorts them chronologically (ASC) for the caller.
|
||||
cur = self.connection.execute(
|
||||
"""
|
||||
SELECT role, content, name, created_at FROM (
|
||||
SELECT role, content, name, created_at
|
||||
FROM messages
|
||||
WHERE session_scope = ?
|
||||
ORDER BY created_at DESC
|
||||
LIMIT ?
|
||||
) ORDER BY created_at ASC
|
||||
""",
|
||||
(session_scope, limit),
|
||||
)
|
||||
rows = cur.fetchall()
|
||||
|
||||
return [
|
||||
{
|
||||
"role": r[0],
|
||||
"content": r[1],
|
||||
"name": r[2],
|
||||
"created_at": r[3],
|
||||
}
|
||||
for r in rows
|
||||
]
|
||||
|
||||
def reset(self) -> None:
|
||||
"""Drop and recreate the history table."""
|
||||
"""Drop and recreate the history and messages tables."""
|
||||
with self._lock:
|
||||
try:
|
||||
self.connection.execute("BEGIN")
|
||||
self.connection.execute("DROP TABLE IF EXISTS history")
|
||||
self.connection.execute("DROP TABLE IF EXISTS messages")
|
||||
self.connection.execute("COMMIT")
|
||||
self._create_history_table()
|
||||
except Exception as e:
|
||||
self.connection.execute("ROLLBACK")
|
||||
logger.error(f"Failed to reset history table: {e}")
|
||||
logger.error(f"Failed to reset tables: {e}")
|
||||
raise
|
||||
self._create_history_table()
|
||||
self._create_messages_table()
|
||||
|
||||
def close(self) -> None:
|
||||
if self.connection:
|
||||
|
||||
@@ -169,9 +169,6 @@ def capture_event(event_name, memory_instance, additional_data=None):
|
||||
"collection": memory_instance.collection_name,
|
||||
"vector_size": memory_instance.embedding_model.config.embedding_dims,
|
||||
"history_store": "sqlite",
|
||||
"graph_store": f"{memory_instance.graph.__class__.__module__}.{memory_instance.graph.__class__.__name__}"
|
||||
if memory_instance.config.graph_store.config
|
||||
else None,
|
||||
"vector_store": f"{memory_instance.vector_store.__class__.__module__}.{memory_instance.vector_store.__class__.__name__}",
|
||||
"llm": f"{memory_instance.llm.__class__.__module__}.{memory_instance.llm.__class__.__name__}",
|
||||
"embedding_model": f"{memory_instance.embedding_model.__class__.__module__}.{memory_instance.embedding_model.__class__.__name__}",
|
||||
|
||||
@@ -0,0 +1,357 @@
|
||||
"""
|
||||
Entity extraction from text using spaCy NLP.
|
||||
|
||||
Extracts four types of entities from text:
|
||||
- **Proper nouns**: Capitalized multi-word sequences (person names, places, brands)
|
||||
- **Quoted text**: Text in single or double quotes (titles, specific terms)
|
||||
- **Noun compounds**: Multi-word noun phrases with specific modifiers (e.g., "machine learning")
|
||||
- **Noun fallback**: Single nouns from circumstantial compound patterns
|
||||
|
||||
Public API:
|
||||
extract_entities(text: str) -> List[Tuple[str, str]]
|
||||
|
||||
Internal:
|
||||
_extract_entities_from_doc(doc) -> List[Tuple[str, str]]
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import re
|
||||
from typing import List, Tuple
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# Words that are too generic to be useful as entity heads
|
||||
_GENERIC_HEADS = {
|
||||
"thing", "stuff", "way", "time", "experience", "situation", "case",
|
||||
"fact", "matter", "issue", "idea", "thought", "feeling", "place",
|
||||
"area", "part", "kind", "type", "sort", "lot", "bit", "day", "year",
|
||||
"week", "month", "moment", "instance", "example", "technique",
|
||||
"method", "approach", "process", "step", "tool", "result", "outcome",
|
||||
"goal", "task", "item", "topic", "scale", "size", "level", "degree",
|
||||
"amount", "number", "style", "look", "color", "colour", "shape",
|
||||
"form", "piece", "section", "side", "end", "edge", "surface", "point",
|
||||
}
|
||||
|
||||
# Modifiers that describe circumstance, not content
|
||||
_CIRCUMSTANTIAL_MODS = {
|
||||
"solo", "individual", "team", "group", "joint", "collaborative",
|
||||
"first", "last", "next", "previous", "final", "initial", "main", "side",
|
||||
}
|
||||
|
||||
# Adjectives too vague to make a compound entity specific
|
||||
_NON_SPECIFIC_ADJ = {
|
||||
"many", "few", "several", "some", "any", "all", "most", "more",
|
||||
"less", "much", "little", "enough", "various", "numerous", "multiple",
|
||||
"countless", "great", "good", "bad", "nice", "terrible", "awful",
|
||||
"awesome", "amazing", "wonderful", "horrible", "excellent", "poor",
|
||||
"best", "worst", "fine", "okay", "new", "old", "recent", "past",
|
||||
"future", "current", "previous", "next", "last", "first", "latest",
|
||||
"early", "late", "former", "modern", "ancient", "big", "small",
|
||||
"large", "tiny", "huge", "enormous", "long", "short", "tall", "high",
|
||||
"low", "wide", "narrow", "thick", "thin", "deep", "shallow",
|
||||
"similar", "different", "same", "other", "another", "such", "certain",
|
||||
"important", "main", "major", "minor", "key", "primary", "real",
|
||||
"actual", "true", "whole", "entire", "full", "complete", "total",
|
||||
"basic", "simple", "interesting", "boring", "exciting", "special",
|
||||
"particular", "general", "common", "unique", "rare", "typical",
|
||||
"usual", "normal", "regular", "possible", "likely", "potential",
|
||||
"available", "necessary", "only", "solo", "individual", "team",
|
||||
"group", "joint", "collaborative", "final", "initial", "side",
|
||||
}
|
||||
|
||||
# Generic tail words to strip from compound entities
|
||||
_GENERIC_ENDINGS = {
|
||||
"work", "works", "job", "jobs", "task", "tasks", "stuff", "things",
|
||||
"thing", "info", "information", "details", "data", "content",
|
||||
"material", "materials", "activities", "activity", "efforts", "effort",
|
||||
"options", "option", "choices", "choice", "results", "result",
|
||||
"output", "outputs", "products", "product", "items", "item",
|
||||
}
|
||||
|
||||
# Capitalized single words that are too generic to be proper nouns
|
||||
_GENERIC_CAPS = {
|
||||
"works", "items", "things", "stuff", "resources", "options", "tips",
|
||||
"ideas", "steps", "ways", "methods", "tools", "features", "benefits",
|
||||
"examples", "details", "notes", "instructions", "guidelines",
|
||||
"recommendations", "suggestions", "overview", "summary", "conclusion",
|
||||
"introduction", "pros", "cons", "advantages", "disadvantages",
|
||||
}
|
||||
|
||||
# Markdown/formatting markers to skip during extraction
|
||||
_FORMATTING_MARKERS = {"*", "-", "+", "\u2022", "\u2013", "\u2014", "#", "##", "###", "**", "__"}
|
||||
|
||||
|
||||
def _is_sentence_start(tokens: list, idx: int) -> bool:
|
||||
"""Check if a token is at the start of a sentence or after formatting."""
|
||||
if idx == 0:
|
||||
return True
|
||||
tok = tokens[idx]
|
||||
if tok.is_sent_start:
|
||||
return True
|
||||
prev = tokens[idx - 1].text
|
||||
return prev in ".!?:" or prev in _FORMATTING_MARKERS or "\n" in prev
|
||||
|
||||
|
||||
def _strip_generic_ending(toks: list) -> list:
|
||||
"""Remove generic trailing words from compound token sequences."""
|
||||
if len(toks) <= 1:
|
||||
return toks
|
||||
last = toks[-1].lemma_.lower() if hasattr(toks[-1], "lemma_") else toks[-1].lower()
|
||||
return toks[:-1] if last in _GENERIC_ENDINGS and len(toks) > 2 else toks
|
||||
|
||||
|
||||
def _lemmatize_compound(toks: list) -> str:
|
||||
"""Join compound tokens, lemmatizing nouns."""
|
||||
return " ".join(t.lemma_ if t.pos_ == "NOUN" else t.text for t in toks)
|
||||
|
||||
|
||||
def _has_artifacts(txt: str) -> bool:
|
||||
"""Check for formatting artifacts that indicate non-entity text."""
|
||||
return any(
|
||||
[
|
||||
"**" in txt or "__" in txt or ":*" in txt,
|
||||
re.search(r"\s\*\s|\s\*$|^\*\s", txt),
|
||||
" " in txt or "\n" in txt or "\t" in txt,
|
||||
len(txt) > 100,
|
||||
txt.startswith(("\u2022", "-", "+", "\u2013", "\u2014")),
|
||||
]
|
||||
)
|
||||
|
||||
|
||||
def extract_entities(text: str) -> List[Tuple[str, str]]:
|
||||
"""Extract named entities, quoted text, and noun compounds from text.
|
||||
|
||||
This is the public API that accepts a string. It loads the spaCy model
|
||||
internally and delegates to _extract_entities_from_doc().
|
||||
|
||||
Args:
|
||||
text: Input text to extract entities from.
|
||||
|
||||
Returns:
|
||||
Deduplicated list of (entity_type, entity_text) tuples.
|
||||
Entity types: PROPER, QUOTED, COMPOUND, NOUN.
|
||||
Returns empty list if spaCy is unavailable.
|
||||
"""
|
||||
from mem0.utils.spacy_models import get_nlp_full
|
||||
|
||||
nlp = get_nlp_full()
|
||||
if nlp is None:
|
||||
return []
|
||||
|
||||
doc = nlp(text)
|
||||
return _extract_entities_from_doc(doc)
|
||||
|
||||
|
||||
def extract_entities_batch(texts: List[str], batch_size: int = 32) -> List[List[Tuple[str, str]]]:
|
||||
"""Extract entities from multiple texts using spaCy's nlp.pipe() for batched NER.
|
||||
|
||||
Uses spaCy's efficient batch processing pipeline instead of calling
|
||||
nlp() individually per text. Significantly faster for multiple texts.
|
||||
|
||||
Args:
|
||||
texts: List of input texts to extract entities from.
|
||||
batch_size: Number of texts to process in each spaCy batch.
|
||||
|
||||
Returns:
|
||||
List of entity lists, one per input text. Each entity list contains
|
||||
(entity_type, entity_text) tuples. Returns list of empty lists if
|
||||
spaCy is unavailable.
|
||||
"""
|
||||
if not texts:
|
||||
return []
|
||||
|
||||
from mem0.utils.spacy_models import get_nlp_full
|
||||
|
||||
nlp = get_nlp_full()
|
||||
if nlp is None:
|
||||
return [[] for _ in texts]
|
||||
|
||||
results = []
|
||||
for doc in nlp.pipe(texts, batch_size=batch_size):
|
||||
results.append(_extract_entities_from_doc(doc))
|
||||
return results
|
||||
|
||||
|
||||
def _extract_entities_from_doc(doc) -> List[Tuple[str, str]]:
|
||||
"""Extract entities from a spaCy Doc object.
|
||||
|
||||
Ported from platform's shared.core.utils.entity_extraction.extract_entities().
|
||||
"""
|
||||
entities: List[Tuple[str, str]] = []
|
||||
text = doc.text
|
||||
tokens = list(doc)
|
||||
|
||||
# === PROPER NOUN SEQUENCES ===
|
||||
i = 0
|
||||
while i < len(tokens):
|
||||
tok = tokens[i]
|
||||
if tok.text in _FORMATTING_MARKERS:
|
||||
i += 1
|
||||
continue
|
||||
is_cap = tok.text and tok.text[0].isupper()
|
||||
is_label = i + 1 < len(tokens) and tokens[i + 1].text == ":"
|
||||
|
||||
if is_cap and not is_label and tok.pos_ in {"PROPN", "NOUN", "ADJ"}:
|
||||
seq = [(tok, i)]
|
||||
j = i + 1
|
||||
while j < len(tokens):
|
||||
t = tokens[j]
|
||||
if (t.text and t.text[0].isupper()) or t.text.lower() in {
|
||||
"'s", "of", "the", "in", "and", "for", "at", "is",
|
||||
}:
|
||||
seq.append((t, j))
|
||||
j += 1
|
||||
else:
|
||||
break
|
||||
# Strip trailing function words
|
||||
while seq and seq[-1][0].text.lower() in {"of", "the", "in", "and", "for", "at", "is", "'s"}:
|
||||
seq.pop()
|
||||
if seq:
|
||||
has_mid_cap = any(
|
||||
not _is_sentence_start(tokens, idx)
|
||||
for (t, idx) in seq
|
||||
if t.text[0].isupper() and t.text.lower() not in {"'s", "of", "the", "in", "and", "for", "at", "is"}
|
||||
)
|
||||
if has_mid_cap:
|
||||
phrase = "".join(t.text_with_ws for (t, idx) in seq).strip()
|
||||
if len(phrase) > 2:
|
||||
entities.append(("PROPER", phrase))
|
||||
i = j
|
||||
else:
|
||||
i += 1
|
||||
|
||||
# === QUOTED TEXT ===
|
||||
for m in re.finditer(r'"([^"]+)"', text):
|
||||
if len(m.group(1).strip()) > 2:
|
||||
entities.append(("QUOTED", m.group(1).strip()))
|
||||
for m in re.finditer(r"(?:^|[\s\(\[{,;])'([^']+)'(?=[\s\.,;:!?\)\]]|$)", text):
|
||||
if len(m.group(1).strip()) > 2:
|
||||
entities.append(("QUOTED", m.group(1).strip()))
|
||||
|
||||
# === NOUN-NOUN COMPOUNDS ===
|
||||
for chunk in doc.noun_chunks:
|
||||
chunk_tokens = list(chunk)
|
||||
split_indices: list = []
|
||||
poss_splits: list = []
|
||||
for idx, tok in enumerate(chunk_tokens):
|
||||
if tok.dep_ == "case" and tok.text in {"'s", "\u2019s", "'"}:
|
||||
split_indices.append(idx)
|
||||
poss_splits.append(idx)
|
||||
elif tok.pos_ == "PUNCT" and tok.text in {"'", '"', "\u2018", "\u2019", "\u201c", "\u201d"}:
|
||||
split_indices.append(idx)
|
||||
|
||||
if split_indices:
|
||||
groups: list = []
|
||||
prev = 0
|
||||
for split_idx in split_indices:
|
||||
if split_idx > prev:
|
||||
groups.append(chunk_tokens[prev:split_idx])
|
||||
if split_idx in poss_splits:
|
||||
next_split = next((s for s in split_indices if s > split_idx), None)
|
||||
owned = chunk_tokens[split_idx + 1: next_split if next_split else len(chunk_tokens)]
|
||||
if owned:
|
||||
first_content = next((t for t in owned if t.pos_ not in {"PUNCT", "PART"}), None)
|
||||
if not (first_content and first_content.text and first_content.text[0].isupper()):
|
||||
prev = next_split if next_split else len(chunk_tokens)
|
||||
continue
|
||||
prev = split_idx + 1
|
||||
if prev < len(chunk_tokens):
|
||||
groups.append(chunk_tokens[prev:])
|
||||
else:
|
||||
groups = [chunk_tokens]
|
||||
|
||||
for group in groups:
|
||||
if not group:
|
||||
continue
|
||||
head = next((t for t in reversed(group) if t.pos_ in {"NOUN", "PROPN"}), None)
|
||||
if not head:
|
||||
continue
|
||||
head_generic = head.lemma_.lower() in _GENERIC_HEADS
|
||||
content = [
|
||||
t
|
||||
for t in group
|
||||
if t.pos_ not in {"DET", "PRON", "PUNCT", "PART", "ADP", "SCONJ", "NUM"} and (t.pos_ == "ADJ" or not t.is_stop)
|
||||
]
|
||||
if not content:
|
||||
continue
|
||||
|
||||
compound_toks = [t for t in content if t.dep_ == "compound"]
|
||||
adj_toks = [t for t in content if t.pos_ == "ADJ" or t.dep_ == "amod"]
|
||||
has_spec_adj = any(t.lemma_.lower() not in _NON_SPECIFIC_ADJ for t in adj_toks)
|
||||
if head_generic and not has_spec_adj and not compound_toks:
|
||||
continue
|
||||
|
||||
if compound_toks:
|
||||
is_circ = any(t.lemma_.lower() in _CIRCUMSTANTIAL_MODS for t in compound_toks)
|
||||
if is_circ:
|
||||
val = head.lemma_ if head.pos_ == "NOUN" else head.text
|
||||
if len(val) > 2:
|
||||
entities.append(("NOUN", val))
|
||||
else:
|
||||
filtered = _strip_generic_ending(
|
||||
[t for t in content if not (t.pos_ == "ADJ" and t.lemma_.lower() in _NON_SPECIFIC_ADJ)]
|
||||
)
|
||||
if filtered:
|
||||
phrase = _lemmatize_compound(filtered)
|
||||
if len(phrase) > 3 and " " in phrase:
|
||||
entities.append(("COMPOUND", phrase))
|
||||
elif len(content) > 1 and has_spec_adj:
|
||||
filtered = _strip_generic_ending(
|
||||
[t for t in content if not ((t.pos_ == "ADJ" or t.dep_ == "amod") and t.lemma_.lower() in _NON_SPECIFIC_ADJ)]
|
||||
)
|
||||
if filtered:
|
||||
phrase = _lemmatize_compound(filtered)
|
||||
if len(phrase) > 3 and " " in phrase:
|
||||
entities.append(("COMPOUND", phrase))
|
||||
|
||||
# === FALLBACK: Mis-tagged VERB heads ===
|
||||
processed = {e[1].lower() for e in entities if e[0] == "COMPOUND"}
|
||||
generic_verb_heads = _GENERIC_HEADS | {"find", "buy", "purchase", "sale", "deal", "trip", "visit"}
|
||||
|
||||
def collect_compounds(head):
|
||||
return [t for t in doc if t.head == head and t.dep_ == "compound"]
|
||||
|
||||
for tok in doc:
|
||||
if tok.pos_ == "VERB" and tok.dep_ in {"pobj", "dobj", "nsubj"}:
|
||||
comps = sorted(collect_compounds(tok), key=lambda t: t.i)
|
||||
if comps:
|
||||
phrase_toks = comps if tok.lemma_.lower() in generic_verb_heads else comps + [tok]
|
||||
phrase = " ".join(t.text for t in phrase_toks)
|
||||
if phrase.lower() not in processed and len(phrase) > 3 and " " in phrase:
|
||||
entities.append(("COMPOUND", phrase))
|
||||
processed.add(phrase.lower())
|
||||
|
||||
# === DEDUPLICATION & CLEANUP ===
|
||||
seen: set = set()
|
||||
deduped = []
|
||||
for t, e in entities:
|
||||
k = e.lower().strip()
|
||||
if k not in seen and len(k) > 2:
|
||||
seen.add(k)
|
||||
deduped.append((t, e))
|
||||
|
||||
cleaned: List[Tuple[str, str]] = []
|
||||
for etype, etext in deduped:
|
||||
txt = re.sub(r"^\*+\s*|\s*\*+$", "", etext.strip())
|
||||
txt = re.sub(r"\s*:+$", "", txt)
|
||||
txt = re.sub(r"^\d+\s*\.\s*", "", txt)
|
||||
if not txt or len(txt) <= 2 or _has_artifacts(txt):
|
||||
continue
|
||||
if etype == "PROPER" and " " not in txt and txt.lower() in _GENERIC_CAPS:
|
||||
continue
|
||||
cleaned.append((etype, txt))
|
||||
|
||||
# Keep best type per entity (PROPER > COMPOUND > QUOTED > NOUN)
|
||||
type_pri = {"PROPER": 0, "COMPOUND": 1, "QUOTED": 2, "NOUN": 3, "VERB": 4}
|
||||
best: dict = {}
|
||||
for t, e in cleaned:
|
||||
k = e.lower()
|
||||
if k not in best or type_pri.get(t, 99) < type_pri.get(best[k][0], 99):
|
||||
best[k] = (t, e)
|
||||
deduped = list(best.values())
|
||||
|
||||
# Remove entities that are substrings of longer entities
|
||||
all_lower = [e[1].lower() for e in deduped]
|
||||
return [(t, e) for t, e in deduped if not any(e.lower() != o and e.lower() in o for o in all_lower)]
|
||||
@@ -209,30 +209,6 @@ class VectorStoreFactory:
|
||||
return instance
|
||||
|
||||
|
||||
class GraphStoreFactory:
|
||||
"""
|
||||
Factory for creating MemoryGraph instances for different graph store providers.
|
||||
Usage: GraphStoreFactory.create(provider_name, config)
|
||||
"""
|
||||
|
||||
provider_to_class = {
|
||||
"memgraph": "mem0.memory.memgraph_memory.MemoryGraph",
|
||||
"neptune": "mem0.graphs.neptune.neptunegraph.MemoryGraph",
|
||||
"neptunedb": "mem0.graphs.neptune.neptunedb.MemoryGraph",
|
||||
"kuzu": "mem0.memory.kuzu_memory.MemoryGraph",
|
||||
"apache_age": "mem0.memory.apache_age_memory.MemoryGraph",
|
||||
"default": "mem0.memory.graph_memory.MemoryGraph",
|
||||
}
|
||||
|
||||
@classmethod
|
||||
def create(cls, provider_name, config):
|
||||
class_type = cls.provider_to_class.get(provider_name, cls.provider_to_class["default"])
|
||||
try:
|
||||
GraphClass = load_class(class_type)
|
||||
except (ImportError, AttributeError) as e:
|
||||
raise ImportError(f"Could not import MemoryGraph for provider '{provider_name}': {e}")
|
||||
return GraphClass(config)
|
||||
|
||||
|
||||
class RerankerFactory:
|
||||
"""
|
||||
|
||||
@@ -0,0 +1,50 @@
|
||||
"""
|
||||
BM25 lemmatization for consistent keyword matching.
|
||||
|
||||
Uses spaCy's lemmatizer for better handling of:
|
||||
- Verb forms: attending/attends/attended -> attend
|
||||
- Comparatives/superlatives: older/oldest -> old
|
||||
- Plurals: memories -> memory
|
||||
- Avoids over-stemming: organization != organize
|
||||
|
||||
Also includes original -ing forms alongside lemmas to handle cases
|
||||
where spaCy's context-dependent lemmatization produces inconsistent
|
||||
results (e.g., "meeting" as noun vs verb -> different lemmas).
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def lemmatize_for_bm25(text: str) -> str:
|
||||
"""Lemmatize text for BM25 matching.
|
||||
|
||||
Returns space-joined lemmas for full-text search. Falls back to
|
||||
the original text if spaCy is unavailable.
|
||||
"""
|
||||
from mem0.utils.spacy_models import get_nlp_lemma
|
||||
|
||||
nlp = get_nlp_lemma()
|
||||
if nlp is None:
|
||||
return text
|
||||
|
||||
doc = nlp(text.lower())
|
||||
tokens = []
|
||||
|
||||
for token in doc:
|
||||
if token.is_punct or token.is_stop:
|
||||
continue
|
||||
|
||||
lemma = token.lemma_
|
||||
if lemma.isalnum():
|
||||
tokens.append(lemma)
|
||||
|
||||
# Also add original if it ends in -ing and differs from lemma.
|
||||
# This handles noun/verb ambiguity (meeting/meet, attending/attend).
|
||||
if token.text.endswith("ing") and token.text != lemma and token.text.isalnum():
|
||||
tokens.append(token.text)
|
||||
|
||||
return " ".join(tokens)
|
||||
@@ -0,0 +1,126 @@
|
||||
"""
|
||||
Scoring utilities for hybrid retrieval.
|
||||
|
||||
Provides:
|
||||
- **BM25 normalization**: Sigmoid normalization of raw BM25 scores to [0, 1].
|
||||
- **BM25 parameter selection**: Query-length-adaptive sigmoid parameters.
|
||||
- **Additive scoring**: Combined scoring with semantic + BM25 + entity boost.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
|
||||
def get_bm25_params(query: str, *, lemmatized: Optional[str] = None) -> tuple:
|
||||
"""Get BM25 sigmoid parameters based on query length.
|
||||
|
||||
Longer queries tend to have higher raw BM25 scores, so we adjust
|
||||
the sigmoid midpoint and steepness accordingly.
|
||||
|
||||
Returns:
|
||||
(midpoint, steepness) for sigmoid normalization.
|
||||
"""
|
||||
if lemmatized is None:
|
||||
from mem0.utils.lemmatization import lemmatize_for_bm25
|
||||
|
||||
lemmatized = lemmatize_for_bm25(query)
|
||||
num_terms = len(lemmatized.split()) if lemmatized else 1
|
||||
|
||||
if num_terms <= 3:
|
||||
return 5.0, 0.7
|
||||
elif num_terms <= 6:
|
||||
return 7.0, 0.6
|
||||
elif num_terms <= 9:
|
||||
return 9.0, 0.5
|
||||
elif num_terms <= 15:
|
||||
return 10.0, 0.5
|
||||
else:
|
||||
return 12.0, 0.5
|
||||
|
||||
|
||||
def normalize_bm25(raw_score: float, midpoint: float, steepness: float) -> float:
|
||||
"""Normalize BM25 score to [0, 1] using logistic sigmoid.
|
||||
|
||||
Args:
|
||||
raw_score: Raw BM25 score (unbounded, typically 0-20+).
|
||||
midpoint: Score at which sigmoid outputs 0.5.
|
||||
steepness: Controls how quickly sigmoid transitions.
|
||||
|
||||
Returns:
|
||||
Normalized score in range [0, 1].
|
||||
"""
|
||||
return 1.0 / (1.0 + math.exp(-steepness * (raw_score - midpoint)))
|
||||
|
||||
|
||||
ENTITY_BOOST_WEIGHT = 0.5
|
||||
|
||||
|
||||
def score_and_rank(
|
||||
semantic_results: List[Dict[str, Any]],
|
||||
bm25_scores: Dict[str, float],
|
||||
entity_boosts: Dict[str, float],
|
||||
threshold: float,
|
||||
top_k: int,
|
||||
) -> List[Dict[str, Any]]:
|
||||
"""Score candidates additively and return top-k results.
|
||||
|
||||
For each candidate:
|
||||
semantic_score is taken from the result's score field.
|
||||
combined = (semantic + bm25 + entity_boost) / max_possible
|
||||
|
||||
Threshold gates the semantic score BEFORE combining -- candidates
|
||||
below the threshold are excluded even if BM25/entity would boost them.
|
||||
|
||||
The divisor adapts based on which signals are active:
|
||||
- Semantic only: max_possible = 1.0
|
||||
- Semantic + BM25: max_possible = 2.0
|
||||
- Semantic + BM25 + entity: max_possible = 2.5
|
||||
- Semantic + entity (no BM25): max_possible = 1.5
|
||||
|
||||
Returns:
|
||||
List of scored result dicts sorted by combined score descending.
|
||||
"""
|
||||
has_bm25 = bool(bm25_scores)
|
||||
has_entity = bool(entity_boosts)
|
||||
|
||||
max_possible = 1.0
|
||||
if has_bm25:
|
||||
max_possible += 1.0
|
||||
if has_entity:
|
||||
max_possible += ENTITY_BOOST_WEIGHT
|
||||
|
||||
scored: List[Dict[str, Any]] = []
|
||||
|
||||
for result in semantic_results:
|
||||
mem_id = result.get("id")
|
||||
if mem_id is None:
|
||||
continue
|
||||
|
||||
semantic_score = result.get("score", 0.0)
|
||||
if semantic_score < threshold:
|
||||
continue
|
||||
|
||||
mem_id_str = str(mem_id)
|
||||
bm25_score = bm25_scores.get(mem_id_str, 0.0)
|
||||
entity_boost = entity_boosts.get(mem_id_str, 0.0)
|
||||
|
||||
raw_combined = semantic_score + bm25_score + entity_boost
|
||||
combined = min(raw_combined / max_possible, 1.0)
|
||||
|
||||
scored.append(
|
||||
{
|
||||
"id": mem_id_str,
|
||||
"score": combined,
|
||||
"score_breakdown": {
|
||||
"semantic": semantic_score,
|
||||
"bm25": bm25_score,
|
||||
"entity_boost": entity_boost,
|
||||
},
|
||||
"payload": result.get("payload"),
|
||||
}
|
||||
)
|
||||
|
||||
scored.sort(key=lambda x: x["score"], reverse=True)
|
||||
return scored[:top_k]
|
||||
@@ -0,0 +1,91 @@
|
||||
"""
|
||||
Shared spaCy model loader.
|
||||
|
||||
Consolidates spaCy model loading into a single module so that
|
||||
entity_extraction and lemmatization share one instance instead of
|
||||
each loading their own copy from disk.
|
||||
"""
|
||||
|
||||
import logging
|
||||
import threading
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_nlp_full = None
|
||||
_nlp_lemma = None
|
||||
_load_failed_full = False
|
||||
_load_failed_lemma = False
|
||||
_lock = threading.Lock()
|
||||
|
||||
|
||||
def _ensure_model_available():
|
||||
"""Download en_core_web_sm if spaCy is installed but model is missing."""
|
||||
try:
|
||||
import spacy
|
||||
except ImportError:
|
||||
raise ImportError(
|
||||
"spaCy is not installed. Install it with: pip install mem0ai[nlp]"
|
||||
)
|
||||
|
||||
if not spacy.util.is_package("en_core_web_sm"):
|
||||
logger.info("Downloading spaCy model en_core_web_sm...")
|
||||
try:
|
||||
from spacy.cli import download
|
||||
|
||||
download("en_core_web_sm")
|
||||
logger.info("spaCy model en_core_web_sm downloaded successfully")
|
||||
except Exception as e:
|
||||
raise RuntimeError(
|
||||
f"Failed to download spaCy model en_core_web_sm: {e}. "
|
||||
"Please install manually: python -m spacy download en_core_web_sm"
|
||||
) from e
|
||||
|
||||
|
||||
def get_nlp_full():
|
||||
"""Return spaCy model with all pipelines (NER, tagger, etc.) for entity extraction."""
|
||||
global _nlp_full, _load_failed_full
|
||||
if _load_failed_full:
|
||||
return None
|
||||
if _nlp_full is not None:
|
||||
return _nlp_full
|
||||
with _lock:
|
||||
if _nlp_full is not None:
|
||||
return _nlp_full
|
||||
if _load_failed_full:
|
||||
return None
|
||||
try:
|
||||
_ensure_model_available()
|
||||
import spacy
|
||||
|
||||
_nlp_full = spacy.load("en_core_web_sm")
|
||||
logger.info("spaCy full model loaded")
|
||||
except Exception as e:
|
||||
logger.warning(f"Failed to load spaCy full model: {e}")
|
||||
_load_failed_full = True
|
||||
return None
|
||||
return _nlp_full
|
||||
|
||||
|
||||
def get_nlp_lemma():
|
||||
"""Return spaCy model with only lemmatizer for BM25 text processing."""
|
||||
global _nlp_lemma, _load_failed_lemma
|
||||
if _load_failed_lemma:
|
||||
return None
|
||||
if _nlp_lemma is not None:
|
||||
return _nlp_lemma
|
||||
with _lock:
|
||||
if _nlp_lemma is not None:
|
||||
return _nlp_lemma
|
||||
if _load_failed_lemma:
|
||||
return None
|
||||
try:
|
||||
_ensure_model_available()
|
||||
import spacy
|
||||
|
||||
_nlp_lemma = spacy.load("en_core_web_sm", disable=["ner", "parser"])
|
||||
logger.info("spaCy lemma model loaded")
|
||||
except Exception as e:
|
||||
logger.warning(f"Failed to load spaCy lemma model: {e}")
|
||||
_load_failed_lemma = True
|
||||
return None
|
||||
return _nlp_lemma
|
||||
@@ -246,6 +246,34 @@ class AzureAISearch(VectorStoreBase):
|
||||
results.append(OutputData(id=result["id"], score=result["@search.score"], payload=payload))
|
||||
return results
|
||||
|
||||
def keyword_search(self, query, top_k=5, filters=None):
|
||||
"""Search for memories using keyword/BM25 text matching (no vector queries).
|
||||
|
||||
Args:
|
||||
query (str): The text query to search for.
|
||||
top_k (int): Maximum number of results to return. Defaults to 5.
|
||||
filters (Dict, optional): Filters to apply to the search.
|
||||
|
||||
Returns:
|
||||
List[OutputData]: Search results with id, score, and payload.
|
||||
"""
|
||||
filter_expression = None
|
||||
if filters:
|
||||
filter_expression = self._build_filter_expression(filters)
|
||||
|
||||
search_results = self.search_client.search(
|
||||
search_text=query,
|
||||
filter=filter_expression,
|
||||
top=top_k,
|
||||
search_fields=["payload"],
|
||||
)
|
||||
|
||||
results = []
|
||||
for result in search_results:
|
||||
payload = json.loads(extract_json(result["payload"]))
|
||||
results.append(OutputData(id=result["id"], score=result["@search.score"], payload=payload))
|
||||
return results
|
||||
|
||||
def delete(self, vector_id):
|
||||
"""
|
||||
Delete a vector by ID.
|
||||
|
||||
@@ -178,17 +178,32 @@ class AzureMySQL(VectorStoreBase):
|
||||
dims = vector_size or self.embedding_model_dims
|
||||
|
||||
with self._get_cursor(commit=True) as cur:
|
||||
# Create table with vector column
|
||||
# Create table with vector column and a generated column for fulltext keyword search
|
||||
cur.execute(f"""
|
||||
CREATE TABLE IF NOT EXISTS `{table_name}` (
|
||||
id VARCHAR(255) PRIMARY KEY,
|
||||
vector JSON,
|
||||
payload JSON,
|
||||
text_lemmatized VARCHAR(1000) GENERATED ALWAYS AS
|
||||
(CAST(payload->>'$.text_lemmatized' AS CHAR(1000))) STORED,
|
||||
INDEX idx_payload_keys ((CAST(payload AS CHAR(255)) ARRAY))
|
||||
)
|
||||
""")
|
||||
logger.info(f"Created collection '{table_name}' with vector dimension {dims}")
|
||||
|
||||
# Add FULLTEXT index on text_lemmatized for keyword_search()
|
||||
try:
|
||||
cur.execute(f"""
|
||||
CREATE FULLTEXT INDEX ft_text_lemmatized
|
||||
ON `{table_name}` (text_lemmatized)
|
||||
""")
|
||||
logger.info(f"Created FULLTEXT index on '{table_name}.text_lemmatized'")
|
||||
except Exception as e:
|
||||
logger.debug(
|
||||
f"Could not create FULLTEXT index on '{table_name}.text_lemmatized': {e}. "
|
||||
"It may already exist or FULLTEXT may not be supported."
|
||||
)
|
||||
|
||||
def insert(self, vectors: List[List[float]], payloads: Optional[List[Dict]] = None, ids: Optional[List[str]] = None):
|
||||
"""
|
||||
Insert vectors into the collection.
|
||||
@@ -300,6 +315,61 @@ class AzureMySQL(VectorStoreBase):
|
||||
for r in scored_results
|
||||
]
|
||||
|
||||
def keyword_search(self, query, top_k=5, filters=None):
|
||||
"""
|
||||
Search for memories using MySQL FULLTEXT search via MATCH() AGAINST().
|
||||
|
||||
This method attempts to use a FULLTEXT index on the text_lemmatized column.
|
||||
If the column or index does not exist, it returns None gracefully.
|
||||
|
||||
Args:
|
||||
query (str): The text query for keyword-based search.
|
||||
top_k (int, optional): Number of results to return. Defaults to 5.
|
||||
filters (dict, optional): Filters to apply to the search. Defaults to None.
|
||||
|
||||
Returns:
|
||||
list: Search results in the same format as search(), or None if FULLTEXT
|
||||
search is not supported on this collection.
|
||||
"""
|
||||
try:
|
||||
filter_conditions = []
|
||||
filter_params = []
|
||||
|
||||
if filters:
|
||||
for k, v in filters.items():
|
||||
filter_conditions.append("JSON_EXTRACT(payload, %s) = %s")
|
||||
filter_params.extend([f"$.{k}", json.dumps(v)])
|
||||
|
||||
filter_clause = ""
|
||||
if filter_conditions:
|
||||
filter_clause = " AND " + " AND ".join(filter_conditions)
|
||||
|
||||
with self._get_cursor() as cur:
|
||||
query_sql = f"""
|
||||
SELECT id, payload,
|
||||
MATCH(text_lemmatized) AGAINST(%s IN NATURAL LANGUAGE MODE) AS score
|
||||
FROM `{self.collection_name}`
|
||||
WHERE MATCH(text_lemmatized) AGAINST(%s IN NATURAL LANGUAGE MODE)
|
||||
{filter_clause}
|
||||
ORDER BY score DESC
|
||||
LIMIT %s
|
||||
"""
|
||||
params = [query, query] + filter_params + [top_k]
|
||||
cur.execute(query_sql, params)
|
||||
results = cur.fetchall()
|
||||
|
||||
return [
|
||||
OutputData(
|
||||
id=r['id'],
|
||||
score=float(r['score']),
|
||||
payload=json.loads(r['payload']) if isinstance(r['payload'], str) else r['payload'],
|
||||
)
|
||||
for r in results
|
||||
]
|
||||
except Exception as e:
|
||||
logger.debug(f"Keyword search not available for collection {self.collection_name}: {e}")
|
||||
return None
|
||||
|
||||
def delete(self, vector_id: str):
|
||||
"""
|
||||
Delete a vector by ID.
|
||||
|
||||
@@ -27,6 +27,7 @@ try:
|
||||
VectorIndex,
|
||||
)
|
||||
from pymochow.model.table import (
|
||||
BM25SearchRequest,
|
||||
FloatVector,
|
||||
Partition,
|
||||
Row,
|
||||
@@ -227,6 +228,48 @@ class BaiduDB(VectorStoreBase):
|
||||
|
||||
return output
|
||||
|
||||
def keyword_search(self, query, top_k=5, filters=None):
|
||||
"""
|
||||
Perform keyword-based search using Baidu Mochow's BM25 search.
|
||||
|
||||
Args:
|
||||
query (str): The text query to search for.
|
||||
top_k (int, optional): Number of results to return. Defaults to 5.
|
||||
filters (Dict, optional): Filters to apply to the search.
|
||||
|
||||
Returns:
|
||||
list: Search results, or None if the table lacks an inverted index.
|
||||
"""
|
||||
try:
|
||||
search_filter = None
|
||||
if filters:
|
||||
search_filter = self._create_filter(filters)
|
||||
|
||||
request = BM25SearchRequest(
|
||||
index_name="data_bm25_idx",
|
||||
search_text=query,
|
||||
limit=top_k,
|
||||
filter=search_filter,
|
||||
)
|
||||
|
||||
projections = ["id", "metadata"]
|
||||
res = self._table.bm25_search(request=request, projections=projections)
|
||||
|
||||
output = []
|
||||
for row in res.rows:
|
||||
row_data = row.get("row", {})
|
||||
output_data = OutputData(
|
||||
id=row_data.get("id"),
|
||||
score=row.get("score", 0.0),
|
||||
payload=row_data.get("metadata", {}),
|
||||
)
|
||||
output.append(output_data)
|
||||
|
||||
return output
|
||||
except Exception as e:
|
||||
logger.error(f"Error during keyword search for query '{query}': {e}")
|
||||
return None
|
||||
|
||||
def delete(self, vector_id):
|
||||
"""
|
||||
Delete a vector by ID.
|
||||
|
||||
@@ -56,3 +56,37 @@ class VectorStoreBase(ABC):
|
||||
def reset(self):
|
||||
"""Reset by delete the collection and recreate it."""
|
||||
pass
|
||||
|
||||
def keyword_search(self, query: str, top_k: int = 5, filters: dict = None):
|
||||
"""Keyword/BM25 full-text search. Returns None if not supported by this store.
|
||||
|
||||
Override in subclasses that support native keyword/BM25 search.
|
||||
Returns results in the same format as search() -- list of objects with
|
||||
id, score, and payload attributes.
|
||||
|
||||
Args:
|
||||
query: The search query text (should be lemmatized for best results).
|
||||
top_k: Maximum number of results to return.
|
||||
filters: Optional metadata filters (same format as search filters).
|
||||
|
||||
Returns:
|
||||
List of search results with id, score, payload, or None if not supported.
|
||||
"""
|
||||
return None
|
||||
|
||||
def search_batch(self, queries: list, vectors_list: list, top_k: int = 1, filters: dict = None):
|
||||
"""Batch search for multiple queries at once.
|
||||
|
||||
Default implementation calls search() sequentially. Override in subclasses
|
||||
with native batch support (e.g., Qdrant query_batch_points).
|
||||
|
||||
Args:
|
||||
queries: List of query texts.
|
||||
vectors_list: List of query vectors (one per query).
|
||||
top_k: Maximum results per query.
|
||||
filters: Optional metadata filters applied to all queries.
|
||||
|
||||
Returns:
|
||||
List of result lists, one per query.
|
||||
"""
|
||||
return [self.search(q, v, top_k=top_k, filters=filters) for q, v in zip(queries, vectors_list)]
|
||||
|
||||
@@ -515,6 +515,55 @@ class Databricks(VectorStoreBase):
|
||||
logger.error(f"Search failed: {e}")
|
||||
raise
|
||||
|
||||
def keyword_search(self, query, top_k=5, filters=None):
|
||||
"""
|
||||
Search for memories using full-text keyword search.
|
||||
|
||||
Only supported for DELTA_SYNC index type. Returns None for DIRECT_ACCESS indexes.
|
||||
|
||||
Args:
|
||||
query (str): Search query text.
|
||||
top_k (int): Maximum number of results. Defaults to 5.
|
||||
filters (dict, optional): Filters to apply.
|
||||
|
||||
Returns:
|
||||
List[MemoryResult] or None: Search results, or None if index type is DIRECT_ACCESS.
|
||||
"""
|
||||
if self.index_type == VectorIndexType.DIRECT_ACCESS:
|
||||
logger.warning("keyword_search is not supported for DIRECT_ACCESS index type.")
|
||||
return None
|
||||
|
||||
try:
|
||||
filters_json = json.dumps(filters) if filters else None
|
||||
|
||||
sdk_results = self.client.vector_search_indexes.query_index(
|
||||
index_name=self.fully_qualified_index_name,
|
||||
columns=self.column_names,
|
||||
query_text=query,
|
||||
num_results=top_k,
|
||||
query_type="FULL_TEXT",
|
||||
filters_json=filters_json,
|
||||
)
|
||||
|
||||
result_data = sdk_results.result if hasattr(sdk_results, "result") else sdk_results
|
||||
data_array = result_data.data_array if getattr(result_data, "data_array", None) else []
|
||||
|
||||
memory_results = []
|
||||
for row in data_array:
|
||||
row_dict = dict(zip(self.column_names, row)) if isinstance(row, (list, tuple)) else row
|
||||
score = row_dict.get("score") or (
|
||||
row[-1] if isinstance(row, (list, tuple)) and len(row) > len(self.column_names) else None
|
||||
)
|
||||
payload = {k: row_dict.get(k) for k in self.column_names}
|
||||
payload["data"] = payload.get("memory", "")
|
||||
memory_id = row_dict.get("memory_id") or row_dict.get("id")
|
||||
memory_results.append(MemoryResult(id=memory_id, score=score, payload=payload))
|
||||
return memory_results
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Keyword search failed: {e}")
|
||||
raise
|
||||
|
||||
def delete(self, vector_id):
|
||||
"""
|
||||
Delete a vector by ID from the Delta table.
|
||||
|
||||
@@ -158,6 +158,49 @@ class ElasticsearchDB(VectorStoreBase):
|
||||
|
||||
return results
|
||||
|
||||
def keyword_search(self, query, top_k=5, filters=None):
|
||||
"""Search for memories using BM25 keyword matching.
|
||||
|
||||
Args:
|
||||
query (str): The text query to search for.
|
||||
top_k (int): Maximum number of results to return. Defaults to 5.
|
||||
filters (Dict, optional): Filters to apply to the search.
|
||||
|
||||
Returns:
|
||||
List[OutputData]: Search results with id, score, and payload.
|
||||
"""
|
||||
# Build a multi_match query across text fields in metadata
|
||||
should_clauses = [
|
||||
{"match": {"metadata.data": query}},
|
||||
{"match": {"metadata.text_lemmatized": query}},
|
||||
]
|
||||
|
||||
bool_query = {
|
||||
"should": should_clauses,
|
||||
"minimum_should_match": 1,
|
||||
}
|
||||
|
||||
if filters:
|
||||
filter_conditions = []
|
||||
for key, value in filters.items():
|
||||
filter_conditions.append({"term": {f"metadata.{key}": value}})
|
||||
bool_query["filter"] = filter_conditions
|
||||
|
||||
search_query = {
|
||||
"size": top_k,
|
||||
"query": {"bool": bool_query},
|
||||
}
|
||||
|
||||
response = self.client.search(index=self.collection_name, body=search_query)
|
||||
|
||||
results = []
|
||||
for hit in response["hits"]["hits"]:
|
||||
results.append(
|
||||
OutputData(id=hit["_id"], score=hit["_score"], payload=hit.get("_source", {}).get("metadata", {}))
|
||||
)
|
||||
|
||||
return results
|
||||
|
||||
def delete(self, vector_id: str) -> None:
|
||||
"""Delete a vector by ID."""
|
||||
self.client.delete(index=self.collection_name, id=vector_id)
|
||||
|
||||
@@ -11,7 +11,7 @@ try:
|
||||
except ImportError:
|
||||
raise ImportError("The 'pymilvus' library is required. Please install it using 'pip install pymilvus'.")
|
||||
|
||||
from pymilvus import CollectionSchema, DataType, FieldSchema, MilvusClient
|
||||
from pymilvus import CollectionSchema, DataType, FieldSchema, Function, FunctionType, MilvusClient
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -73,14 +73,34 @@ class MilvusDB(VectorStoreBase):
|
||||
FieldSchema(name="id", dtype=DataType.VARCHAR, is_primary=True, max_length=512),
|
||||
FieldSchema(name="vectors", dtype=DataType.FLOAT_VECTOR, dim=vector_size),
|
||||
FieldSchema(name="metadata", dtype=DataType.JSON),
|
||||
# Text field for BM25 full-text search (auto-tokenized by Milvus analyzer)
|
||||
FieldSchema(name="text", dtype=DataType.VARCHAR, max_length=65535, enable_analyzer=True),
|
||||
# Sparse vector field populated automatically by the BM25 function below
|
||||
FieldSchema(name="sparse", dtype=DataType.SPARSE_FLOAT_VECTOR),
|
||||
]
|
||||
|
||||
schema = CollectionSchema(fields, enable_dynamic_field=True)
|
||||
|
||||
index = self.client.prepare_index_params(
|
||||
# Add BM25 function so Milvus auto-generates sparse vectors from the text field
|
||||
bm25_function = Function(
|
||||
name="bm25",
|
||||
input_field_names=["text"],
|
||||
output_field_names=["sparse"],
|
||||
function_type=FunctionType.BM25,
|
||||
)
|
||||
schema.add_function(bm25_function)
|
||||
|
||||
index_params = self.client.prepare_index_params()
|
||||
index_params.add_index(
|
||||
field_name="vectors", metric_type=metric_type, index_type="AUTOINDEX", index_name="vector_index"
|
||||
)
|
||||
self.client.create_collection(collection_name=collection_name, schema=schema, index_params=index)
|
||||
index_params.add_index(
|
||||
field_name="sparse",
|
||||
index_type="SPARSE_INVERTED_INDEX",
|
||||
metric_type="BM25",
|
||||
index_name="sparse_index",
|
||||
)
|
||||
self.client.create_collection(collection_name=collection_name, schema=schema, index_params=index_params)
|
||||
|
||||
def insert(self, ids, vectors, payloads, **kwargs: Optional[dict[str, any]]):
|
||||
"""Insert vectors into a collection.
|
||||
@@ -92,7 +112,13 @@ class MilvusDB(VectorStoreBase):
|
||||
"""
|
||||
# Batch insert all records at once for better performance and consistency
|
||||
data = [
|
||||
{"id": idx, "vectors": embedding, "metadata": metadata}
|
||||
{
|
||||
"id": idx,
|
||||
"vectors": embedding,
|
||||
"metadata": metadata,
|
||||
# Populate the text field for BM25 sparse search; prefer lemmatized text, fall back to raw data
|
||||
"text": (metadata.get("text_lemmatized") or metadata.get("data", ""))[:65535] if metadata else "",
|
||||
}
|
||||
for idx, embedding, metadata in zip(ids, vectors, payloads)
|
||||
]
|
||||
self.client.insert(collection_name=self.collection_name, data=data, **kwargs)
|
||||
@@ -163,6 +189,39 @@ class MilvusDB(VectorStoreBase):
|
||||
result = self._parse_output(data=hits[0])
|
||||
return result
|
||||
|
||||
def keyword_search(self, query, top_k=5, filters=None):
|
||||
"""
|
||||
Search for memories using BM25-based full-text search via Milvus sparse vector support.
|
||||
|
||||
Milvus 2.5+ supports native BM25 via full-text search with a SPARSE_FLOAT_VECTOR field.
|
||||
This method attempts to use that capability. If the collection does not have a sparse
|
||||
field configured, it returns None gracefully.
|
||||
|
||||
Args:
|
||||
query (str): The text query for keyword-based search.
|
||||
top_k (int, optional): Number of results to return. Defaults to 5.
|
||||
filters (dict, optional): Filters to apply to the search. Defaults to None.
|
||||
|
||||
Returns:
|
||||
list: Search results in the same format as search(), or None if sparse search
|
||||
is not supported on this collection.
|
||||
"""
|
||||
try:
|
||||
query_filter = self._create_filter(filters) if filters else None
|
||||
hits = self.client.search(
|
||||
collection_name=self.collection_name,
|
||||
data=[query],
|
||||
anns_field="sparse",
|
||||
limit=top_k,
|
||||
filter=query_filter,
|
||||
output_fields=["*"],
|
||||
)
|
||||
result = self._parse_output(data=hits[0])
|
||||
return result
|
||||
except Exception as e:
|
||||
logger.debug(f"Keyword search not available for collection {self.collection_name}: {e}")
|
||||
return None
|
||||
|
||||
def delete(self, vector_id):
|
||||
"""
|
||||
Delete a vector by ID.
|
||||
@@ -192,7 +251,10 @@ class MilvusDB(VectorStoreBase):
|
||||
if payload is None:
|
||||
payload = existing[0].get("metadata")
|
||||
|
||||
schema = {"id": vector_id, "vectors": vector, "metadata": payload}
|
||||
text = ""
|
||||
if payload:
|
||||
text = (payload.get("text_lemmatized") or payload.get("data", ""))[:65535]
|
||||
schema = {"id": vector_id, "vectors": vector, "metadata": payload, "text": text}
|
||||
self.client.upsert(collection_name=self.collection_name, data=schema)
|
||||
|
||||
def get(self, vector_id):
|
||||
|
||||
@@ -85,6 +85,41 @@ class MongoDB(VectorStoreBase):
|
||||
logger.info(
|
||||
f"Search index '{self.index_name}' created successfully for collection '{self.collection_name}'."
|
||||
)
|
||||
|
||||
# Create Atlas Search text index for keyword_search()
|
||||
text_index_name = f"{self.collection_name}_text_search_index"
|
||||
try:
|
||||
found_text_indexes = list(collection.list_search_indexes(name=text_index_name))
|
||||
if not found_text_indexes:
|
||||
text_search_index_model = SearchIndexModel(
|
||||
name=text_index_name,
|
||||
definition={
|
||||
"mappings": {
|
||||
"dynamic": False,
|
||||
"fields": {
|
||||
"payload": {
|
||||
"type": "document",
|
||||
"fields": {
|
||||
"data": {"type": "string"},
|
||||
"text_lemmatized": {"type": "string"},
|
||||
},
|
||||
}
|
||||
},
|
||||
}
|
||||
},
|
||||
)
|
||||
collection.create_search_index(text_search_index_model)
|
||||
logger.info(
|
||||
f"Text search index '{text_index_name}' created successfully for collection '{self.collection_name}'."
|
||||
)
|
||||
else:
|
||||
logger.info(f"Text search index '{text_index_name}' already exists in collection '{self.collection_name}'.")
|
||||
except Exception as e:
|
||||
logger.warning(
|
||||
f"Could not create text search index '{text_index_name}': {e}. "
|
||||
"Atlas Search may not be available. keyword_search() will not work."
|
||||
)
|
||||
|
||||
return collection
|
||||
except PyMongoError as e:
|
||||
logger.error(f"Error creating collection and search index: {e}")
|
||||
@@ -168,6 +203,58 @@ class MongoDB(VectorStoreBase):
|
||||
output = [OutputData(id=str(doc["_id"]), score=doc.get("score"), payload=doc.get("payload")) for doc in results]
|
||||
return output
|
||||
|
||||
def keyword_search(self, query, top_k=5, filters=None):
|
||||
"""
|
||||
Perform keyword-based search using MongoDB Atlas Search.
|
||||
|
||||
Args:
|
||||
query (str): The text query to search for.
|
||||
top_k (int, optional): Number of results to return. Defaults to 5.
|
||||
filters (Dict, optional): Filters to apply to the search.
|
||||
|
||||
Returns:
|
||||
List[OutputData]: Search results, or None if Atlas Search index is not available.
|
||||
"""
|
||||
try:
|
||||
collection = self.client[self.db_name][self.collection_name]
|
||||
search_index_name = f"{self.collection_name}_text_search_index"
|
||||
|
||||
pipeline = [
|
||||
{
|
||||
"$search": {
|
||||
"index": search_index_name,
|
||||
"text": {
|
||||
"query": query,
|
||||
"path": ["payload.data", "payload.text_lemmatized"],
|
||||
},
|
||||
}
|
||||
},
|
||||
{"$set": {"score": {"$meta": "searchScore"}}},
|
||||
{"$project": {"embedding": 0}},
|
||||
]
|
||||
|
||||
# Add filter stage if filters are provided
|
||||
if filters:
|
||||
filter_conditions = []
|
||||
for key, value in filters.items():
|
||||
filter_conditions.append({"payload." + key: value})
|
||||
if filter_conditions:
|
||||
pipeline.insert(1, {"$match": {"$and": filter_conditions}})
|
||||
|
||||
pipeline.append({"$limit": top_k})
|
||||
|
||||
results = list(collection.aggregate(pipeline))
|
||||
logger.info(f"Keyword search completed. Found {len(results)} documents.")
|
||||
|
||||
output = [
|
||||
OutputData(id=str(doc["_id"]), score=doc.get("score"), payload=doc.get("payload"))
|
||||
for doc in results
|
||||
]
|
||||
return output
|
||||
except Exception as e:
|
||||
logger.error(f"Error during keyword search for query '{query}': {e}")
|
||||
return None
|
||||
|
||||
def delete(self, vector_id: str) -> None:
|
||||
"""
|
||||
Delete a vector by ID.
|
||||
|
||||
@@ -203,6 +203,57 @@ class OpenSearchDB(VectorStoreBase):
|
||||
logger.error(f"Error during search: {e}", exc_info=True)
|
||||
return []
|
||||
|
||||
def keyword_search(self, query, top_k=5, filters=None):
|
||||
"""Search for memories using BM25 keyword matching.
|
||||
|
||||
Args:
|
||||
query (str): The text query to search for.
|
||||
top_k (int): Maximum number of results to return. Defaults to 5.
|
||||
filters (Dict, optional): Filters to apply to the search.
|
||||
|
||||
Returns:
|
||||
List[OutputData]: Search results with id, score, and payload.
|
||||
"""
|
||||
# Build a multi_match query across text fields in payload
|
||||
should_clauses = [
|
||||
{"match": {"payload.data": query}},
|
||||
{"match": {"payload.text_lemmatized": query}},
|
||||
]
|
||||
|
||||
bool_query = {
|
||||
"should": should_clauses,
|
||||
"minimum_should_match": 1,
|
||||
}
|
||||
|
||||
# Apply filters consistently with the existing search() method
|
||||
filter_clauses = []
|
||||
if filters:
|
||||
for key in ["user_id", "run_id", "agent_id"]:
|
||||
value = filters.get(key)
|
||||
if value:
|
||||
filter_clauses.append({"term": {f"payload.{key}.keyword": value}})
|
||||
|
||||
if filter_clauses:
|
||||
bool_query["filter"] = filter_clauses
|
||||
|
||||
query_body = {
|
||||
"size": top_k,
|
||||
"query": {"bool": bool_query},
|
||||
}
|
||||
|
||||
try:
|
||||
response = self.client.search(index=self.collection_name, body=query_body)
|
||||
|
||||
hits = response["hits"]["hits"]
|
||||
results = [
|
||||
OutputData(id=hit["_source"].get("id"), score=hit["_score"], payload=hit["_source"].get("payload", {}))
|
||||
for hit in hits[:top_k]
|
||||
]
|
||||
return results
|
||||
except Exception as e:
|
||||
logger.error(f"Error during keyword search: {e}")
|
||||
return []
|
||||
|
||||
def delete(self, vector_id: str) -> None:
|
||||
"""Delete a vector by custom ID."""
|
||||
# First, find the document by custom ID
|
||||
|
||||
@@ -179,6 +179,13 @@ class PGVector(VectorStoreBase):
|
||||
USING hnsw (vector vector_cosine_ops)
|
||||
"""
|
||||
)
|
||||
cur.execute(
|
||||
f"""
|
||||
CREATE INDEX IF NOT EXISTS {self.collection_name}_text_lemmatized_idx
|
||||
ON {self.collection_name}
|
||||
USING gin(to_tsvector('simple', payload->>'text_lemmatized'));
|
||||
"""
|
||||
)
|
||||
|
||||
def insert(self, vectors: list[list[float]], payloads=None, ids=None) -> None:
|
||||
logger.info(f"Inserting {len(vectors)} vectors into collection {self.collection_name}")
|
||||
@@ -243,6 +250,50 @@ class PGVector(VectorStoreBase):
|
||||
results = cur.fetchall()
|
||||
return [OutputData(id=str(r[0]), score=float(r[1]), payload=r[2]) for r in results]
|
||||
|
||||
def keyword_search(self, query, top_k=5, filters=None):
|
||||
"""
|
||||
Search using PostgreSQL full-text search on lemmatized text.
|
||||
|
||||
Args:
|
||||
query (str): The search query text.
|
||||
top_k (int, optional): Number of results to return. Defaults to 5.
|
||||
filters (dict, optional): Filters to apply to the search. Defaults to None.
|
||||
|
||||
Returns:
|
||||
List[OutputData]: Search results ranked by text relevance.
|
||||
"""
|
||||
filter_conditions = []
|
||||
filter_params = []
|
||||
|
||||
if filters:
|
||||
for k, v in filters.items():
|
||||
filter_conditions.append("payload->>%s = %s")
|
||||
filter_params.extend([k, str(v)])
|
||||
|
||||
filter_clause = ""
|
||||
if filter_conditions:
|
||||
filter_clause = "AND " + " AND ".join(filter_conditions)
|
||||
|
||||
try:
|
||||
with self._get_cursor() as cur:
|
||||
cur.execute(
|
||||
f"""
|
||||
SELECT id, ts_rank_cd(to_tsvector('simple', payload->>'text_lemmatized'), plainto_tsquery('simple', %s)) AS score, payload
|
||||
FROM {self.collection_name}
|
||||
WHERE to_tsvector('simple', payload->>'text_lemmatized') @@ plainto_tsquery('simple', %s)
|
||||
{filter_clause}
|
||||
ORDER BY score DESC
|
||||
LIMIT %s
|
||||
""",
|
||||
(query, query, *filter_params, top_k),
|
||||
)
|
||||
|
||||
results = cur.fetchall()
|
||||
return [OutputData(id=str(r[0]), score=float(r[1]), payload=r[2]) for r in results]
|
||||
except Exception as e:
|
||||
logger.debug(f"Keyword search failed: {e}")
|
||||
return None
|
||||
|
||||
def delete(self, vector_id: str) -> None:
|
||||
"""
|
||||
Delete a vector by ID.
|
||||
|
||||
@@ -241,6 +241,42 @@ class PineconeDB(VectorStoreBase):
|
||||
results = self._parse_output(response.matches)
|
||||
return results
|
||||
|
||||
def keyword_search(self, query, top_k=5, filters=None):
|
||||
"""
|
||||
Search using BM25 sparse vectors for keyword-based retrieval.
|
||||
|
||||
Args:
|
||||
query (str): The search query text.
|
||||
top_k (int, optional): Number of results to return. Defaults to 5.
|
||||
filters (dict, optional): Filters to apply to the search. Defaults to None.
|
||||
|
||||
Returns:
|
||||
List[OutputData]: Search results, or None if hybrid search is not configured.
|
||||
"""
|
||||
if not self.hybrid_search or self.sparse_encoder is None:
|
||||
return None
|
||||
|
||||
try:
|
||||
filter_dict = self._create_filter(filters) if filters else None
|
||||
|
||||
sparse_vector = self.sparse_encoder.encode_queries(query)
|
||||
|
||||
query_params = {
|
||||
"sparse_vector": sparse_vector,
|
||||
"top_k": top_k,
|
||||
"include_metadata": True,
|
||||
"include_values": False,
|
||||
}
|
||||
|
||||
if filter_dict:
|
||||
query_params["filter"] = filter_dict
|
||||
|
||||
response = self.index.query(**query_params, namespace=self.namespace)
|
||||
|
||||
return self._parse_output(response.matches)
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
def delete(self, vector_id: Union[str, int]):
|
||||
"""
|
||||
Delete a vector by ID.
|
||||
|
||||
+121
-15
@@ -2,7 +2,7 @@ import logging
|
||||
import re
|
||||
from typing import Optional
|
||||
|
||||
from qdrant_client import QdrantClient
|
||||
from qdrant_client import QdrantClient, models
|
||||
from qdrant_client.models import (
|
||||
DatetimeRange,
|
||||
Distance,
|
||||
@@ -16,6 +16,8 @@ from qdrant_client.models import (
|
||||
PointStruct,
|
||||
PointVectors,
|
||||
Range,
|
||||
SparseVector,
|
||||
SparseVectorParams,
|
||||
VectorParams,
|
||||
)
|
||||
|
||||
@@ -64,7 +66,7 @@ class Qdrant(VectorStoreBase):
|
||||
if host and port:
|
||||
params["host"] = host
|
||||
params["port"] = port
|
||||
|
||||
|
||||
if not params:
|
||||
params["path"] = path
|
||||
self.is_local = True
|
||||
@@ -76,11 +78,44 @@ class Qdrant(VectorStoreBase):
|
||||
self.collection_name = collection_name
|
||||
self.embedding_model_dims = embedding_model_dims
|
||||
self.on_disk = on_disk
|
||||
self._bm25_encoder = None
|
||||
self.create_col(embedding_model_dims, on_disk)
|
||||
|
||||
def _get_bm25_encoder(self):
|
||||
"""Lazy-load the BM25 sparse text encoder (fastembed)."""
|
||||
if self._bm25_encoder is None:
|
||||
try:
|
||||
from fastembed import SparseTextEmbedding
|
||||
self._bm25_encoder = SparseTextEmbedding(model_name="Qdrant/bm25")
|
||||
logger.info("BM25 encoder loaded (fastembed Qdrant/bm25)")
|
||||
except ImportError:
|
||||
logger.warning("fastembed not installed — BM25 keyword search disabled. Install with: pip install fastembed")
|
||||
self._bm25_encoder = False # sentinel: tried and failed
|
||||
except Exception as e:
|
||||
logger.warning(f"Failed to load BM25 encoder: {e}")
|
||||
self._bm25_encoder = False
|
||||
return self._bm25_encoder if self._bm25_encoder is not False else None
|
||||
|
||||
def _encode_bm25(self, text: str) -> SparseVector | None:
|
||||
"""Encode text into a BM25 sparse vector."""
|
||||
encoder = self._get_bm25_encoder()
|
||||
if encoder is None:
|
||||
return None
|
||||
try:
|
||||
results = list(encoder.embed([text]))
|
||||
if results:
|
||||
sparse = results[0]
|
||||
return SparseVector(
|
||||
indices=sparse.indices.tolist(),
|
||||
values=sparse.values.tolist(),
|
||||
)
|
||||
except Exception as e:
|
||||
logger.debug(f"BM25 encoding failed: {e}")
|
||||
return None
|
||||
|
||||
def create_col(self, vector_size: int, on_disk: bool, distance: Distance = Distance.COSINE):
|
||||
"""
|
||||
Create a new collection.
|
||||
Create a new collection with dense vectors and BM25 sparse vectors.
|
||||
|
||||
Args:
|
||||
vector_size (int): Size of the vectors to be stored.
|
||||
@@ -98,6 +133,11 @@ class Qdrant(VectorStoreBase):
|
||||
self.client.create_collection(
|
||||
collection_name=self.collection_name,
|
||||
vectors_config=VectorParams(size=vector_size, distance=distance, on_disk=on_disk),
|
||||
sparse_vectors_config={
|
||||
"bm25": SparseVectorParams(
|
||||
modifier=models.Modifier.IDF,
|
||||
),
|
||||
},
|
||||
)
|
||||
self._create_filter_indexes()
|
||||
|
||||
@@ -107,9 +147,9 @@ class Qdrant(VectorStoreBase):
|
||||
if self.is_local:
|
||||
logger.debug("Skipping payload index creation for local Qdrant (not supported)")
|
||||
return
|
||||
|
||||
|
||||
common_fields = ["user_id", "agent_id", "run_id", "actor_id"]
|
||||
|
||||
|
||||
for field in common_fields:
|
||||
try:
|
||||
self.client.create_payload_index(
|
||||
@@ -123,7 +163,8 @@ class Qdrant(VectorStoreBase):
|
||||
|
||||
def insert(self, vectors: list, payloads: list = None, ids: list = None):
|
||||
"""
|
||||
Insert vectors into a collection.
|
||||
Insert vectors into a collection, including BM25 sparse vectors
|
||||
computed from the text_lemmatized payload field.
|
||||
|
||||
Args:
|
||||
vectors (list): List of vectors to insert.
|
||||
@@ -131,14 +172,21 @@ class Qdrant(VectorStoreBase):
|
||||
ids (list, optional): List of IDs corresponding to vectors. Defaults to None.
|
||||
"""
|
||||
logger.info(f"Inserting {len(vectors)} vectors into collection {self.collection_name}")
|
||||
points = [
|
||||
PointStruct(
|
||||
id=idx if ids is None else ids[idx],
|
||||
vector=vector,
|
||||
payload=payloads[idx] if payloads else {},
|
||||
)
|
||||
for idx, vector in enumerate(vectors)
|
||||
]
|
||||
points = []
|
||||
for idx, vector in enumerate(vectors):
|
||||
payload = payloads[idx] if payloads else {}
|
||||
point_id = idx if ids is None else ids[idx]
|
||||
|
||||
# Build named vectors: dense + optional BM25 sparse
|
||||
named_vectors = {"": vector}
|
||||
text_for_bm25 = payload.get("text_lemmatized") or payload.get("data", "")
|
||||
if text_for_bm25:
|
||||
sparse = self._encode_bm25(text_for_bm25)
|
||||
if sparse is not None:
|
||||
named_vectors["bm25"] = sparse
|
||||
|
||||
points.append(PointStruct(id=point_id, vector=named_vectors, payload=payload))
|
||||
|
||||
self.client.upsert(collection_name=self.collection_name, points=points)
|
||||
|
||||
# ISO 8601 datetime pattern for detecting datetime strings in range filters
|
||||
@@ -330,6 +378,53 @@ class Qdrant(VectorStoreBase):
|
||||
)
|
||||
return hits.points
|
||||
|
||||
def search_batch(self, queries: list, vectors_list: list, top_k: int = 1, filters: dict = None):
|
||||
"""Batch search using Qdrant's query_batch_points for efficiency."""
|
||||
query_filter = self._create_filter(filters) if filters else None
|
||||
requests = [
|
||||
models.QueryRequest(query=vec, filter=query_filter, limit=top_k)
|
||||
for vec in vectors_list
|
||||
]
|
||||
try:
|
||||
results = self.client.query_batch_points(
|
||||
collection_name=self.collection_name,
|
||||
requests=requests,
|
||||
)
|
||||
return [r.points for r in results]
|
||||
except Exception as e:
|
||||
logger.warning(f"Batch search failed, falling back to sequential: {e}")
|
||||
return [self.search(q, v, top_k=top_k, filters=filters) for q, v in zip(queries, vectors_list)]
|
||||
|
||||
def keyword_search(self, query, top_k=5, filters=None):
|
||||
"""
|
||||
Search using BM25 sparse vectors for keyword-based retrieval.
|
||||
|
||||
Args:
|
||||
query (str): The search query text.
|
||||
top_k (int, optional): Number of results to return. Defaults to 5.
|
||||
filters (dict, optional): Filters to apply to the search. Defaults to None.
|
||||
|
||||
Returns:
|
||||
list: Search results, or None if BM25 is not available.
|
||||
"""
|
||||
sparse_query = self._encode_bm25(query)
|
||||
if sparse_query is None:
|
||||
return None
|
||||
|
||||
try:
|
||||
query_filter = self._create_filter(filters) if filters else None
|
||||
hits = self.client.query_points(
|
||||
collection_name=self.collection_name,
|
||||
query=sparse_query,
|
||||
using="bm25",
|
||||
query_filter=query_filter,
|
||||
limit=top_k,
|
||||
)
|
||||
return hits.points
|
||||
except Exception as e:
|
||||
logger.debug(f"BM25 keyword search failed: {e}")
|
||||
return None
|
||||
|
||||
def delete(self, vector_id: int):
|
||||
"""
|
||||
Delete a vector by ID.
|
||||
@@ -354,9 +449,20 @@ class Qdrant(VectorStoreBase):
|
||||
payload (dict, optional): Updated payload. Defaults to None.
|
||||
"""
|
||||
if vector is not None and payload is not None:
|
||||
point = PointStruct(id=vector_id, vector=vector, payload=payload)
|
||||
# Full update: attach BM25 sparse vector alongside dense vector
|
||||
named_vectors = {"": vector}
|
||||
text_for_bm25 = payload.get("text_lemmatized") or payload.get("data", "")
|
||||
if text_for_bm25:
|
||||
sparse = self._encode_bm25(text_for_bm25)
|
||||
if sparse is not None:
|
||||
named_vectors["bm25"] = sparse
|
||||
point = PointStruct(id=vector_id, vector=named_vectors, payload=payload)
|
||||
self.client.upsert(collection_name=self.collection_name, points=[point])
|
||||
else:
|
||||
# Partial update: use Qdrant's dedicated endpoints.
|
||||
# Note: BM25 sparse vector cannot be refreshed via set_payload alone;
|
||||
# payload-only updates will leave any existing BM25 vector stale. In
|
||||
# practice v3 re-embeds on memory text change, so this is acceptable.
|
||||
if payload is not None:
|
||||
self.client.set_payload(
|
||||
collection_name=self.collection_name,
|
||||
|
||||
@@ -7,7 +7,7 @@ import numpy as np
|
||||
import redis
|
||||
from redis.commands.search.query import Query
|
||||
from redisvl.index import SearchIndex
|
||||
from redisvl.query import VectorQuery
|
||||
from redisvl.query import TextQuery, VectorQuery
|
||||
from redisvl.query.filter import Tag
|
||||
|
||||
from mem0.memory.utils import extract_json
|
||||
@@ -181,6 +181,60 @@ class RedisDB(VectorStoreBase):
|
||||
for result in results
|
||||
]
|
||||
|
||||
def keyword_search(self, query, top_k=5, filters=None):
|
||||
"""
|
||||
Search for memories using BM25 keyword search on the memory field.
|
||||
|
||||
Args:
|
||||
query (str): Search query text.
|
||||
top_k (int): Maximum number of results. Defaults to 5.
|
||||
filters (dict, optional): Filters to apply (user_id, agent_id, run_id).
|
||||
|
||||
Returns:
|
||||
List[MemoryResult]: Search results.
|
||||
"""
|
||||
filter_expression = None
|
||||
if filters:
|
||||
conditions = [Tag(key) == value for key, value in filters.items() if value is not None]
|
||||
if conditions:
|
||||
filter_expression = reduce(lambda x, y: x & y, conditions)
|
||||
|
||||
t = TextQuery(
|
||||
text=query,
|
||||
text_field_name="memory",
|
||||
return_fields=["memory_id", "hash", "agent_id", "run_id", "user_id", "memory", "metadata", "created_at"],
|
||||
filter_expression=filter_expression,
|
||||
num_results=top_k,
|
||||
)
|
||||
|
||||
results = self.index.query(t)
|
||||
|
||||
return [
|
||||
MemoryResult(
|
||||
id=result["memory_id"],
|
||||
score=result.get("text_score", 1.0),
|
||||
payload={
|
||||
"hash": result["hash"],
|
||||
"data": result["memory"],
|
||||
"created_at": datetime.fromtimestamp(
|
||||
int(result["created_at"]), tz=timezone.utc
|
||||
).isoformat(timespec="microseconds"),
|
||||
**(
|
||||
{
|
||||
"updated_at": datetime.fromtimestamp(
|
||||
int(result["updated_at"]), tz=timezone.utc
|
||||
).isoformat(timespec="microseconds")
|
||||
}
|
||||
if "updated_at" in result
|
||||
else {}
|
||||
),
|
||||
**{field: result[field] for field in ["agent_id", "run_id", "user_id"] if field in result},
|
||||
**{k: v for k, v in json.loads(extract_json(result["metadata"])).items()},
|
||||
},
|
||||
)
|
||||
for result in results
|
||||
]
|
||||
|
||||
def delete(self, vector_id):
|
||||
self.index.drop_keys(f"{self.schema['index']['prefix']}:{vector_id}")
|
||||
|
||||
|
||||
@@ -149,6 +149,45 @@ class UpstashVector(VectorStoreBase):
|
||||
for res in response
|
||||
]
|
||||
|
||||
def keyword_search(self, query, top_k=5, filters=None):
|
||||
"""
|
||||
Perform keyword-based search using Upstash's BM25 sparse search.
|
||||
|
||||
Args:
|
||||
query (str): The text query to search for.
|
||||
top_k (int, optional): Number of results to return. Defaults to 5.
|
||||
filters (Dict, optional): Filters to apply to the search.
|
||||
|
||||
Returns:
|
||||
List[OutputData]: Search results, or None if sparse/BM25 search is not supported.
|
||||
"""
|
||||
try:
|
||||
filters_str = (
|
||||
" AND ".join([f"{k} = {self._stringify(v)}" for k, v in filters.items()])
|
||||
if filters
|
||||
else None
|
||||
)
|
||||
|
||||
response = self.client.query(
|
||||
data=query,
|
||||
top_k=top_k,
|
||||
filter=filters_str or "",
|
||||
include_metadata=True,
|
||||
namespace=self.collection_name,
|
||||
)
|
||||
|
||||
return [
|
||||
OutputData(
|
||||
id=res.id,
|
||||
score=res.score,
|
||||
payload=res.metadata,
|
||||
)
|
||||
for res in response
|
||||
]
|
||||
except Exception as e:
|
||||
logger.error(f"Error during keyword search for query '{query}': {e}")
|
||||
return None
|
||||
|
||||
def delete(self, vector_id: int):
|
||||
"""
|
||||
Delete a vector by ID.
|
||||
|
||||
@@ -276,6 +276,11 @@ class GoogleMatchingEngine(VectorStoreBase):
|
||||
logger.error("Stack trace: %s", traceback.format_exc())
|
||||
raise
|
||||
|
||||
def keyword_search(self, query, top_k=5, filters=None):
|
||||
# Vertex AI hybrid search requires sparse embeddings configuration.
|
||||
# Not yet supported - requires HybridQuery with sparse encoder setup.
|
||||
return None
|
||||
|
||||
def delete(self, vector_id: Optional[str] = None, ids: Optional[List[str]] = None) -> bool:
|
||||
"""
|
||||
Delete vectors from the Matching Engine index.
|
||||
|
||||
@@ -223,6 +223,55 @@ class Weaviate(VectorStoreBase):
|
||||
)
|
||||
return results
|
||||
|
||||
def keyword_search(self, query, top_k=5, filters=None):
|
||||
"""
|
||||
Search for memories using BM25 keyword search.
|
||||
|
||||
Args:
|
||||
query (str): Search query text.
|
||||
top_k (int): Maximum number of results. Defaults to 5.
|
||||
filters (dict, optional): Filters to apply (user_id, agent_id, run_id).
|
||||
|
||||
Returns:
|
||||
List[OutputData]: Search results.
|
||||
"""
|
||||
collection = self.client.collections.get(str(self.collection_name))
|
||||
filter_conditions = []
|
||||
if filters:
|
||||
for key, value in filters.items():
|
||||
if value and key in ["user_id", "agent_id", "run_id"]:
|
||||
filter_conditions.append(Filter.by_property(key).equal(value))
|
||||
combined_filter = Filter.all_of(filter_conditions) if filter_conditions else None
|
||||
response = collection.query.bm25(
|
||||
query=query,
|
||||
query_properties=["data"],
|
||||
limit=top_k,
|
||||
filters=combined_filter,
|
||||
return_properties=["hash", "created_at", "updated_at", "user_id", "agent_id", "run_id", "data", "category"],
|
||||
return_metadata=MetadataQuery(score=True),
|
||||
)
|
||||
results = []
|
||||
for obj in response.objects:
|
||||
payload = obj.properties.copy()
|
||||
|
||||
for id_field in ["run_id", "agent_id", "user_id"]:
|
||||
if id_field in payload and payload[id_field] is None:
|
||||
del payload[id_field]
|
||||
|
||||
payload["id"] = str(obj.uuid).split("'")[0]
|
||||
if obj.metadata.score is not None:
|
||||
score = obj.metadata.score
|
||||
else:
|
||||
score = 1.0
|
||||
results.append(
|
||||
OutputData(
|
||||
id=str(obj.uuid),
|
||||
score=score,
|
||||
payload=payload,
|
||||
)
|
||||
)
|
||||
return results
|
||||
|
||||
def delete(self, vector_id):
|
||||
"""
|
||||
Delete a vector by ID.
|
||||
|
||||
+5
-14
@@ -14,7 +14,7 @@ license = "Apache-2.0"
|
||||
license-files = ["LICENSE"]
|
||||
requires-python = ">=3.9,<4.0"
|
||||
dependencies = [
|
||||
"qdrant-client>=1.9.1",
|
||||
"qdrant-client>=1.12.0",
|
||||
"pydantic>=2.7.3",
|
||||
"openai>=1.90.0",
|
||||
"posthog>=4.5.0",
|
||||
@@ -24,14 +24,8 @@ dependencies = [
|
||||
]
|
||||
|
||||
[project.optional-dependencies]
|
||||
graph = [
|
||||
"langchain-neo4j>=0.4.0",
|
||||
"langchain-aws>=0.2.23",
|
||||
"langchain-memgraph>=0.1.0",
|
||||
"neo4j>=5.23.1",
|
||||
"rank-bm25>=0.2.2",
|
||||
"kuzu>=0.11.0",
|
||||
"apache-age-python>=0.0.6",
|
||||
nlp = [
|
||||
"spacy>=3.7.0",
|
||||
]
|
||||
vector_stores = [
|
||||
"vecs>=0.4.0",
|
||||
@@ -41,7 +35,7 @@ vector_stores = [
|
||||
"pinecone<=7.3.0",
|
||||
"pinecone-text>=0.10.0",
|
||||
"faiss-cpu>=1.7.4",
|
||||
"upstash-vector>=0.1.0",
|
||||
"upstash-vector>=0.6.0",
|
||||
"azure-search-documents>=11.4.0b8",
|
||||
"psycopg>=3.2.8",
|
||||
"psycopg-pool>=3.2.6,<4.0.0",
|
||||
@@ -70,6 +64,7 @@ llms = [
|
||||
]
|
||||
extras = [
|
||||
"boto3>=1.34.0",
|
||||
"langchain>=0.1.0",
|
||||
"langchain-community>=0.0.0",
|
||||
"sentence-transformers>=5.0.0",
|
||||
"elasticsearch>=8.0.0,<9.0.0",
|
||||
@@ -107,7 +102,6 @@ only-include = ["mem0"]
|
||||
python = "3.9"
|
||||
features = [
|
||||
"test",
|
||||
"graph",
|
||||
"vector_stores",
|
||||
"llms",
|
||||
"extras",
|
||||
@@ -117,7 +111,6 @@ features = [
|
||||
python = "3.10"
|
||||
features = [
|
||||
"test",
|
||||
"graph",
|
||||
"vector_stores",
|
||||
"llms",
|
||||
"extras",
|
||||
@@ -127,7 +120,6 @@ features = [
|
||||
python = "3.11"
|
||||
features = [
|
||||
"test",
|
||||
"graph",
|
||||
"vector_stores",
|
||||
"llms",
|
||||
"extras",
|
||||
@@ -137,7 +129,6 @@ features = [
|
||||
python = "3.12"
|
||||
features = [
|
||||
"test",
|
||||
"graph",
|
||||
"vector_stores",
|
||||
"llms",
|
||||
"extras",
|
||||
|
||||
@@ -40,14 +40,6 @@ POSTGRES_USER = os.environ.get("POSTGRES_USER", "postgres")
|
||||
POSTGRES_PASSWORD = os.environ.get("POSTGRES_PASSWORD", "postgres")
|
||||
POSTGRES_COLLECTION_NAME = os.environ.get("POSTGRES_COLLECTION_NAME", "memories")
|
||||
|
||||
NEO4J_URI = os.environ.get("NEO4J_URI", "bolt://neo4j:7687")
|
||||
NEO4J_USERNAME = os.environ.get("NEO4J_USERNAME", "neo4j")
|
||||
NEO4J_PASSWORD = os.environ.get("NEO4J_PASSWORD", "mem0graph")
|
||||
|
||||
MEMGRAPH_URI = os.environ.get("MEMGRAPH_URI", "bolt://localhost:7687")
|
||||
MEMGRAPH_USERNAME = os.environ.get("MEMGRAPH_USERNAME", "memgraph")
|
||||
MEMGRAPH_PASSWORD = os.environ.get("MEMGRAPH_PASSWORD", "mem0graph")
|
||||
|
||||
OPENAI_API_KEY = os.environ.get("OPENAI_API_KEY")
|
||||
HISTORY_DB_PATH = os.environ.get("HISTORY_DB_PATH", "/app/history/history.db")
|
||||
|
||||
@@ -64,10 +56,6 @@ DEFAULT_CONFIG = {
|
||||
"collection_name": POSTGRES_COLLECTION_NAME,
|
||||
},
|
||||
},
|
||||
"graph_store": {
|
||||
"provider": "neo4j",
|
||||
"config": {"url": NEO4J_URI, "username": NEO4J_USERNAME, "password": NEO4J_PASSWORD},
|
||||
},
|
||||
"llm": {"provider": "openai", "config": {"api_key": OPENAI_API_KEY, "temperature": 0.2, "model": "gpt-4.1-nano-2025-04-14"}},
|
||||
"embedder": {"provider": "openai", "config": {"api_key": OPENAI_API_KEY, "model": "text-embedding-3-small"}},
|
||||
"history_db_path": HISTORY_DB_PATH,
|
||||
|
||||
@@ -2,6 +2,8 @@ from unittest.mock import Mock
|
||||
|
||||
import pytest
|
||||
|
||||
pytest.importorskip("langchain", reason="langchain not installed")
|
||||
|
||||
from mem0.configs.llms.base import BaseLlmConfig
|
||||
from mem0.llms.langchain import LangchainLLM
|
||||
|
||||
|
||||
+26
-20
@@ -144,6 +144,7 @@ def create_mocked_memory():
|
||||
|
||||
mock_embedder = MagicMock()
|
||||
mock_embedder.embed.return_value = [0.1, 0.2, 0.3]
|
||||
mock_embedder.embed_batch.return_value = [[0.1, 0.2, 0.3]]
|
||||
mock_embedder_factory.return_value = mock_embedder
|
||||
|
||||
mock_vector_store = MagicMock()
|
||||
@@ -151,9 +152,12 @@ def create_mocked_memory():
|
||||
mock_vector_store.add.return_value = None
|
||||
mock_vector_factory.return_value = mock_vector_store
|
||||
|
||||
mock_sqlite.return_value = MagicMock()
|
||||
mock_db = MagicMock()
|
||||
mock_db.get_last_messages.return_value = []
|
||||
mock_sqlite.return_value = mock_db
|
||||
|
||||
memory = Memory()
|
||||
memory.custom_instructions = None
|
||||
memory.api_version = "v1.0"
|
||||
return memory, mock_llm, mock_vector_store
|
||||
|
||||
@@ -170,6 +174,7 @@ def create_mocked_async_memory():
|
||||
|
||||
mock_embedder = MagicMock()
|
||||
mock_embedder.embed.return_value = [0.1, 0.2, 0.3]
|
||||
mock_embedder.embed_batch.return_value = [[0.1, 0.2, 0.3]]
|
||||
mock_embedder_factory.return_value = mock_embedder
|
||||
|
||||
mock_vector_store = MagicMock()
|
||||
@@ -177,9 +182,12 @@ def create_mocked_async_memory():
|
||||
mock_vector_store.add.return_value = None
|
||||
mock_vector_factory.return_value = mock_vector_store
|
||||
|
||||
mock_sqlite.return_value = MagicMock()
|
||||
mock_db = MagicMock()
|
||||
mock_db.get_last_messages.return_value = []
|
||||
mock_sqlite.return_value = mock_db
|
||||
|
||||
memory = AsyncMemory()
|
||||
memory.custom_instructions = None
|
||||
memory.api_version = "v1.0"
|
||||
return memory, mock_llm, mock_vector_store
|
||||
|
||||
@@ -187,22 +195,21 @@ def create_mocked_async_memory():
|
||||
def test_thinking_tags_sync():
|
||||
"""Test thinking tags handling in Memory._add_to_vector_store (sync)."""
|
||||
memory, mock_llm, mock_vector_store = create_mocked_memory()
|
||||
|
||||
# Mock LLM responses for both phases
|
||||
|
||||
# v3 pipeline: single LLM call returning ADD-only memories
|
||||
mock_llm.generate_response.side_effect = [
|
||||
' <think>Sync fact extraction</think> \n{"facts": ["User loves sci-fi"]}',
|
||||
' <think>Sync memory actions</think> \n{"memory": [{"text": "Loves sci-fi", "event": "ADD"}]}'
|
||||
'<think>Sync extraction</think>\n{"memory": [{"text": "Loves sci-fi", "attributed_to": "user"}]}'
|
||||
]
|
||||
|
||||
|
||||
mock_vector_store.search.return_value = []
|
||||
|
||||
|
||||
result = memory._add_to_vector_store(
|
||||
messages=[{"role": "user", "content": "I love sci-fi movies"}],
|
||||
metadata={},
|
||||
filters={},
|
||||
metadata={},
|
||||
filters={},
|
||||
infer=True
|
||||
)
|
||||
|
||||
|
||||
assert len(result) == 1
|
||||
assert result[0]["memory"] == "Loves sci-fi"
|
||||
assert result[0]["event"] == "ADD"
|
||||
@@ -213,13 +220,12 @@ def test_thinking_tags_sync():
|
||||
async def test_async_thinking_tags_async():
|
||||
"""Test thinking tags handling in AsyncMemory._add_to_vector_store."""
|
||||
memory, mock_llm, mock_vector_store = create_mocked_async_memory()
|
||||
|
||||
# Directly mock llm.generate_response instead of via asyncio.to_thread
|
||||
|
||||
# v3 pipeline: single LLM call returning ADD-only memories
|
||||
mock_llm.generate_response.side_effect = [
|
||||
' <think>Async fact extraction</think> \n{"facts": ["User loves sci-fi"]}',
|
||||
' <think>Async memory actions</think> \n{"memory": [{"text": "Loves sci-fi", "event": "ADD"}]}'
|
||||
'<think>Async extraction</think>\n{"memory": [{"text": "Loves sci-fi", "attributed_to": "user"}]}'
|
||||
]
|
||||
|
||||
|
||||
# Mock asyncio.to_thread to call the function directly (bypass threading)
|
||||
async def mock_to_thread(func, *args, **kwargs):
|
||||
if func == mock_llm.generate_response:
|
||||
@@ -230,15 +236,15 @@ async def test_async_thinking_tags_async():
|
||||
return []
|
||||
else:
|
||||
return func(*args, **kwargs)
|
||||
|
||||
|
||||
with patch('mem0.memory.main.asyncio.to_thread', side_effect=mock_to_thread):
|
||||
result = await memory._add_to_vector_store(
|
||||
messages=[{"role": "user", "content": "I love sci-fi movies"}],
|
||||
metadata={},
|
||||
effective_filters={},
|
||||
metadata={},
|
||||
effective_filters={},
|
||||
infer=True
|
||||
)
|
||||
|
||||
|
||||
assert len(result) == 1
|
||||
assert result[0]["memory"] == "Loves sci-fi"
|
||||
assert result[0]["event"] == "ADD"
|
||||
File diff suppressed because it is too large
Load Diff
@@ -1,226 +0,0 @@
|
||||
from unittest.mock import MagicMock, Mock, patch
|
||||
|
||||
# age and rank_bm25 are optional deps — mock them so tests run without install
|
||||
_age_mock = Mock()
|
||||
patch.dict("sys.modules", {
|
||||
"age": _age_mock,
|
||||
"age.models": Mock(),
|
||||
"rank_bm25": Mock(),
|
||||
}).start()
|
||||
|
||||
from mem0.memory.apache_age_memory import MemoryGraph, _cosine_similarity # noqa: E402
|
||||
|
||||
|
||||
def _make_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
|
||||
instance.ag = MagicMock()
|
||||
instance.graph_name = "test_graph"
|
||||
instance.threshold = 0.7
|
||||
return instance
|
||||
|
||||
|
||||
class TestCosineSimilarity:
|
||||
"""Tests for the _cosine_similarity helper."""
|
||||
|
||||
def test_identical_vectors(self):
|
||||
assert abs(_cosine_similarity([1, 0, 0], [1, 0, 0]) - 1.0) < 1e-6
|
||||
|
||||
def test_orthogonal_vectors(self):
|
||||
assert abs(_cosine_similarity([1, 0, 0], [0, 1, 0])) < 1e-6
|
||||
|
||||
def test_zero_vector(self):
|
||||
assert _cosine_similarity([0, 0, 0], [1, 2, 3]) == 0.0
|
||||
|
||||
|
||||
class TestRetrieveNodesFromData:
|
||||
"""Tests for _retrieve_nodes_from_data in Apache AGE MemoryGraph."""
|
||||
|
||||
def test_normal_entities_extracted(self):
|
||||
instance = _make_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_malformed_entity_missing_entity_type_is_skipped(self):
|
||||
instance = _make_instance()
|
||||
instance.llm.generate_response.return_value = {
|
||||
"tool_calls": [{"name": "extract_entities", "arguments": {"entities": [
|
||||
{"entity": "matrix multiplication", "entity_type": "task"},
|
||||
{"entity": "task"},
|
||||
{"entity": "ReLU", "entity_type": "task"},
|
||||
]}}]
|
||||
}
|
||||
result = instance._retrieve_nodes_from_data("some text", {"user_id": "u1"})
|
||||
assert "matrix_multiplication" in result
|
||||
assert "relu" in result
|
||||
assert "task" not in result
|
||||
|
||||
def test_missing_entities_key_returns_empty(self):
|
||||
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}
|
||||
result = instance._retrieve_nodes_from_data("hello world", {"user_id": "u1"})
|
||||
assert result == {}
|
||||
|
||||
|
||||
class TestEstablishNodesRelationsFromData:
|
||||
"""Tests for _establish_nodes_relations_from_data in Apache AGE MemoryGraph."""
|
||||
|
||||
def test_none_response_does_not_crash(self):
|
||||
instance = _make_instance()
|
||||
instance.llm.generate_response.return_value = None
|
||||
result = instance._establish_nodes_relations_from_data(
|
||||
"Hello world", {"user_id": "u1"}, {}
|
||||
)
|
||||
assert result == []
|
||||
|
||||
def test_empty_tool_calls_returns_empty(self):
|
||||
instance = _make_instance()
|
||||
instance.llm.generate_response.return_value = {"tool_calls": []}
|
||||
result = instance._establish_nodes_relations_from_data(
|
||||
"Hello world", {"user_id": "u1"}, {}
|
||||
)
|
||||
assert result == []
|
||||
|
||||
def test_valid_entities_returned(self):
|
||||
instance = _make_instance()
|
||||
instance.llm.generate_response.return_value = {
|
||||
"tool_calls": [{"name": "add_entities", "arguments": {"entities": [
|
||||
{"source": "alice", "relationship": "loves", "destination": "hiking"}
|
||||
]}}]
|
||||
}
|
||||
result = instance._establish_nodes_relations_from_data(
|
||||
"Alice loves hiking", {"user_id": "u1"}, {"alice": "person"}
|
||||
)
|
||||
assert len(result) == 1
|
||||
assert result[0]["source"] == "alice"
|
||||
|
||||
|
||||
class TestRemoveSpacesFromEntities:
|
||||
"""Tests for _remove_spaces_from_entities."""
|
||||
|
||||
def test_spaces_and_case(self):
|
||||
instance = _make_instance()
|
||||
entities = [{"source": "Alice Smith", "relationship": "Works At", "destination": "Big Corp"}]
|
||||
result = instance._remove_spaces_from_entities(entities)
|
||||
assert result[0]["source"] == "alice_smith"
|
||||
assert result[0]["relationship"] == "works_at"
|
||||
assert result[0]["destination"] == "big_corp"
|
||||
|
||||
|
||||
class TestFindSimilarNode:
|
||||
"""Tests for _find_similar_node."""
|
||||
|
||||
def test_returns_none_when_no_nodes(self):
|
||||
instance = _make_instance()
|
||||
instance._exec_cypher = MagicMock(return_value=[])
|
||||
result = instance._find_similar_node([1.0, 0.0], {"user_id": "u1"}, threshold=0.9)
|
||||
assert result is None
|
||||
|
||||
def test_returns_best_match_above_threshold(self):
|
||||
instance = _make_instance()
|
||||
instance._exec_cypher = MagicMock(return_value=[
|
||||
{"name": "alice", "embedding": [1.0, 0.0], "user_id": "u1"},
|
||||
{"name": "bob", "embedding": [0.0, 1.0], "user_id": "u1"},
|
||||
])
|
||||
result = instance._find_similar_node([1.0, 0.0], {"user_id": "u1"}, threshold=0.9)
|
||||
assert result["name"] == "alice"
|
||||
|
||||
def test_filters_by_agent_id(self):
|
||||
instance = _make_instance()
|
||||
instance._exec_cypher = MagicMock(return_value=[
|
||||
{"name": "alice", "embedding": [1.0, 0.0], "user_id": "u1", "agent_id": "a2"},
|
||||
])
|
||||
result = instance._find_similar_node(
|
||||
[1.0, 0.0], {"user_id": "u1", "agent_id": "a1"}, threshold=0.9
|
||||
)
|
||||
assert result is None
|
||||
|
||||
|
||||
class TestDeleteAll:
|
||||
"""Tests for delete_all."""
|
||||
|
||||
def test_calls_exec_cypher_and_commits(self):
|
||||
instance = _make_instance()
|
||||
instance._exec_cypher = MagicMock(return_value=[])
|
||||
instance.delete_all({"user_id": "u1"})
|
||||
instance._exec_cypher.assert_called_once()
|
||||
instance.ag.commit.assert_called_once()
|
||||
|
||||
|
||||
class TestGetAll:
|
||||
"""Tests for get_all."""
|
||||
|
||||
def test_returns_formatted_results(self):
|
||||
instance = _make_instance()
|
||||
instance._exec_cypher = MagicMock(return_value=[
|
||||
{"source": "alice", "relationship": "KNOWS", "target": "bob"},
|
||||
{"source": "alice", "relationship": "LIKES", "target": "hiking"},
|
||||
])
|
||||
results = instance.get_all({"user_id": "u1"}, top_k=10)
|
||||
assert len(results) == 2
|
||||
assert results[0]["source"] == "alice"
|
||||
assert results[0]["relationship"] == "KNOWS"
|
||||
assert results[0]["target"] == "bob"
|
||||
|
||||
def test_passes_limit_to_cypher(self):
|
||||
"""Limit is enforced via LIMIT in the Cypher query, not Python slicing."""
|
||||
instance = _make_instance()
|
||||
instance._exec_cypher = MagicMock(return_value=[
|
||||
{"source": "n0", "relationship": "R", "target": "m0"},
|
||||
])
|
||||
instance.get_all({"user_id": "u1"}, top_k=3)
|
||||
# Verify limit was passed as a parameter to the query
|
||||
cypher_stmt = instance._exec_cypher.call_args[0][0]
|
||||
assert "LIMIT %s" in cypher_stmt
|
||||
params = instance._exec_cypher.call_args[1].get("params") or instance._exec_cypher.call_args[0][2]
|
||||
assert 3 in params
|
||||
|
||||
|
||||
class TestAdd:
|
||||
"""Tests for the add orchestration method."""
|
||||
|
||||
def test_add_returns_added_and_deleted(self):
|
||||
instance = _make_instance()
|
||||
instance._retrieve_nodes_from_data = MagicMock(return_value={"alice": "person"})
|
||||
instance._establish_nodes_relations_from_data = MagicMock(return_value=[
|
||||
{"source": "alice", "relationship": "knows", "destination": "bob"}
|
||||
])
|
||||
instance._search_graph_db = MagicMock(return_value=[])
|
||||
instance._get_delete_entities_from_search_output = MagicMock(return_value=[])
|
||||
instance._delete_entities = MagicMock(return_value=[])
|
||||
instance._add_entities = MagicMock(return_value=["added"])
|
||||
|
||||
result = instance.add("Alice knows Bob", {"user_id": "u1"})
|
||||
assert "deleted_entities" in result
|
||||
assert "added_entities" in result
|
||||
assert result["added_entities"] == ["added"]
|
||||
|
||||
|
||||
class TestSearch:
|
||||
"""Tests for the search method."""
|
||||
|
||||
def test_returns_empty_when_no_search_output(self):
|
||||
instance = _make_instance()
|
||||
instance._retrieve_nodes_from_data = MagicMock(return_value={"alice": "person"})
|
||||
instance._search_graph_db = MagicMock(return_value=[])
|
||||
result = instance.search("Who is Alice?", {"user_id": "u1"})
|
||||
assert result == []
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user