Add support for graph memory using Kuzu (#2934)
This commit is contained in:
@@ -13,7 +13,7 @@ install:
|
||||
install_all:
|
||||
pip install ruff==0.6.9 groq together boto3 litellm ollama chromadb weaviate weaviate-client sentence_transformers vertexai \
|
||||
google-generativeai elasticsearch opensearch-py vecs "pinecone<7.0.0" pinecone-text faiss-cpu langchain-community \
|
||||
upstash-vector azure-search-documents langchain-memgraph langchain-neo4j langchain-aws rank-bm25 pymochow pymongo psycopg
|
||||
upstash-vector azure-search-documents langchain-memgraph langchain-neo4j langchain-aws rank-bm25 pymochow pymongo psycopg kuzu
|
||||
|
||||
# Format code with ruff
|
||||
format:
|
||||
|
||||
@@ -6,7 +6,7 @@ iconType: "solid"
|
||||
---
|
||||
|
||||
Mem0 now supports **Graph Memory**.
|
||||
With Graph Memory, users can now create and utilize complex relationships between pieces of information, allowing for more nuanced and context-aware responses.
|
||||
With Graph Memory, users can now create and utilize complex relationships between pieces of information, allowing for more nuanced and context-aware responses.
|
||||
This integration enables users to leverage the strengths of both vector-based and graph-based approaches, resulting in more accurate and comprehensive information retrieval and generation.
|
||||
|
||||
<Note>
|
||||
@@ -48,8 +48,7 @@ allowfullscreen
|
||||
## Initialize Graph Memory
|
||||
|
||||
To initialize Graph Memory you'll need to set up your configuration with graph
|
||||
store providers. Currently, we support [Neo4j](#initialize-neo4j), [Memgraph](#initialize-memgraph),
|
||||
and [Neptune Analytics](#initialize-neptune-analytics) as graph store providers.
|
||||
store providers. Currently, we support [Neo4j](#initialize-neo4j), [Memgraph](#initialize-memgraph), [Neptune Analytics](#initialize-neptune-analytics), and [Kuzu](#initialize-kuzu) as graph store providers.
|
||||
|
||||
|
||||
### Initialize Neo4j
|
||||
@@ -181,7 +180,7 @@ The `--schema-info-enabled` flag is set to `True` for more performant schema
|
||||
generation.
|
||||
|
||||
Additional information can be found on [Memgraph
|
||||
documentation](https://memgraph.com/docs).
|
||||
documentation](https://memgraph.com/docs).
|
||||
|
||||
User can also customize the LLM for Graph Memory from the [Supported LLM list](https://docs.mem0.ai/components/llms/overview) with three levels of configuration:
|
||||
|
||||
@@ -288,8 +287,54 @@ m = Memory.from_config(config_dict=config)
|
||||
|
||||
- For more details on how to connect, configure, and use the graph_memory graph store, see the [Neptune Analytics example notebook](examples/graph-db-demo/neptune-analytics-example.ipynb).
|
||||
|
||||
### Initialize Kuzu
|
||||
|
||||
[Kuzu](https://kuzudb.com) is a fully local in-process graph database system that runs openCypher queries.
|
||||
Kuzu comes embedded into the Python package and there is no additional setup required.
|
||||
|
||||
Kuzu needs a path to a file where it will store the graph database. For example:
|
||||
|
||||
<CodeGroup>
|
||||
```python Python
|
||||
config = {
|
||||
"graph_store": {
|
||||
"provider": "kuzu",
|
||||
"config": {
|
||||
"db": "/tmp/mem0-example.kuzu"
|
||||
}
|
||||
}
|
||||
}
|
||||
```
|
||||
</CodeGroup>
|
||||
|
||||
Kuzu can also store its database in memory. Note that in this mode, all stored memories will be lost
|
||||
after the program has finished executing.
|
||||
|
||||
<CodeGroup>
|
||||
```python Python
|
||||
config = {
|
||||
"graph_store": {
|
||||
"provider": "kuzu",
|
||||
"config": {
|
||||
"db": ":memory:"
|
||||
}
|
||||
}
|
||||
}
|
||||
```
|
||||
</CodeGroup>
|
||||
|
||||
You can then use the above configuration in the usual way:
|
||||
|
||||
<CodeGroup>
|
||||
```python Python
|
||||
from mem0 import Memory
|
||||
|
||||
m = Memory.from_config(config_dict=config)
|
||||
```
|
||||
</CodeGroup>
|
||||
|
||||
## Graph Operations
|
||||
The Mem0's graph supports the following operations:
|
||||
Mem0's graph memory supports the following operations:
|
||||
|
||||
### Add Memories
|
||||
|
||||
|
||||
@@ -0,0 +1,319 @@
|
||||
{
|
||||
"cells": [
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "ApdaLD4Qi30H"
|
||||
},
|
||||
"source": [
|
||||
"# Kuzu as Graph Memory"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "l7bi3i21i30I"
|
||||
},
|
||||
"source": [
|
||||
"## Prerequisites\n",
|
||||
"\n",
|
||||
"### Install Mem0 with Graph Memory support\n",
|
||||
"\n",
|
||||
"To use Mem0 with Graph Memory support, install it using pip:\n",
|
||||
"\n",
|
||||
"```bash\n",
|
||||
"pip install \"mem0ai[graph]\"\n",
|
||||
"```\n",
|
||||
"\n",
|
||||
"This command installs Mem0 along with the necessary dependencies for graph functionality.\n",
|
||||
"\n",
|
||||
"### Kuzu setup\n",
|
||||
"\n",
|
||||
"Kuzu comes embedded into the Python package that gets installed with the above command. There is no extra setup required.\n",
|
||||
"Just pick an empty directory where Kuzu should persist its database.\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "DkeBdFEpi30I"
|
||||
},
|
||||
"source": [
|
||||
"## Configuration\n",
|
||||
"\n",
|
||||
"Do all the imports and configure OpenAI (enter your OpenAI API key):"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"id": "d99EfBpii30I"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"from mem0 import Memory\n",
|
||||
"from openai import OpenAI\n",
|
||||
"\n",
|
||||
"import os\n",
|
||||
"\n",
|
||||
"os.environ[\"OPENAI_API_KEY\"] = \"\"\n",
|
||||
"openai_client = OpenAI()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "QTucZJjIi30J"
|
||||
},
|
||||
"source": [
|
||||
"Set up configuration to use the embedder model and Neo4j as a graph store:"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 19,
|
||||
"metadata": {
|
||||
"id": "QSE0RFoSi30J"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"config = {\n",
|
||||
" \"embedder\": {\n",
|
||||
" \"provider\": \"openai\",\n",
|
||||
" \"config\": {\"model\": \"text-embedding-3-large\", \"embedding_dims\": 1536},\n",
|
||||
" },\n",
|
||||
" \"graph_store\": {\n",
|
||||
" \"provider\": \"kuzu\",\n",
|
||||
" \"config\": {\n",
|
||||
" \"db\": \":memory:\",\n",
|
||||
" },\n",
|
||||
" },\n",
|
||||
"}\n",
|
||||
"memory = Memory.from_config(config_dict=config)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 20,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"def print_added_memories(results):\n",
|
||||
" print(\"::: Saved the following memories:\")\n",
|
||||
" print(\" embeddings:\")\n",
|
||||
" for r in results['results']:\n",
|
||||
" print(\" \",r)\n",
|
||||
" print(\" relations:\")\n",
|
||||
" for k,v in results['relations'].items():\n",
|
||||
" print(\" \",k)\n",
|
||||
" for e in v:\n",
|
||||
" print(\" \",e)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "kr1fVMwEi30J"
|
||||
},
|
||||
"source": [
|
||||
"## Store memories\n",
|
||||
"\n",
|
||||
"Create memories:"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 21,
|
||||
"metadata": {
|
||||
"id": "sEfogqp_i30J"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"user = \"myuser\"\n",
|
||||
"\n",
|
||||
"messages = [\n",
|
||||
" {\"role\": \"user\", \"content\": \"I'm planning to watch a movie tonight. Any recommendations?\"},\n",
|
||||
" {\"role\": \"assistant\", \"content\": \"How about a thriller movies? They can be quite engaging.\"},\n",
|
||||
" {\"role\": \"user\", \"content\": \"I'm not a big fan of thriller movies but I love sci-fi movies.\"},\n",
|
||||
" {\"role\": \"assistant\", \"content\": \"Got it! I'll avoid thriller recommendations and suggest sci-fi movies in the future.\"}\n",
|
||||
"]"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "gtBHCyIgi30J"
|
||||
},
|
||||
"source": [
|
||||
"Store memories in Kuzu:"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 22,
|
||||
"metadata": {
|
||||
"id": "BMVGgZMFi30K"
|
||||
},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"::: Saved the following memories:\n",
|
||||
" embeddings:\n",
|
||||
" {'id': 'd3e63d11-5f84-4d08-94d8-402959f7b059', 'memory': 'Planning to watch a movie tonight', 'event': 'ADD'}\n",
|
||||
" {'id': 'be561168-56df-4493-ab35-a5e2f0966274', 'memory': 'Not a big fan of thriller movies', 'event': 'ADD'}\n",
|
||||
" {'id': '9bd3db2d-7233-4d82-a257-a5397cb78473', 'memory': 'Loves sci-fi movies', 'event': 'ADD'}\n",
|
||||
" relations:\n",
|
||||
" deleted_entities\n",
|
||||
" added_entities\n",
|
||||
" [{'source': 'myuser', 'relationship': 'plans_to_watch', 'target': 'movie'}]\n",
|
||||
" [{'source': 'movie', 'relationship': 'is_genre', 'target': 'thriller'}]\n",
|
||||
" [{'source': 'movie', 'relationship': 'is_genre', 'target': 'sci-fi'}]\n",
|
||||
" [{'source': 'myuser', 'relationship': 'has_preference', 'target': 'sci-fi'}]\n",
|
||||
" [{'source': 'myuser', 'relationship': 'does_not_prefer', 'target': 'thriller'}]\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"results = memory.add(messages, user_id=user, metadata={\"category\": \"movie_recommendations\"})\n",
|
||||
"print_added_memories(results)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "LBXW7Gv-i30K"
|
||||
},
|
||||
"source": [
|
||||
"## Search memories"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 23,
|
||||
"metadata": {
|
||||
"colab": {
|
||||
"base_uri": "https://localhost:8080/"
|
||||
},
|
||||
"id": "UHFDeQBEi30K",
|
||||
"outputId": "2c69de7d-a79a-48f6-e3c4-bd743067857c"
|
||||
},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"Loves sci-fi movies 0.31536642873409\n",
|
||||
"Planning to watch a movie tonight 0.0967911158879874\n",
|
||||
"Not a big fan of thriller movies 0.09468540071789472\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"for result in memory.search(\"what does alice love?\", user_id=user)[\"results\"]:\n",
|
||||
" print(result[\"memory\"], result[\"score\"])"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"## Chatbot"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 24,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"def chat_with_memories(message: str, user_id: str = user) -> str:\n",
|
||||
" # Retrieve relevant memories\n",
|
||||
" relevant_memories = memory.search(query=message, user_id=user_id, limit=3)\n",
|
||||
" memories_str = \"\\n\".join(f\"- {entry['memory']}\" for entry in relevant_memories[\"results\"])\n",
|
||||
" print(\"::: Using memories:\")\n",
|
||||
" print(memories_str)\n",
|
||||
"\n",
|
||||
" # Generate Assistant response\n",
|
||||
" system_prompt = f\"You are a helpful AI. Answer the question based on query and memories.\\nUser Memories:\\n{memories_str}\"\n",
|
||||
" messages = [{\"role\": \"system\", \"content\": system_prompt}, {\"role\": \"user\", \"content\": message}]\n",
|
||||
" response = openai_client.chat.completions.create(model=\"gpt-4o-mini\", messages=messages)\n",
|
||||
" assistant_response = response.choices[0].message.content\n",
|
||||
"\n",
|
||||
" # Create new memories from the conversation\n",
|
||||
" messages.append({\"role\": \"assistant\", \"content\": assistant_response})\n",
|
||||
" results = memory.add(messages, user_id=user_id)\n",
|
||||
" print_added_memories(results)\n",
|
||||
"\n",
|
||||
" return assistant_response"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 25,
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"Chat with AI (type 'exit' to quit)\n",
|
||||
"::: Using memories:\n",
|
||||
"- Planning to watch a movie tonight\n",
|
||||
"- Not a big fan of thriller movies\n",
|
||||
"- Loves sci-fi movies\n",
|
||||
"::: Saved the following memories:\n",
|
||||
" embeddings:\n",
|
||||
" relations:\n",
|
||||
" deleted_entities\n",
|
||||
" []\n",
|
||||
" added_entities\n",
|
||||
" [{'source': 'myuser', 'relationship': 'loves', 'target': 'sci-fi'}]\n",
|
||||
" [{'source': 'myuser', 'relationship': 'wants_to_avoid', 'target': 'thrillers'}]\n",
|
||||
" [{'source': 'myuser', 'relationship': 'recommends', 'target': 'interstellar'}]\n",
|
||||
" [{'source': 'myuser', 'relationship': 'recommends', 'target': 'the_martian'}]\n",
|
||||
" [{'source': 'interstellar', 'relationship': 'is_a', 'target': 'sci-fi'}]\n",
|
||||
" [{'source': 'the_martian', 'relationship': 'is_a', 'target': 'sci-fi'}]\n",
|
||||
"<<< AI: Since you love sci-fi movies and want to avoid thrillers, I recommend watching \"Interstellar\" if you haven't seen it yet. It's a visually stunning film that explores space travel, time, and love. Another great option is \"The Martian,\" which is more of a fun survival story set on Mars. Both films offer engaging stories and impressive visuals that are characteristic of the sci-fi genre!\n",
|
||||
"Goodbye!\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"print(\"Chat with AI (type 'exit' to quit)\")\n",
|
||||
"while True:\n",
|
||||
" user_input = input(\">>> You: \").strip()\n",
|
||||
" if user_input.lower() == 'exit':\n",
|
||||
" print(\"Goodbye!\")\n",
|
||||
" break\n",
|
||||
" print(f\"<<< AI response:\\n{chat_with_memories(user_input)}\")"
|
||||
]
|
||||
}
|
||||
],
|
||||
"metadata": {
|
||||
"colab": {
|
||||
"provenance": []
|
||||
},
|
||||
"kernelspec": {
|
||||
"display_name": "mem0ai-sQeqgA1d-py3.12",
|
||||
"language": "python",
|
||||
"name": "python3"
|
||||
},
|
||||
"language_info": {
|
||||
"codemirror_mode": {
|
||||
"name": "ipython",
|
||||
"version": 3
|
||||
},
|
||||
"file_extension": ".py",
|
||||
"mimetype": "text/x-python",
|
||||
"name": "python",
|
||||
"nbconvert_exporter": "python",
|
||||
"pygments_lexer": "ipython3",
|
||||
"version": "3.12.10"
|
||||
}
|
||||
},
|
||||
"nbformat": 4,
|
||||
"nbformat_minor": 0
|
||||
}
|
||||
@@ -70,12 +70,16 @@ class NeptuneConfig(BaseModel):
|
||||
)
|
||||
|
||||
|
||||
class KuzuConfig(BaseModel):
|
||||
db: Optional[str] = Field(":memory:", description="Path to a Kuzu database file")
|
||||
|
||||
|
||||
class GraphStoreConfig(BaseModel):
|
||||
provider: str = Field(
|
||||
description="Provider of the data store (e.g., 'neo4j', 'memgraph', 'neptune')",
|
||||
description="Provider of the data store (e.g., 'neo4j', 'memgraph', 'neptune', 'kuzu')",
|
||||
default="neo4j",
|
||||
)
|
||||
config: Union[Neo4jConfig, MemgraphConfig, NeptuneConfig] = Field(
|
||||
config: Union[Neo4jConfig, MemgraphConfig, NeptuneConfig, KuzuConfig] = Field(
|
||||
description="Configuration for the specific data store", default=None
|
||||
)
|
||||
llm: Optional[LlmConfig] = Field(description="LLM configuration for querying the graph store", default=None)
|
||||
@@ -92,5 +96,7 @@ class GraphStoreConfig(BaseModel):
|
||||
return MemgraphConfig(**v.model_dump())
|
||||
elif provider == "neptune":
|
||||
return NeptuneConfig(**v.model_dump())
|
||||
elif provider == "kuzu":
|
||||
return KuzuConfig(**v.model_dump())
|
||||
else:
|
||||
raise ValueError(f"Unsupported graph store provider: {provider}")
|
||||
|
||||
@@ -0,0 +1,710 @@
|
||||
import logging
|
||||
|
||||
from mem0.memory.utils import format_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
|
||||
|
||||
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
|
||||
self.threshold = 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, limit=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.
|
||||
limit (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=limit)
|
||||
|
||||
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_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, limit=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.
|
||||
limit (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": limit,
|
||||
}
|
||||
# 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["arguments"]["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, limit=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": limit,
|
||||
}
|
||||
# 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)[:limit])
|
||||
|
||||
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=0.9)
|
||||
destination_node_search_result = self._search_destination_node(dest_embedding, filters, threshold=0.9)
|
||||
|
||||
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):
|
||||
for item in entity_list:
|
||||
item["source"] = item["source"].lower().replace(" ", "_")
|
||||
item["relationship"] = item["relationship"].lower().replace(" ", "_")
|
||||
item["destination"] = item["destination"].lower().replace(" ", "_")
|
||||
return entity_list
|
||||
|
||||
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)
|
||||
@@ -200,6 +200,7 @@ class GraphStoreFactory:
|
||||
provider_to_class = {
|
||||
"memgraph": "mem0.memory.memgraph_memory.MemoryGraph",
|
||||
"neptune": "mem0.graphs.neptune.main.MemoryGraph",
|
||||
"kuzu": "mem0.memory.kuzu_memory.MemoryGraph",
|
||||
"default": "mem0.memory.graph_memory.MemoryGraph",
|
||||
}
|
||||
|
||||
|
||||
@@ -27,6 +27,7 @@ graph = [
|
||||
"langchain-aws>=0.2.23",
|
||||
"neo4j>=5.23.1",
|
||||
"rank-bm25>=0.2.2",
|
||||
"kuzu>=0.11.0",
|
||||
]
|
||||
vector_stores = [
|
||||
"vecs>=0.4.0",
|
||||
@@ -69,6 +70,9 @@ dev = [
|
||||
"pytest>=8.2.2",
|
||||
]
|
||||
|
||||
[tool.pytest.ini_options]
|
||||
pythonpath = ["."]
|
||||
|
||||
[tool.hatch.build]
|
||||
include = [
|
||||
"mem0/**/*.py",
|
||||
|
||||
@@ -0,0 +1,167 @@
|
||||
import numpy as np
|
||||
import pytest
|
||||
from unittest.mock import Mock, patch
|
||||
from mem0.memory.kuzu_memory import MemoryGraph
|
||||
|
||||
|
||||
class TestKuzu:
|
||||
"""Test that Kuzu memory works correctly"""
|
||||
|
||||
embeddings = {
|
||||
"alice": np.random.uniform(0.0, 0.9, 384).tolist(),
|
||||
"bob": np.random.uniform(0.0, 0.9, 384).tolist(),
|
||||
"charlie": np.random.uniform(0.0, 0.9, 384).tolist(),
|
||||
}
|
||||
|
||||
@pytest.fixture
|
||||
def mock_config(self):
|
||||
"""Create a mock configuration for testing"""
|
||||
config = Mock()
|
||||
|
||||
# Mock embedder config
|
||||
config.embedder.provider = "mock_embedder"
|
||||
config.embedder.config = {"model": "mock_model"}
|
||||
config.vector_store.config = {"dimensions": 384}
|
||||
|
||||
# Mock graph store config
|
||||
config.graph_store.config.db = ":memory:"
|
||||
|
||||
# Mock LLM config
|
||||
config.llm.provider = "mock_llm"
|
||||
config.llm.config = {"api_key": "test_key"}
|
||||
|
||||
return config
|
||||
|
||||
@pytest.fixture
|
||||
def mock_embedding_model(self):
|
||||
"""Create a mock embedding model"""
|
||||
mock_model = Mock()
|
||||
mock_model.config.embedding_dims = 384
|
||||
|
||||
def mock_embed(text):
|
||||
return self.embeddings[text]
|
||||
|
||||
mock_model.embed.side_effect = mock_embed
|
||||
return mock_model
|
||||
|
||||
@pytest.fixture
|
||||
def mock_llm(self):
|
||||
"""Create a mock LLM"""
|
||||
mock_llm = Mock()
|
||||
mock_llm.generate_response.return_value = {
|
||||
"tool_calls": [
|
||||
{
|
||||
"name": "extract_entities",
|
||||
"arguments": {"entities": [{"entity": "test_entity", "entity_type": "test_type"}]},
|
||||
}
|
||||
]
|
||||
}
|
||||
return mock_llm
|
||||
|
||||
@patch("mem0.memory.kuzu_memory.EmbedderFactory")
|
||||
@patch("mem0.memory.kuzu_memory.LlmFactory")
|
||||
def test_kuzu_memory_initialization(
|
||||
self, mock_llm_factory, mock_embedder_factory, mock_config, mock_embedding_model, mock_llm
|
||||
):
|
||||
"""Test that Kuzu memory initializes correctly"""
|
||||
# Setup mocks
|
||||
mock_embedder_factory.create.return_value = mock_embedding_model
|
||||
mock_llm_factory.create.return_value = mock_llm
|
||||
|
||||
# Create instance
|
||||
kuzu_memory = MemoryGraph(mock_config)
|
||||
|
||||
# Verify initialization
|
||||
assert kuzu_memory.config == mock_config
|
||||
assert kuzu_memory.embedding_model == mock_embedding_model
|
||||
assert kuzu_memory.embedding_dims == 384
|
||||
assert kuzu_memory.llm == mock_llm
|
||||
assert kuzu_memory.threshold == 0.7
|
||||
|
||||
@patch("mem0.memory.kuzu_memory.EmbedderFactory")
|
||||
@patch("mem0.memory.kuzu_memory.LlmFactory")
|
||||
def test_kuzu(self, mock_llm_factory, mock_embedder_factory, mock_config, mock_embedding_model, mock_llm):
|
||||
"""Test adding memory to the graph"""
|
||||
mock_embedder_factory.create.return_value = mock_embedding_model
|
||||
mock_llm_factory.create.return_value = mock_llm
|
||||
|
||||
kuzu_memory = MemoryGraph(mock_config)
|
||||
|
||||
filters = {"user_id": "test_user", "agent_id": "test_agent", "run_id": "test_run"}
|
||||
data1 = [
|
||||
{"source": "alice", "destination": "bob", "relationship": "knows"},
|
||||
{"source": "bob", "destination": "charlie", "relationship": "knows"},
|
||||
{"source": "charlie", "destination": "alice", "relationship": "knows"},
|
||||
]
|
||||
data2 = [
|
||||
{"source": "charlie", "destination": "alice", "relationship": "likes"},
|
||||
]
|
||||
|
||||
result = kuzu_memory._add_entities(data1, filters, {})
|
||||
assert result[0] == [{"source": "alice", "relationship": "knows", "target": "bob"}]
|
||||
assert result[1] == [{"source": "bob", "relationship": "knows", "target": "charlie"}]
|
||||
assert result[2] == [{"source": "charlie", "relationship": "knows", "target": "alice"}]
|
||||
assert get_node_count(kuzu_memory) == 3
|
||||
assert get_edge_count(kuzu_memory) == 3
|
||||
|
||||
result = kuzu_memory._add_entities(data2, filters, {})
|
||||
assert result[0] == [{"source": "charlie", "relationship": "likes", "target": "alice"}]
|
||||
assert get_node_count(kuzu_memory) == 3
|
||||
assert get_edge_count(kuzu_memory) == 4
|
||||
|
||||
results = kuzu_memory.get_all(filters)
|
||||
assert set([f"{result['source']}_{result['relationship']}_{result['target']}" for result in results]) == set([
|
||||
"alice_knows_bob",
|
||||
"bob_knows_charlie",
|
||||
"charlie_likes_alice",
|
||||
"charlie_knows_alice"
|
||||
])
|
||||
|
||||
results = kuzu_memory._search_graph_db(["bob"], filters, threshold=0.8)
|
||||
assert set([f"{result['source']}_{result['relationship']}_{result['destination']}" for result in results]) == set([
|
||||
"alice_knows_bob",
|
||||
"bob_knows_charlie",
|
||||
])
|
||||
|
||||
result = kuzu_memory._delete_entities(data2, filters)
|
||||
assert result[0] == [{"source": "charlie", "relationship": "likes", "target": "alice"}]
|
||||
assert get_node_count(kuzu_memory) == 3
|
||||
assert get_edge_count(kuzu_memory) == 3
|
||||
|
||||
result = kuzu_memory._delete_entities(data1, filters)
|
||||
assert result[0] == [{"source": "alice", "relationship": "knows", "target": "bob"}]
|
||||
assert result[1] == [{"source": "bob", "relationship": "knows", "target": "charlie"}]
|
||||
assert result[2] == [{"source": "charlie", "relationship": "knows", "target": "alice"}]
|
||||
assert get_node_count(kuzu_memory) == 3
|
||||
assert get_edge_count(kuzu_memory) == 0
|
||||
|
||||
result = kuzu_memory.delete_all(filters)
|
||||
assert get_node_count(kuzu_memory) == 0
|
||||
assert get_edge_count(kuzu_memory) == 0
|
||||
|
||||
result = kuzu_memory._add_entities(data2, filters, {})
|
||||
assert result[0] == [{"source": "charlie", "relationship": "likes", "target": "alice"}]
|
||||
assert get_node_count(kuzu_memory) == 2
|
||||
assert get_edge_count(kuzu_memory) == 1
|
||||
|
||||
result = kuzu_memory.reset()
|
||||
assert get_node_count(kuzu_memory) == 0
|
||||
assert get_edge_count(kuzu_memory) == 0
|
||||
|
||||
def get_node_count(kuzu_memory):
|
||||
results = kuzu_memory.kuzu_execute(
|
||||
"""
|
||||
MATCH (n)
|
||||
RETURN COUNT(n) as count
|
||||
"""
|
||||
)
|
||||
return int(results[0]['count'])
|
||||
|
||||
def get_edge_count(kuzu_memory):
|
||||
results = kuzu_memory.kuzu_execute(
|
||||
"""
|
||||
MATCH (n)-[e]->(m)
|
||||
RETURN COUNT(e) as count
|
||||
"""
|
||||
)
|
||||
return int(results[0]['count'])
|
||||
Reference in New Issue
Block a user