diff --git a/docs/open-source/graph_memory/overview.mdx b/docs/open-source/graph_memory/overview.mdx index e7ecd10b4..26b5c8293 100644 --- a/docs/open-source/graph_memory/overview.mdx +++ b/docs/open-source/graph_memory/overview.mdx @@ -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: ```python Python @@ -279,13 +271,61 @@ m = Memory.from_config(config_dict=config) ``` -#### 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: + + +```python Python +from mem0 import Memory + +config = { + "graph_store": { + "provider": "neptunedb", + "config": { + "collection_name": "", + "endpoint": "neptune-graph://", + }, + }, +} + +m = Memory.from_config(config_dict=config) +``` + + +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 diff --git a/examples/graph-db-demo/neptune-db-example.ipynb b/examples/graph-db-demo/neptune-db-example.ipynb new file mode 100644 index 000000000..e3839f531 --- /dev/null +++ b/examples/graph-db-demo/neptune-db-example.ipynb @@ -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 +} diff --git a/mem0/graphs/configs.py b/mem0/graphs/configs.py index 0b5e1e168..a79d3f31a 100644 --- a/mem0/graphs/configs.py +++ b/mem0/graphs/configs.py @@ -46,18 +46,20 @@ class NeptuneConfig(BaseModel): endpoint: Optional[str] = ( Field( None, - description="Endpoint to connect to a Neptune Analytics Server as neptune-graph://", + description="Endpoint to connect to a Neptune-DB Cluster as 'neptune-db://' or Neptune Analytics Server as 'neptune-graph://'", ), ) 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://'.") + raise ValueError("Please provide 'endpoint' with the format as 'neptune-db://' or 'neptune-graph://'.") 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()) diff --git a/mem0/graphs/neptune/base.py b/mem0/graphs/neptune/base.py index d427d2945..552220a73 100644 --- a/mem0/graphs/neptune/base.py +++ b/mem0/graphs/neptune/base.py @@ -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): diff --git a/mem0/graphs/neptune/main.py b/mem0/graphs/neptune/main.py deleted file mode 100644 index c47280b65..000000000 --- a/mem0/graphs/neptune/main.py +++ /dev/null @@ -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 diff --git a/mem0/graphs/neptune/neptunedb.py b/mem0/graphs/neptune/neptunedb.py new file mode 100644 index 000000000..d0ca68faf --- /dev/null +++ b/mem0/graphs/neptune/neptunedb.py @@ -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 diff --git a/mem0/graphs/neptune/neptunegraph.py b/mem0/graphs/neptune/neptunegraph.py new file mode 100644 index 000000000..c9264485d --- /dev/null +++ b/mem0/graphs/neptune/neptunegraph.py @@ -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 diff --git a/mem0/memory/main.py b/mem0/memory/main.py index 2419c47df..f4b2f6a3c 100644 --- a/mem0/memory/main.py +++ b/mem0/memory/main.py @@ -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") diff --git a/mem0/utils/factory.py b/mem0/utils/factory.py index 3ac56c824..c70c1bee1 100644 --- a/mem0/utils/factory.py +++ b/mem0/utils/factory.py @@ -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", } diff --git a/tests/memory/test_neptune_analytics_memory.py b/tests/memory/test_neptune_analytics_memory.py new file mode 100644 index 000000000..c69f9c115 --- /dev/null +++ b/tests/memory/test_neptune_analytics_memory.py @@ -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() diff --git a/tests/memory/test_neptune_memory.py b/tests/memory/test_neptune_memory.py index a269e003a..e690ac721 100644 --- a/tests/memory/test_neptune_memory.py +++ b/tests/memory/test_neptune_memory.py @@ -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."""