Add Neptune-DB graph store with vector store (#3443)

Signed-off-by: Andrew Carbonetto <andrew.carbonetto@improving.com>
Co-authored-by: Siddhartha Sahu <dev@sdht.in>
This commit is contained in:
Andrew Carbonetto
2025-09-18 14:24:33 -07:00
committed by GitHub
parent d4e98dba38
commit a015e2ff4a
11 changed files with 2006 additions and 444 deletions
+63 -23
View File
@@ -48,7 +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), [Neptune Analytics](#initialize-neptune-analytics), and [Kuzu](#initialize-kuzu) as graph store providers.
store providers. Currently, we support [Neo4j](#initialize-neo4j), [Memgraph](#initialize-memgraph), [Neptune Analytics](#initialize-neptune-analytics), [Neptune DB Cluster](#initialize-neptune-db),and [Kuzu](#initialize-kuzu) as graph store providers.
### Initialize Neo4j
@@ -231,35 +231,27 @@ m = Memory.from_config(config_dict=config)
### Initialize Neptune Analytics
Mem0 now supports Amazon Neptune Analytics as a graph store provider. This integration allows you to use Neptune Analytics for storing and querying graph-based memories.
Note: You can use Neptune Analytics as part of an Amazon tech stack [Setup AWS Bedrock, AOSS, and Neptune](https://docs.mem0.ai/examples/aws_example#aws-bedrock-and-aoss)
You can use Neptune Analytics as part of an Amazon tech stack [Setup AWS Bedrock, AOSS, and Neptune](https://docs.mem0.ai/examples/aws_example#aws-bedrock-and-aoss)
#### Instance Setup
Create an Amazon Neptune Analytics instance in your AWS account following the [AWS documentation](https://docs.aws.amazon.com/neptune-analytics/latest/userguide/get-started.html).
Create an instance of Amazon Neptune Analytics in your AWS account following the [AWS documentation](https://docs.aws.amazon.com/neptune-analytics/latest/userguide/get-started.html).
- Public connectivity is not enabled by default, and if accessing from outside a VPC, it needs to be enabled.
- Once the Amazon Neptune Analytics instance is available, you will need the graph-identifier to connect.
- The Neptune Analytics instance must be created using the same vector dimensions as the embedding model creates. See: https://docs.aws.amazon.com/neptune-analytics/latest/userguide/vector-index.html
- The Neptune Analytics instance must be created using the same vector dimensions as the embedding model creates. See: [Vector indexing in Neptune Analytics](https://docs.aws.amazon.com/neptune-analytics/latest/userguide/vector-index.html).
#### Attach Credentials
Ensure that you attach your AWS credentials with access to your Amazon Neptune Analytics resources by following the [Configuration and credentials precedence](https://docs.aws.amazon.com/cli/v1/userguide/cli-chap-configure.html#configure-precedence).
Configure your AWS credentials with access to your Amazon Neptune Analytics resources by following the [Configuration and credentials precedence](https://docs.aws.amazon.com/cli/v1/userguide/cli-chap-configure.html#configure-precedence).
- For example, add your SSH access key session token via environment variables:
```bash
export AWS_ACCESS_KEY_ID=your-access-key
export AWS_SECRET_ACCESS_KEY=your-secret-key
export AWS_SESSION_TOKEN=your-session-token
export AWS_DEFAULT_REGION=your-region
```
- The IAM user or role making the request must have a policy attached that allows one of the following IAM actions in that neptune-graph:
The IAM user or role making the request must have a policy attached that allows one of the following IAM actions in that neptune-graph:
- neptune-graph:ReadDataViaQuery
- neptune-graph:WriteDataViaQuery
- neptune-graph:DeleteDataViaQuery
#### Usage
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:
The Neptune memory store uses AWS LangChain Python API to connect to Neptune instances. For additional configuration options for connecting to your Amazon Neptune Analytics instance see [AWS LangChain API documentation](https://python.langchain.com/api_reference/aws/graphs/langchain_aws.graphs.neptune_graph.NeptuneAnalyticsGraph.html).
1. **Main Configuration**: If `llm` is set in the main config, it will be used for all graph operations.
2. **Graph Store Configuration**: If `llm` is set in the graph_store config, it will override the main config `llm` and be used specifically for graph operations.
3. **Default Configuration**: If no custom LLM is set, the default LLM (`gpt-4o-2024-08-06`) will be used for all graph operations.
Here's how you can do it:
<CodeGroup>
```python Python
@@ -279,13 +271,61 @@ m = Memory.from_config(config_dict=config)
```
</CodeGroup>
#### Troubleshooting
Troubleshooting:
- For issues connecting to Amazon Neptune Analytics, please refer to the [Connecting to a graph guide](https://docs.aws.amazon.com/neptune-analytics/latest/userguide/gettingStarted-connecting.html).
- For issues related to authentication, refer to the [boto3 client configuration options](https://boto3.amazonaws.com/v1/documentation/api/latest/guide/configuration.html).
- For more details on how to connect, configure, and use the graph_memory graph store, see the Neptune Analytics example in our [AWS example guide](/examples/aws_example#aws-bedrock-and-aoss).
- The Neptune memory store uses AWS LangChain Python API to connect to Neptune instances. For additional configuration options for connecting to your Amazon Neptune Analytics instance, see [AWS LangChain API documentation](https://python.langchain.com/api_reference/aws/graphs/langchain_aws.graphs.neptune_graph.NeptuneAnalyticsGraph.html).
### Initialize Neptune DB
Note that Neptune DB does not support vectors, and this graph store provider requires a collection in the vector store to save entity vectors.
Create a cluster of Amazon DB instances in your AWS account following the [AWS documentation](https://docs.aws.amazon.com/neptune/latest/userguide/graph-get-started.html).
- Public connectivity is not enabled by default. To access the instance from outside a VPC, public connectivity needs to be enabled on the Neptune DB instance by following [Neptune Public Endpoints](https://docs.aws.amazon.com/neptune/latest/userguide/neptune-public-endpoints.html).
- Once the Amazon Neptune Cluster instance is available, you will need the graph host endpoint to connect.
- Neptune DB doesn't support vectors. The `collection_name` config field can be used to specify the vector store collection used to store vectors for the Neptune entities.
Ensure that you attach your AWS credentials with access to your Amazon Neptune Analytics resources by following the [Configuration and credentials precedence](https://docs.aws.amazon.com/cli/v1/userguide/cli-chap-configure.html#configure-precedence).
The IAM user or role making the request must have a policy attached that allows one of the following IAM actions in that neptune-db:
- neptune-db:ReadDataViaQuery
- neptune-db:WriteDataViaQuery
- neptune-db:DeleteDataViaQuery
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:
1. **Main Configuration**: If `llm` is set in the main config, it will be used for all graph operations.
2. **Graph Store Configuration**: If `llm` is set in the graph_store config, it will override the main config `llm` and be used specifically for graph operations.
3. **Default Configuration**: If no custom LLM is set, the default LLM (`gpt-4o-2024-08-06`) will be used for all graph operations.
Here's how you can do it:
<CodeGroup>
```python Python
from mem0 import Memory
config = {
"graph_store": {
"provider": "neptunedb",
"config": {
"collection_name": "<VECTOR_COLLECTION_NAME>",
"endpoint": "neptune-graph://<HOST_ENDPOINT>",
},
},
}
m = Memory.from_config(config_dict=config)
```
</CodeGroup>
Troubleshooting:
- For issues connecting to Amazon Neptune Analytics, please refer to the [Accessing graph data in Amazon Neptune](https://docs.aws.amazon.com/neptune/latest/userguide/get-started-access-graph.html).
- For issues related to authentication, refer to the [boto3 client configuration options](https://boto3.amazonaws.com/v1/documentation/api/latest/guide/configuration.html).
- For more details on how to connect, configure, and use the graph_memory graph store, see the [Neptune DB example notebook](examples/graph-db-demo/neptune-example.ipynb).
- The Neptune memory store uses AWS LangChain Python API to connect to Neptune instances. For additional configuration options for connecting to your Amazon Neptune Analytics instance, see [AWS LangChain API documentation](https://python.langchain.com/api_reference/aws/graphs/langchain_aws.graphs.neptune_graph.NeptuneGraph.html).
### Initialize Kuzu
@@ -0,0 +1,459 @@
{
"cells": [
{
"cell_type": "markdown",
"metadata": {},
"source": [
"# Neptune as Graph Memory\n",
"\n",
"In this notebook, we will be connecting using an Amazon Neptune DC Cluster instance as our memory graph storage for Mem0. Unlike other graph stores, Neptune DB doesn't store vectors itself. To detect vector similary in nodes, we store the node vectors in our defined vector store, and use vector search to retrieve similar nodes.\n",
"\n",
"For this reason, a vector store is required to configure neptune-db.\n",
"\n",
"The Graph Memory storage persists memories in a graph or relationship form when performing `m.add` memory operations."
]
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"## Prerequisites\n",
"\n",
"### 1. Install Mem0 with Graph Memory support \n",
"\n",
"To use Mem0 with Graph Memory support (as well as other Amazon services), use pip install:\n",
"\n",
"```bash\n",
"pip install \"mem0ai[graph,vector_stores,extras]\"\n",
"```\n",
"\n",
"This command installs Mem0 along with the necessary dependencies for graph functionality (`graph`), vector stores, and other Amazon dependencies (`extras`).\n",
"\n",
"### 2. Connect to Amazon services\n",
"\n",
"For this sample notebook, configure `mem0ai` with [Amazon Neptune Database Cluster](https://docs.aws.amazon.com/neptune/latest/userguide/intro.html) as the graph store, [Amazon OpenSearch Serverless](https://docs.aws.amazon.com/opensearch-service/latest/developerguide/serverless-overview.html) as the vector store, and [Amazon Bedrock](https://docs.aws.amazon.com/bedrock/latest/userguide/what-is-bedrock.html) for generating embeddings.\n",
"\n",
"Your configuration should look similar to:\n",
"\n",
"```python\n",
"config = {\n",
" \"embedder\": {\n",
" \"provider\": \"aws_bedrock\",\n",
" \"config\": {\n",
" \"model\": \"amazon.titan-embed-text-v2:0\"\n",
" }\n",
" },\n",
" \"llm\": {\n",
" \"provider\": \"aws_bedrock\",\n",
" \"config\": {\n",
" \"model\": \"us.anthropic.claude-3-7-sonnet-20250219-v1:0\",\n",
" \"temperature\": 0.1,\n",
" \"max_tokens\": 2000\n",
" }\n",
" },\n",
" \"vector_store\": {\n",
" \"provider\": \"opensearch\",\n",
" \"config\": {\n",
" \"collection_name\": \"mem0\",\n",
" \"host\": \"your-opensearch-domain.us-west-2.es.amazonaws.com\",\n",
" \"port\": 443,\n",
" \"http_auth\": auth,\n",
" \"connection_class\": RequestsHttpConnection,\n",
" \"pool_maxsize\": 20,\n",
" \"use_ssl\": True,\n",
" \"verify_certs\": True,\n",
" \"embedding_model_dims\": 1024,\n",
" }\n",
" },\n",
" \"graph_store\": {\n",
" \"provider\": \"neptunedb\",\n",
" \"config\": {\n",
" \"\": \"\",\n",
" \"endpoint\": f\"neptune-db://my-graph-host\",\n",
" },\n",
" },\n",
"}\n",
"```"
]
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"## Setup\n",
"\n",
"Import all packages and setup logging"
]
},
{
"cell_type": "code",
"metadata": {},
"source": [
"from mem0 import Memory\n",
"import os\n",
"import logging\n",
"import sys\n",
"import boto3\n",
"from opensearchpy import RequestsHttpConnection, AWSV4SignerAuth\n",
"from dotenv import load_dotenv\n",
"\n",
"load_dotenv()\n",
"\n",
"logging.getLogger(\"mem0.graphs.neptune.neptunedb\").setLevel(logging.DEBUG)\n",
"logging.getLogger(\"mem0.graphs.neptune.base\").setLevel(logging.DEBUG)\n",
"logger = logging.getLogger(__name__)\n",
"logger.setLevel(logging.DEBUG)\n",
"\n",
"logging.basicConfig(\n",
" format=\"%(levelname)s - %(message)s\",\n",
" datefmt=\"%Y-%m-%d %H:%M:%S\",\n",
" stream=sys.stdout, # Explicitly set output to stdout\n",
")"
],
"outputs": [],
"execution_count": null
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"Setup the Mem0 configuration using:\n",
"- Amazon Bedrock as the LLM and embedder\n",
"- Amazon Neptune DB instance as a graph store with node vectors in OpenSearch (collection: `mem0ai_neptune_entities`)\n",
"- OpenSearch as the text summaries vector store (collection: `mem0ai_text_summaries`)"
]
},
{
"cell_type": "code",
"metadata": {},
"source": [
"bedrock_embedder_model = \"amazon.titan-embed-text-v2:0\"\n",
"bedrock_llm_model = \"us.anthropic.claude-3-7-sonnet-20250219-v1:0\"\n",
"embedding_model_dims = 1024\n",
"\n",
"neptune_host = os.environ.get(\"GRAPH_HOST\")\n",
"\n",
"opensearch_host = os.environ.get(\"OS_HOST\")\n",
"opensearch_port = 443\n",
"\n",
"credentials = boto3.Session().get_credentials()\n",
"region = os.environ.get(\"AWS_REGION\")\n",
"auth = AWSV4SignerAuth(credentials, region)\n",
"\n",
"config = {\n",
" \"embedder\": {\n",
" \"provider\": \"aws_bedrock\",\n",
" \"config\": {\n",
" \"model\": bedrock_embedder_model,\n",
" }\n",
" },\n",
" \"llm\": {\n",
" \"provider\": \"aws_bedrock\",\n",
" \"config\": {\n",
" \"model\": bedrock_llm_model,\n",
" \"temperature\": 0.1,\n",
" \"max_tokens\": 2000\n",
" }\n",
" },\n",
" \"vector_store\": {\n",
" \"provider\": \"opensearch\",\n",
" \"config\": {\n",
" \"collection_name\": \"mem0ai_text_summaries\",\n",
" \"host\": opensearch_host,\n",
" \"port\": opensearch_port,\n",
" \"http_auth\": auth,\n",
" \"embedding_model_dims\": embedding_model_dims,\n",
" \"use_ssl\": True,\n",
" \"verify_certs\": True,\n",
" \"connection_class\": RequestsHttpConnection,\n",
" },\n",
" },\n",
" \"graph_store\": {\n",
" \"provider\": \"neptunedb\",\n",
" \"config\": {\n",
" \"collection_name\": \"mem0ai_neptune_entities\",\n",
" \"endpoint\": f\"neptune-db://{neptune_host}\",\n",
" },\n",
" },\n",
"}"
],
"outputs": [],
"execution_count": null
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"## Graph Memory initializiation\n",
"\n",
"Initialize Memgraph as a Graph Memory store:"
]
},
{
"cell_type": "code",
"metadata": {},
"source": [
"m = Memory.from_config(config_dict=config)\n",
"\n",
"app_id = \"movies\"\n",
"user_id = \"alice\"\n",
"\n",
"m.delete_all(user_id=user_id)"
],
"outputs": [],
"execution_count": null
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"## Store memories\n",
"\n",
"Create memories and store one at a time:"
]
},
{
"cell_type": "code",
"metadata": {},
"source": [
"messages = [\n",
" {\n",
" \"role\": \"user\",\n",
" \"content\": \"I'm planning to watch a movie tonight. Any recommendations?\",\n",
" },\n",
"]\n",
"\n",
"# Store inferred memories (default behavior)\n",
"result = m.add(messages, user_id=user_id, metadata={\"category\": \"movie_recommendations\"})\n",
"\n",
"all_results = m.get_all(user_id=user_id)\n",
"for n in all_results[\"results\"]:\n",
" print(f\"node \\\"{n['memory']}\\\": [hash: {n['hash']}]\")\n",
"\n",
"for e in all_results[\"relations\"]:\n",
" print(f\"edge \\\"{e['source']}\\\" --{e['relationship']}--> \\\"{e['target']}\\\"\")"
],
"outputs": [],
"execution_count": null
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"## Graph Explorer Visualization\n",
"\n",
"You can visualize the graph using a Graph Explorer connection to Neptune-DB in Neptune Notebooks in the Amazon console. See [Using Amazon Neptune with graph notebooks](https://docs.aws.amazon.com/neptune/latest/userguide/graph-notebooks.html) for instructions on how to setup a Neptune Notebook with Graph Explorer.\n",
"\n",
"Once the graph has been generated, you can open the visualization in the Neptune > Notebooks and click on Actions > Open Graph Explorer. This will automatically connect to your neptune db graph that was provided in the notebook setup.\n",
"\n",
"Once in Graph Explorer, visit Open Connections and send all the available nodes and edges to Explorer. Visit Open Graph Explorer to see the nodes and edges in the graph.\n",
"\n",
"### Graph Explorer Visualization Example\n",
"\n",
"_Note that the visualization given below represents only a single example of the possible results generated by the LLM._\n",
"\n",
"Visualization for the relationship:\n",
"```\n",
"\"alice\" --plans_to_watch--> \"movie\"\n",
"```\n",
"\n",
"![neptune-example-visualization-1.png](./neptune-example-visualization-1.png)"
]
},
{
"cell_type": "code",
"metadata": {},
"source": [
"messages = [\n",
" {\n",
" \"role\": \"assistant\",\n",
" \"content\": \"How about a thriller movies? They can be quite engaging.\",\n",
" },\n",
"]\n",
"\n",
"# Store inferred memories (default behavior)\n",
"result = m.add(messages, user_id=user_id, metadata={\"category\": \"movie_recommendations\"})\n",
"\n",
"all_results = m.get_all(user_id=user_id)\n",
"for n in all_results[\"results\"]:\n",
" print(f\"node \\\"{n['memory']}\\\": [hash: {n['hash']}]\")\n",
"\n",
"for e in all_results[\"relations\"]:\n",
" print(f\"edge \\\"{e['source']}\\\" --{e['relationship']}--> \\\"{e['target']}\\\"\")"
],
"outputs": [],
"execution_count": null
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"### Graph Explorer Visualization Example\n",
"\n",
"_Note that the visualization given below represents only a single example of the possible results generated by the LLM._\n",
"\n",
"Visualization for the relationship:\n",
"```\n",
"\"alice\" --plans_to_watch--> \"movie\"\n",
"\"thriller\" --type_of--> \"movie\"\n",
"\"movie\" --can_be--> \"engaging\"\n",
"```\n",
"\n",
"![neptune-example-visualization-2.png](./neptune-example-visualization-2.png)"
]
},
{
"cell_type": "code",
"metadata": {},
"source": [
"messages = [\n",
" {\n",
" \"role\": \"user\",\n",
" \"content\": \"I'm not a big fan of thriller movies but I love sci-fi movies.\",\n",
" },\n",
"]\n",
"\n",
"# Store inferred memories (default behavior)\n",
"result = m.add(messages, user_id=user_id, metadata={\"category\": \"movie_recommendations\"})\n",
"\n",
"all_results = m.get_all(user_id=user_id)\n",
"for n in all_results[\"results\"]:\n",
" print(f\"node \\\"{n['memory']}\\\": [hash: {n['hash']}]\")\n",
"\n",
"for e in all_results[\"relations\"]:\n",
" print(f\"edge \\\"{e['source']}\\\" --{e['relationship']}--> \\\"{e['target']}\\\"\")"
],
"outputs": [],
"execution_count": null
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"### Graph Explorer Visualization Example\n",
"\n",
"_Note that the visualization given below represents only a single example of the possible results generated by the LLM._\n",
"\n",
"Visualization for the relationship:\n",
"```\n",
"\"alice\" --dislikes--> \"thriller_movies\"\n",
"\"alice\" --loves--> \"sci-fi_movies\"\n",
"\"alice\" --plans_to_watch--> \"movie\"\n",
"\"thriller\" --type_of--> \"movie\"\n",
"\"movie\" --can_be--> \"engaging\"\n",
"```\n",
"\n",
"![neptune-example-visualization-3.png](./neptune-example-visualization-3.png)"
]
},
{
"cell_type": "code",
"metadata": {},
"source": [
"messages = [\n",
" {\n",
" \"role\": \"assistant\",\n",
" \"content\": \"Got it! I'll avoid thriller recommendations and suggest sci-fi movies in the future.\",\n",
" },\n",
"]\n",
"\n",
"# Store inferred memories (default behavior)\n",
"result = m.add(messages, user_id=user_id, metadata={\"category\": \"movie_recommendations\"})\n",
"\n",
"all_results = m.get_all(user_id=user_id)\n",
"for n in all_results[\"results\"]:\n",
" print(f\"node \\\"{n['memory']}\\\": [hash: {n['hash']}]\")\n",
"\n",
"for e in all_results[\"relations\"]:\n",
" print(f\"edge \\\"{e['source']}\\\" --{e['relationship']}--> \\\"{e['target']}\\\"\")"
],
"outputs": [],
"execution_count": null
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"### Graph Explorer Visualization Example\n",
"\n",
"_Note that the visualization given below represents only a single example of the possible results generated by the LLM._\n",
"\n",
"Visualization for the relationship:\n",
"```\n",
"\"alice\" --recommends--> \"sci-fi\"\n",
"\"alice\" --dislikes--> \"thriller_movies\"\n",
"\"alice\" --loves--> \"sci-fi_movies\"\n",
"\"alice\" --plans_to_watch--> \"movie\"\n",
"\"alice\" --avoids--> \"thriller\"\n",
"\"thriller\" --type_of--> \"movie\"\n",
"\"movie\" --can_be--> \"engaging\"\n",
"\"sci-fi\" --type_of--> \"movie\"\n",
"```\n",
"\n",
"![neptune-example-visualization-4.png](./neptune-example-visualization-4.png)"
]
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"## Search memories\n",
"\n",
"Search all memories for \"what does alice love?\". Since \"alice\" the user, this will search for a relationship that fits the users love of \"sci-fi\" movies and dislike of \"thriller\" movies."
]
},
{
"cell_type": "code",
"metadata": {},
"source": [
"search_results = m.search(\"what does alice love?\", user_id=user_id)\n",
"for result in search_results[\"results\"]:\n",
" print(f\"\\\"{result['memory']}\\\" [score: {result['score']}]\")\n",
"for relation in search_results[\"relations\"]:\n",
" print(f\"{relation}\")"
],
"outputs": [],
"execution_count": null
},
{
"cell_type": "code",
"metadata": {},
"source": [
"m.delete_all(user_id)\n",
"m.reset()"
],
"outputs": [],
"execution_count": null
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"## Conclusion\n",
"\n",
"In this example we demonstrated how an AWS tech stack can be used to store and retrieve memory context. Bedrock LLM models can be used to interpret given conversations. OpenSearch can store text chunks with vector embeddings. Neptune Database can store the text entities in a graph format with relationship entities."
]
}
],
"metadata": {
"kernelspec": {
"display_name": ".venv",
"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.13.2"
}
},
"nbformat": 4,
"nbformat_minor": 2
}
+6 -4
View File
@@ -46,18 +46,20 @@ class NeptuneConfig(BaseModel):
endpoint: Optional[str] = (
Field(
None,
description="Endpoint to connect to a Neptune Analytics Server as neptune-graph://<graphid>",
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-graph://<graphid>'.")
raise ValueError("Please provide 'endpoint' with the format as 'neptune-db://<endpoint>' or 'neptune-graph://<graphid>'.")
if endpoint.startswith("neptune-db://"):
raise ValueError("neptune-db server is not yet supported")
# 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://", "")
@@ -95,7 +97,7 @@ class GraphStoreConfig(BaseModel):
return Neo4jConfig(**v.model_dump())
elif provider == "memgraph":
return MemgraphConfig(**v.model_dump())
elif provider == "neptune":
elif provider == "neptune" or provider == "neptunedb":
return NeptuneConfig(**v.model_dump())
elif provider == "kuzu":
return KuzuConfig(**v.model_dump())
+89 -2
View File
@@ -17,7 +17,7 @@ from mem0.graphs.tools import (
RELATIONS_TOOL,
)
from mem0.graphs.utils import EXTRACT_RELATIONS_PROMPT, get_delete_messages
from mem0.utils.factory import EmbedderFactory, LlmFactory
from mem0.utils.factory import EmbedderFactory, LlmFactory, VectorStoreFactory
logger = logging.getLogger(__name__)
@@ -46,6 +46,15 @@ class NeptuneBase(ABC):
"""
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.
@@ -244,7 +253,6 @@ class NeptuneBase(ABC):
results.append(result)
return results
@abstractmethod
def _add_entities_cypher(
self,
source_node_list,
@@ -261,6 +269,85 @@ class NeptuneBase(ABC):
"""
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, limit=100):
-402
View File
@@ -1,402 +0,0 @@
import logging
from .base import NeptuneBase
try:
from langchain_aws import NeptuneAnalyticsGraph
from botocore.config import Config
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)
self.llm_provider = "openai_structured"
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
self.threshold = 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_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
:param source_node_list: list of source nodes
: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 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 ""
# Refactor this code with the graph_memory.py implementation
if not destination_node_list and source_node_list:
cypher = f"""
MATCH (source)
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,
}
elif destination_node_list and not source_node_list:
cypher = f"""
MATCH (destination)
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,
}
elif source_node_list and destination_node_list:
cypher = f"""
MATCH (source)
WHERE id(source) = $source_id
SET
source.mentions = coalesce(source.mentions, 0) + 1,
source.updated = timestamp()
WITH source
MATCH (destination)
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,
}
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_entities:\n destination_node_search_result={destination_node_list}\n source_node_search_result={source_node_list}\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, limit):
"""
Returns the OpenCypher query and parameters to get all edges/nodes in the memory store
:param filters: search filters
:param limit: 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": limit}
return cypher, params
def _search_graph_db_cypher(self, n_embedding, filters, limit):
"""
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 limit: 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": limit,
}
logger.debug(f"_search_graph_db\n query={cypher_query}")
return cypher_query, params
+511
View File
@@ -0,0 +1,511 @@
import logging
import uuid
from datetime import datetime
import pytz
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
self.threshold = 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(pytz.timezone("US/Pacific")).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(pytz.timezone("US/Pacific")).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(pytz.timezone("US/Pacific")).isoformat(),
}
destination_id = str(uuid.uuid4())
destination_payload = {
"name": destination,
"type": destination_type,
"user_id": user_id,
"created_at": datetime.now(pytz.timezone("US/Pacific")).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,
limit=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,
limit=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, limit):
"""
Returns the OpenCypher query and parameters to get all edges/nodes in the memory store
:param filters: search filters
:param limit: 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": limit}
return cypher, params
def _search_graph_db_cypher(self, n_embedding, filters, limit):
"""
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 limit: return limit
:return: str, dict
"""
# search vector store for applicable nodes using cosine similarity
search_nodes = self.vector_store.search(
query="",
vectors=n_embedding,
limit=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": limit,
}
logger.debug(f"_search_graph_db\n query={cypher_query}")
return cypher_query, params
+474
View File
@@ -0,0 +1,474 @@
import logging
from .base import NeptuneBase
try:
from langchain_aws import NeptuneAnalyticsGraph
from botocore.config import Config
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
self.threshold = 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, limit):
"""
Returns the OpenCypher query and parameters to get all edges/nodes in the memory store
:param filters: search filters
:param limit: 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": limit}
return cypher, params
def _search_graph_db_cypher(self, n_embedding, filters, limit):
"""
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 limit: 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": limit,
}
logger.debug(f"_search_graph_db\n query={cypher_query}")
return cypher_query, params
+2
View File
@@ -812,9 +812,11 @@ class Memory(MemoryBase):
keys, encoded_ids = process_telemetry_filters(filters)
capture_event("mem0.delete_all", self, {"keys": keys, "encoded_ids": encoded_ids, "sync_type": "sync"})
# delete all vector memories and reset the collections
memories = self.vector_store.list(filters=filters)[0]
for memory in memories:
self._delete_memory(memory.id)
self.vector_store.reset()
logger.info(f"Deleted {len(memories)} memories")
+2 -1
View File
@@ -204,7 +204,8 @@ class GraphStoreFactory:
provider_to_class = {
"memgraph": "mem0.memory.memgraph_memory.MemoryGraph",
"neptune": "mem0.graphs.neptune.main.MemoryGraph",
"neptune": "mem0.graphs.neptune.neptunegraph.MemoryGraph",
"neptunedb": "mem0.graphs.neptune.neptunedb.MemoryGraph",
"kuzu": "mem0.memory.kuzu_memory.MemoryGraph",
"default": "mem0.memory.graph_memory.MemoryGraph",
}
@@ -0,0 +1,335 @@
import unittest
from unittest.mock import MagicMock, patch
import pytest
from mem0.graphs.neptune.neptunegraph import MemoryGraph
from mem0.graphs.neptune.base import NeptuneBase
class TestNeptuneMemory(unittest.TestCase):
"""Test suite for the Neptune Memory implementation."""
def setUp(self):
"""Set up test fixtures before each test method."""
# Create a mock config
self.config = MagicMock()
self.config.graph_store.config.endpoint = "neptune-graph://test-graph"
self.config.graph_store.config.base_label = True
self.config.llm.provider = "openai_structured"
self.config.graph_store.llm = None
self.config.graph_store.custom_prompt = None
# Create mock for NeptuneAnalyticsGraph
self.mock_graph = MagicMock()
self.mock_graph.client.get_graph.return_value = {"status": "AVAILABLE"}
# Create mocks for static methods
self.mock_embedding_model = MagicMock()
self.mock_llm = MagicMock()
# Patch the necessary components
self.neptune_analytics_graph_patcher = patch("mem0.graphs.neptune.neptunegraph.NeptuneAnalyticsGraph")
self.mock_neptune_analytics_graph = self.neptune_analytics_graph_patcher.start()
self.mock_neptune_analytics_graph.return_value = self.mock_graph
# Patch the static methods
self.create_embedding_model_patcher = patch.object(NeptuneBase, "_create_embedding_model")
self.mock_create_embedding_model = self.create_embedding_model_patcher.start()
self.mock_create_embedding_model.return_value = self.mock_embedding_model
self.create_llm_patcher = patch.object(NeptuneBase, "_create_llm")
self.mock_create_llm = self.create_llm_patcher.start()
self.mock_create_llm.return_value = self.mock_llm
# Create the MemoryGraph instance
self.memory_graph = MemoryGraph(self.config)
# Set up common test data
self.user_id = "test_user"
self.test_filters = {"user_id": self.user_id}
def tearDown(self):
"""Tear down test fixtures after each test method."""
self.neptune_analytics_graph_patcher.stop()
self.create_embedding_model_patcher.stop()
self.create_llm_patcher.stop()
def test_initialization(self):
"""Test that the MemoryGraph is initialized correctly."""
self.assertEqual(self.memory_graph.graph, self.mock_graph)
self.assertEqual(self.memory_graph.embedding_model, self.mock_embedding_model)
self.assertEqual(self.memory_graph.llm, self.mock_llm)
self.assertEqual(self.memory_graph.llm_provider, "openai_structured")
self.assertEqual(self.memory_graph.node_label, ":`__Entity__`")
self.assertEqual(self.memory_graph.threshold, 0.7)
def test_init(self):
"""Test the class init functions"""
# Create a mock config with bad endpoint
config_no_endpoint = MagicMock()
config_no_endpoint.graph_store.config.endpoint = None
# Create the MemoryGraph instance
with pytest.raises(ValueError):
MemoryGraph(config_no_endpoint)
# Create a mock config with bad endpoint
config_ndb_endpoint = MagicMock()
config_ndb_endpoint.graph_store.config.endpoint = "neptune-db://test-graph"
with pytest.raises(ValueError):
MemoryGraph(config_ndb_endpoint)
def test_add_method(self):
"""Test the add method with mocked components."""
# Mock the necessary methods that add() calls
self.memory_graph._retrieve_nodes_from_data = MagicMock(return_value={"alice": "person", "bob": "person"})
self.memory_graph._establish_nodes_relations_from_data = MagicMock(
return_value=[{"source": "alice", "relationship": "knows", "destination": "bob"}]
)
self.memory_graph._search_graph_db = MagicMock(return_value=[])
self.memory_graph._get_delete_entities_from_search_output = MagicMock(return_value=[])
self.memory_graph._delete_entities = MagicMock(return_value=[])
self.memory_graph._add_entities = MagicMock(
return_value=[{"source": "alice", "relationship": "knows", "target": "bob"}]
)
# Call the add method
result = self.memory_graph.add("Alice knows Bob", self.test_filters)
# Verify the method calls
self.memory_graph._retrieve_nodes_from_data.assert_called_once_with("Alice knows Bob", self.test_filters)
self.memory_graph._establish_nodes_relations_from_data.assert_called_once()
self.memory_graph._search_graph_db.assert_called_once()
self.memory_graph._get_delete_entities_from_search_output.assert_called_once()
self.memory_graph._delete_entities.assert_called_once_with([], self.user_id)
self.memory_graph._add_entities.assert_called_once()
# Check the result structure
self.assertIn("deleted_entities", result)
self.assertIn("added_entities", result)
def test_search_method(self):
"""Test the search method with mocked components."""
# Mock the necessary methods that search() calls
self.memory_graph._retrieve_nodes_from_data = MagicMock(return_value={"alice": "person"})
# Mock search results
mock_search_results = [
{"source": "alice", "relationship": "knows", "destination": "bob"},
{"source": "alice", "relationship": "works_with", "destination": "charlie"},
]
self.memory_graph._search_graph_db = MagicMock(return_value=mock_search_results)
# Mock BM25Okapi
with patch("mem0.graphs.neptune.base.BM25Okapi") as mock_bm25:
mock_bm25_instance = MagicMock()
mock_bm25.return_value = mock_bm25_instance
# Mock get_top_n to return reranked results
reranked_results = [["alice", "knows", "bob"], ["alice", "works_with", "charlie"]]
mock_bm25_instance.get_top_n.return_value = reranked_results
# Call the search method
result = self.memory_graph.search("Find Alice", self.test_filters, limit=5)
# Verify the method calls
self.memory_graph._retrieve_nodes_from_data.assert_called_once_with("Find Alice", self.test_filters)
self.memory_graph._search_graph_db.assert_called_once_with(node_list=["alice"], filters=self.test_filters)
# Check the result structure
self.assertEqual(len(result), 2)
self.assertEqual(result[0]["source"], "alice")
self.assertEqual(result[0]["relationship"], "knows")
self.assertEqual(result[0]["destination"], "bob")
def test_get_all_method(self):
"""Test the get_all method."""
# Mock the _get_all_cypher method
mock_cypher = "MATCH (n) RETURN n"
mock_params = {"user_id": self.user_id, "limit": 10}
self.memory_graph._get_all_cypher = MagicMock(return_value=(mock_cypher, mock_params))
# Mock the graph.query result
mock_query_result = [
{"source": "alice", "relationship": "knows", "target": "bob"},
{"source": "bob", "relationship": "works_with", "target": "charlie"},
]
self.mock_graph.query.return_value = mock_query_result
# Call the get_all method
result = self.memory_graph.get_all(self.test_filters, limit=10)
# Verify the method calls
self.memory_graph._get_all_cypher.assert_called_once_with(self.test_filters, 10)
self.mock_graph.query.assert_called_once_with(mock_cypher, params=mock_params)
# Check the result structure
self.assertEqual(len(result), 2)
self.assertEqual(result[0]["source"], "alice")
self.assertEqual(result[0]["relationship"], "knows")
self.assertEqual(result[0]["target"], "bob")
def test_delete_all_method(self):
"""Test the delete_all method."""
# Mock the _delete_all_cypher method
mock_cypher = "MATCH (n) DETACH DELETE n"
mock_params = {"user_id": self.user_id}
self.memory_graph._delete_all_cypher = MagicMock(return_value=(mock_cypher, mock_params))
# Call the delete_all method
self.memory_graph.delete_all(self.test_filters)
# Verify the method calls
self.memory_graph._delete_all_cypher.assert_called_once_with(self.test_filters)
self.mock_graph.query.assert_called_once_with(mock_cypher, params=mock_params)
def test_search_source_node(self):
"""Test the _search_source_node method."""
# Mock embedding
mock_embedding = [0.1, 0.2, 0.3]
# Mock the _search_source_node_cypher method
mock_cypher = "MATCH (n) RETURN n"
mock_params = {"source_embedding": mock_embedding, "user_id": self.user_id, "threshold": 0.9}
self.memory_graph._search_source_node_cypher = MagicMock(return_value=(mock_cypher, mock_params))
# Mock the graph.query result
mock_query_result = [{"id(source_candidate)": 123, "cosine_similarity": 0.95}]
self.mock_graph.query.return_value = mock_query_result
# Call the _search_source_node method
result = self.memory_graph._search_source_node(mock_embedding, self.user_id, threshold=0.9)
# Verify the method calls
self.memory_graph._search_source_node_cypher.assert_called_once_with(mock_embedding, self.user_id, 0.9)
self.mock_graph.query.assert_called_once_with(mock_cypher, params=mock_params)
# Check the result
self.assertEqual(result, mock_query_result)
def test_search_destination_node(self):
"""Test the _search_destination_node method."""
# Mock embedding
mock_embedding = [0.1, 0.2, 0.3]
# Mock the _search_destination_node_cypher method
mock_cypher = "MATCH (n) RETURN n"
mock_params = {"destination_embedding": mock_embedding, "user_id": self.user_id, "threshold": 0.9}
self.memory_graph._search_destination_node_cypher = MagicMock(return_value=(mock_cypher, mock_params))
# Mock the graph.query result
mock_query_result = [{"id(destination_candidate)": 456, "cosine_similarity": 0.92}]
self.mock_graph.query.return_value = mock_query_result
# Call the _search_destination_node method
result = self.memory_graph._search_destination_node(mock_embedding, self.user_id, threshold=0.9)
# Verify the method calls
self.memory_graph._search_destination_node_cypher.assert_called_once_with(mock_embedding, self.user_id, 0.9)
self.mock_graph.query.assert_called_once_with(mock_cypher, params=mock_params)
# Check the result
self.assertEqual(result, mock_query_result)
def test_search_graph_db(self):
"""Test the _search_graph_db method."""
# Mock node list
node_list = ["alice", "bob"]
# Mock embedding
mock_embedding = [0.1, 0.2, 0.3]
self.mock_embedding_model.embed.return_value = mock_embedding
# Mock the _search_graph_db_cypher method
mock_cypher = "MATCH (n) RETURN n"
mock_params = {"n_embedding": mock_embedding, "user_id": self.user_id, "threshold": 0.7, "limit": 10}
self.memory_graph._search_graph_db_cypher = MagicMock(return_value=(mock_cypher, mock_params))
# Mock the graph.query results
mock_query_result1 = [{"source": "alice", "relationship": "knows", "destination": "bob"}]
mock_query_result2 = [{"source": "bob", "relationship": "works_with", "destination": "charlie"}]
self.mock_graph.query.side_effect = [mock_query_result1, mock_query_result2]
# Call the _search_graph_db method
result = self.memory_graph._search_graph_db(node_list, self.test_filters, limit=10)
# Verify the method calls
self.assertEqual(self.mock_embedding_model.embed.call_count, 2)
self.assertEqual(self.memory_graph._search_graph_db_cypher.call_count, 2)
self.assertEqual(self.mock_graph.query.call_count, 2)
# Check the result
expected_result = mock_query_result1 + mock_query_result2
self.assertEqual(result, expected_result)
def test_add_entities(self):
"""Test the _add_entities method."""
# Mock data
to_be_added = [{"source": "alice", "relationship": "knows", "destination": "bob"}]
entity_type_map = {"alice": "person", "bob": "person"}
# Mock embeddings
mock_embedding = [0.1, 0.2, 0.3]
self.mock_embedding_model.embed.return_value = mock_embedding
# Mock search results
mock_source_search = [{"id(source_candidate)": 123, "cosine_similarity": 0.95}]
mock_dest_search = [{"id(destination_candidate)": 456, "cosine_similarity": 0.92}]
# Mock the search methods
self.memory_graph._search_source_node = MagicMock(return_value=mock_source_search)
self.memory_graph._search_destination_node = MagicMock(return_value=mock_dest_search)
# Mock the _add_entities_cypher method
mock_cypher = "MATCH (n) RETURN n"
mock_params = {"source_id": 123, "destination_id": 456}
self.memory_graph._add_entities_cypher = MagicMock(return_value=(mock_cypher, mock_params))
# Mock the graph.query result
mock_query_result = [{"source": "alice", "relationship": "knows", "target": "bob"}]
self.mock_graph.query.return_value = mock_query_result
# Call the _add_entities method
result = self.memory_graph._add_entities(to_be_added, self.user_id, entity_type_map)
# Verify the method calls
self.assertEqual(self.mock_embedding_model.embed.call_count, 2)
self.memory_graph._search_source_node.assert_called_once_with(mock_embedding, self.user_id, threshold=0.9)
self.memory_graph._search_destination_node.assert_called_once_with(mock_embedding, self.user_id, threshold=0.9)
self.memory_graph._add_entities_cypher.assert_called_once()
self.mock_graph.query.assert_called_once_with(mock_cypher, params=mock_params)
# Check the result
self.assertEqual(result, [mock_query_result])
def test_delete_entities(self):
"""Test the _delete_entities method."""
# Mock data
to_be_deleted = [{"source": "alice", "relationship": "knows", "destination": "bob"}]
# Mock the _delete_entities_cypher method
mock_cypher = "MATCH (n) RETURN n"
mock_params = {"source_name": "alice", "dest_name": "bob", "user_id": self.user_id}
self.memory_graph._delete_entities_cypher = MagicMock(return_value=(mock_cypher, mock_params))
# Mock the graph.query result
mock_query_result = [{"source": "alice", "relationship": "knows", "target": "bob"}]
self.mock_graph.query.return_value = mock_query_result
# Call the _delete_entities method
result = self.memory_graph._delete_entities(to_be_deleted, self.user_id)
# Verify the method calls
self.memory_graph._delete_entities_cypher.assert_called_once_with("alice", "bob", "knows", self.user_id)
self.mock_graph.query.assert_called_once_with(mock_cypher, params=mock_params)
# Check the result
self.assertEqual(result, [mock_query_result])
if __name__ == "__main__":
unittest.main()
+65 -12
View File
@@ -1,7 +1,7 @@
import unittest
from unittest.mock import MagicMock, patch
import pytest
from mem0.graphs.neptune.main import MemoryGraph
from mem0.graphs.neptune.neptunedb import MemoryGraph
from mem0.graphs.neptune.base import NeptuneBase
@@ -13,24 +13,26 @@ class TestNeptuneMemory(unittest.TestCase):
# Create a mock config
self.config = MagicMock()
self.config.graph_store.config.endpoint = "neptune-graph://test-graph"
self.config.graph_store.config.endpoint = "neptune-db://test-graph"
self.config.graph_store.config.base_label = True
self.config.llm.provider = "openai_structured"
self.config.graph_store.llm = None
self.config.graph_store.custom_prompt = None
self.config.vector_store.provider = "qdrant"
self.config.vector_store.config = MagicMock()
# Create mock for NeptuneAnalyticsGraph
# Create mock for NeptuneGraph
self.mock_graph = MagicMock()
self.mock_graph.client.get_graph.return_value = {"status": "AVAILABLE"}
# Create mocks for static methods
self.mock_embedding_model = MagicMock()
self.mock_llm = MagicMock()
self.mock_vector_store = MagicMock()
# Patch the necessary components
self.neptune_analytics_graph_patcher = patch("mem0.graphs.neptune.main.NeptuneAnalyticsGraph")
self.mock_neptune_analytics_graph = self.neptune_analytics_graph_patcher.start()
self.mock_neptune_analytics_graph.return_value = self.mock_graph
self.neptune_graph_patcher = patch("mem0.graphs.neptune.neptunedb.NeptuneGraph")
self.mock_neptune_graph = self.neptune_graph_patcher.start()
self.mock_neptune_graph.return_value = self.mock_graph
# Patch the static methods
self.create_embedding_model_patcher = patch.object(NeptuneBase, "_create_embedding_model")
@@ -41,6 +43,10 @@ class TestNeptuneMemory(unittest.TestCase):
self.mock_create_llm = self.create_llm_patcher.start()
self.mock_create_llm.return_value = self.mock_llm
self.create_vector_store_patcher = patch.object(NeptuneBase, "_create_vector_store")
self.mock_create_vector_store = self.create_vector_store_patcher.start()
self.mock_create_vector_store.return_value = self.mock_vector_store
# Create the MemoryGraph instance
self.memory_graph = MemoryGraph(self.config)
@@ -50,18 +56,65 @@ class TestNeptuneMemory(unittest.TestCase):
def tearDown(self):
"""Tear down test fixtures after each test method."""
self.neptune_analytics_graph_patcher.stop()
self.neptune_graph_patcher.stop()
self.create_embedding_model_patcher.stop()
self.create_llm_patcher.stop()
self.create_vector_store_patcher.stop()
def test_initialization(self):
"""Test that the MemoryGraph is initialized correctly."""
self.assertEqual(self.memory_graph.graph, self.mock_graph)
self.assertEqual(self.memory_graph.embedding_model, self.mock_embedding_model)
self.assertEqual(self.memory_graph.llm, self.mock_llm)
self.assertEqual(self.memory_graph.vector_store, self.mock_vector_store)
self.assertEqual(self.memory_graph.llm_provider, "openai_structured")
self.assertEqual(self.memory_graph.node_label, ":`__Entity__`")
self.assertEqual(self.memory_graph.threshold, 0.7)
self.assertEqual(self.memory_graph.vector_store_limit, 5)
def test_collection_name_variants(self):
"""Test all collection_name configuration variants."""
# Test 1: graph_store.config.collection_name is set
config1 = MagicMock()
config1.graph_store.config.endpoint = "neptune-db://test-graph"
config1.graph_store.config.base_label = True
config1.graph_store.config.collection_name = "custom_collection"
config1.llm.provider = "openai"
config1.graph_store.llm = None
config1.vector_store.provider = "qdrant"
config1.vector_store.config = MagicMock()
MemoryGraph(config1)
self.assertEqual(config1.vector_store.config.collection_name, "custom_collection")
# Test 2: vector_store.config.collection_name exists, graph_store.config.collection_name is None
config2 = MagicMock()
config2.graph_store.config.endpoint = "neptune-db://test-graph"
config2.graph_store.config.base_label = True
config2.graph_store.config.collection_name = None
config2.llm.provider = "openai"
config2.graph_store.llm = None
config2.vector_store.provider = "qdrant"
config2.vector_store.config = MagicMock()
config2.vector_store.config.collection_name = "existing_collection"
MemoryGraph(config2)
self.assertEqual(config2.vector_store.config.collection_name, "existing_collection_neptune_vector_store")
# Test 3: Neither collection_name is set (default case)
config3 = MagicMock()
config3.graph_store.config.endpoint = "neptune-db://test-graph"
config3.graph_store.config.base_label = True
config3.graph_store.config.collection_name = None
config3.llm.provider = "openai"
config3.graph_store.llm = None
config3.vector_store.provider = "qdrant"
config3.vector_store.config = MagicMock()
config3.vector_store.config.collection_name = None
MemoryGraph(config3)
self.assertEqual(config3.vector_store.config.collection_name, "mem0_neptune_vector_store")
def test_init(self):
"""Test the class init functions"""
@@ -74,12 +127,12 @@ class TestNeptuneMemory(unittest.TestCase):
with pytest.raises(ValueError):
MemoryGraph(config_no_endpoint)
# Create a mock config with bad endpoint
config_ndb_endpoint = MagicMock()
config_ndb_endpoint.graph_store.config.endpoint = "neptune-db://test-graph"
# Create a mock config with wrong endpoint type
config_wrong_endpoint = MagicMock()
config_wrong_endpoint.graph_store.config.endpoint = "neptune-graph://test-graph"
with pytest.raises(ValueError):
MemoryGraph(config_ndb_endpoint)
MemoryGraph(config_wrong_endpoint)
def test_add_method(self):
"""Test the add method with mocked components."""