Feat: Mem0 vector store backend integration for Neptune Analytics (#3453)
Signed-off-by: Andy Kwok <andy.kwok@improving.com>
This commit is contained in:
@@ -0,0 +1,42 @@
|
||||
# Neptune Analytics Vector Store
|
||||
|
||||
[Neptune Analytics](https://docs.aws.amazon.com/neptune-analytics/latest/userguide/what-is-neptune-analytics.html/) is a memory-optimized graph database engine for analytics. With Neptune Analytics, you can get insights and find trends by processing large amounts of graph data in seconds, including vector search.
|
||||
|
||||
|
||||
## Installation
|
||||
|
||||
```bash
|
||||
pip install mem0ai[vector_stores]
|
||||
```
|
||||
|
||||
## Usage
|
||||
|
||||
```python
|
||||
config = {
|
||||
"vector_store": {
|
||||
"provider": "neptune",
|
||||
"config": {
|
||||
"collection_name": "mem0",
|
||||
"endpoint": f"neptune-graph://my-graph-identifier",
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
m = Memory.from_config(config)
|
||||
messages = [
|
||||
{"role": "user", "content": "I'm planning to watch a movie tonight. Any recommendations?"},
|
||||
{"role": "assistant", "content": "How about a thriller movies? They can be quite engaging."},
|
||||
{"role": "user", "content": "I'm not a big fan of thriller movies but I love sci-fi movies."},
|
||||
{"role": "assistant", "content": "Got it! I'll avoid thriller recommendations and suggest sci-fi movies in the future."}
|
||||
]
|
||||
m.add(messages, user_id="alice", metadata={"category": "movies"})
|
||||
```
|
||||
|
||||
## Parameters
|
||||
|
||||
Let's see the available parameters for the `neptune` config:
|
||||
|
||||
| Parameter | Description | Default Value |
|
||||
| --- | --- | --- |
|
||||
| `collection_name` | The name of the collection to store the vectors | `mem0` |
|
||||
| `endpoint` | Connection URL for the Neptune Analytics service | `neptune-graph://my-graph-identifier` |
|
||||
+3
-1
@@ -164,7 +164,8 @@
|
||||
"components/vectordbs/dbs/langchain",
|
||||
"components/vectordbs/dbs/baidu",
|
||||
"components/vectordbs/dbs/s3_vectors",
|
||||
"components/vectordbs/dbs/databricks"
|
||||
"components/vectordbs/dbs/databricks",
|
||||
"components/vectordbs/dbs/neptune_analytics"
|
||||
]
|
||||
}
|
||||
]
|
||||
@@ -223,6 +224,7 @@
|
||||
"pages": [
|
||||
"examples",
|
||||
"examples/aws_example",
|
||||
"examples/aws_neptune_analytics_hybrid_store",
|
||||
"examples/mem0-demo",
|
||||
"examples/ai_companion_js",
|
||||
"examples/collaborative-task-agent",
|
||||
|
||||
@@ -0,0 +1,120 @@
|
||||
---
|
||||
title: "Amazon Stack - Neptune Analytics Hybrid Store: AWS Bedrock and Neptune Analytics"
|
||||
---
|
||||
|
||||
This example demonstrates how to configure and use the `mem0ai` SDK with **AWS Bedrock** and **AWS Neptune Analytics** for persistent memory capabilities in Python.
|
||||
|
||||
## Installation
|
||||
|
||||
Install the required dependencies to include the Amazon data stack, including **boto3** and **langchain-aws**:
|
||||
|
||||
```bash
|
||||
pip install "mem0ai[graph,extras]"
|
||||
```
|
||||
|
||||
## Environment Setup
|
||||
|
||||
Set your AWS environment variables:
|
||||
|
||||
```python
|
||||
import os
|
||||
|
||||
# Set these in your environment or notebook
|
||||
os.environ['AWS_REGION'] = 'us-west-2'
|
||||
os.environ['AWS_ACCESS_KEY_ID'] = 'AK00000000000000000'
|
||||
os.environ['AWS_SECRET_ACCESS_KEY'] = 'AS00000000000000000'
|
||||
|
||||
# Confirm they are set
|
||||
print(os.environ['AWS_REGION'])
|
||||
print(os.environ['AWS_ACCESS_KEY_ID'])
|
||||
print(os.environ['AWS_SECRET_ACCESS_KEY'])
|
||||
```
|
||||
|
||||
## Configuration and Usage
|
||||
|
||||
This sets up Mem0 with:
|
||||
- [AWS Bedrock for LLM](https://docs.mem0.ai/components/llms/models/aws_bedrock)
|
||||
- [AWS Bedrock for embeddings](https://docs.mem0.ai/components/embedders/models/aws_bedrock#aws-bedrock)
|
||||
- [Neptune Analytics as the vector store](https://docs.mem0.ai/components/vectordbs/dbs/neptune_analytics)
|
||||
- [Neptune Analytics as the graph store](https://docs.mem0.ai/open-source/graph_memory/overview#initialize-neptune-analytics).
|
||||
|
||||
```python
|
||||
import boto3
|
||||
from mem0.memory.main import Memory
|
||||
|
||||
region = 'us-west-2'
|
||||
neptune_analytics_endpoint = 'neptune-graph://my-graph-identifier'
|
||||
|
||||
config = {
|
||||
"embedder": {
|
||||
"provider": "aws_bedrock",
|
||||
"config": {
|
||||
"model": "amazon.titan-embed-text-v2:0"
|
||||
}
|
||||
},
|
||||
"llm": {
|
||||
"provider": "aws_bedrock",
|
||||
"config": {
|
||||
"model": "us.anthropic.claude-3-7-sonnet-20250219-v1:0",
|
||||
"temperature": 0.1,
|
||||
"max_tokens": 2000
|
||||
}
|
||||
},
|
||||
"vector_store": {
|
||||
"provider": "neptune",
|
||||
"config": {
|
||||
"collection_name": "mem0",
|
||||
"endpoint": neptune_analytics_endpoint,
|
||||
},
|
||||
},
|
||||
"graph_store": {
|
||||
"provider": "neptune",
|
||||
"config": {
|
||||
"endpoint": neptune_analytics_endpoint,
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
# Initialize the memory system
|
||||
m = Memory.from_config(config)
|
||||
```
|
||||
|
||||
## Usage
|
||||
|
||||
Reference [Notebook example](https://github.com/mem0ai/mem0/blob/main/examples/graph-db-demo/neptune-example.ipynb)
|
||||
|
||||
#### Add a memory:
|
||||
|
||||
```python
|
||||
messages = [
|
||||
{"role": "user", "content": "I'm planning to watch a movie tonight. Any recommendations?"},
|
||||
{"role": "assistant", "content": "How about a thriller movies? They can be quite engaging."},
|
||||
{"role": "user", "content": "I'm not a big fan of thriller movies but I love sci-fi movies."},
|
||||
{"role": "assistant", "content": "Got it! I'll avoid thriller recommendations and suggest sci-fi movies in the future."}
|
||||
]
|
||||
|
||||
# Store inferred memories (default behavior)
|
||||
result = m.add(messages, user_id="alice", metadata={"category": "movie_recommendations"})
|
||||
```
|
||||
|
||||
#### Search a memory:
|
||||
```python
|
||||
relevant_memories = m.search(query, user_id="alice")
|
||||
```
|
||||
|
||||
#### Get all memories:
|
||||
```python
|
||||
all_memories = m.get_all(user_id="alice")
|
||||
```
|
||||
|
||||
#### Get a specific memory:
|
||||
```python
|
||||
memory = m.get(memory_id)
|
||||
```
|
||||
|
||||
|
||||
---
|
||||
|
||||
## Conclusion
|
||||
|
||||
With Mem0 and AWS services like Bedrock and Neptune Analytics, you can build intelligent AI companions that remember, adapt, and personalize their responses over time. This makes them ideal for long-term assistants, tutors, or support bots with persistent memory and natural conversation abilities.
|
||||
@@ -31,10 +31,12 @@
|
||||
"\n",
|
||||
"### 2. Connect to Amazon services\n",
|
||||
"\n",
|
||||
"For this sample notebook, configure `mem0ai` with [Amazon Neptune Analytics](https://docs.aws.amazon.com/neptune-analytics/latest/userguide/what-is-neptune-analytics.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",
|
||||
"For this sample notebook, configure `mem0ai` with [Amazon Neptune Analytics](https://docs.aws.amazon.com/neptune-analytics/latest/userguide/what-is-neptune-analytics.html) as the vector and graph store, and [Amazon Bedrock](https://docs.aws.amazon.com/bedrock/latest/userguide/what-is-bedrock.html) for generating embeddings.\n",
|
||||
"\n",
|
||||
"Use the following guide for setup details: [Setup AWS Bedrock, AOSS, and Neptune](https://docs.mem0.ai/examples/aws_example#aws-bedrock-and-aoss)\n",
|
||||
"\n",
|
||||
"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\n",
|
||||
"\n",
|
||||
"Your configuration should look similar to:\n",
|
||||
"\n",
|
||||
"```python\n",
|
||||
@@ -42,7 +44,8 @@
|
||||
" \"embedder\": {\n",
|
||||
" \"provider\": \"aws_bedrock\",\n",
|
||||
" \"config\": {\n",
|
||||
" \"model\": \"amazon.titan-embed-text-v2:0\"\n",
|
||||
" \"model\": \"amazon.titan-embed-text-v2:0\",\n",
|
||||
" \"embedding_dims\": 1024\n",
|
||||
" }\n",
|
||||
" },\n",
|
||||
" \"llm\": {\n",
|
||||
@@ -54,18 +57,10 @@
|
||||
" }\n",
|
||||
" },\n",
|
||||
" \"vector_store\": {\n",
|
||||
" \"provider\": \"opensearch\",\n",
|
||||
" \"provider\": \"neptune\",\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",
|
||||
" \"endpoint\": f\"neptune-graph://my-graph-identifier\",\n",
|
||||
" },\n",
|
||||
" },\n",
|
||||
" \"graph_store\": {\n",
|
||||
" \"provider\": \"neptune\",\n",
|
||||
@@ -96,14 +91,12 @@
|
||||
"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.main\").setLevel(logging.DEBUG)\n",
|
||||
"logging.getLogger(\"mem0.graphs.neptune.base\").setLevel(logging.DEBUG)\n",
|
||||
"logging.getLogger(\"mem0.graphs.neptune.main\").setLevel(logging.INFO)\n",
|
||||
"logging.getLogger(\"mem0.graphs.neptune.base\").setLevel(logging.INFO)\n",
|
||||
"logger = logging.getLogger(__name__)\n",
|
||||
"logger.setLevel(logging.DEBUG)\n",
|
||||
"\n",
|
||||
@@ -120,8 +113,7 @@
|
||||
"source": [
|
||||
"Setup the Mem0 configuration using:\n",
|
||||
"- Amazon Bedrock as the embedder\n",
|
||||
"- Amazon Neptune Analytics instance as a graph store\n",
|
||||
"- OpenSearch as the vector store"
|
||||
"- Amazon Neptune Analytics instance as a vector / graph store"
|
||||
]
|
||||
},
|
||||
{
|
||||
@@ -136,18 +128,12 @@
|
||||
"\n",
|
||||
"graph_identifier = os.environ.get(\"GRAPH_ID\")\n",
|
||||
"\n",
|
||||
"opensearch_host = os.environ.get(\"OS_HOST\")\n",
|
||||
"opensearch_post = os.environ.get(\"OS_PORT\")\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",
|
||||
" \"embedding_dims\": embedding_model_dims\n",
|
||||
" }\n",
|
||||
" },\n",
|
||||
" \"llm\": {\n",
|
||||
@@ -159,16 +145,9 @@
|
||||
" }\n",
|
||||
" },\n",
|
||||
" \"vector_store\": {\n",
|
||||
" \"provider\": \"opensearch\",\n",
|
||||
" \"provider\": \"neptune\",\n",
|
||||
" \"config\": {\n",
|
||||
" \"collection_name\": \"mem0ai_vector_store\",\n",
|
||||
" \"host\": opensearch_host,\n",
|
||||
" \"port\": opensearch_post,\n",
|
||||
" \"http_auth\": auth,\n",
|
||||
" \"embedding_model_dims\": embedding_model_dims,\n",
|
||||
" \"use_ssl\": True,\n",
|
||||
" \"verify_certs\": True,\n",
|
||||
" \"connection_class\": RequestsHttpConnection,\n",
|
||||
" \"endpoint\": f\"neptune-graph://{graph_identifier}\",\n",
|
||||
" },\n",
|
||||
" },\n",
|
||||
" \"graph_store\": {\n",
|
||||
@@ -431,13 +410,13 @@
|
||||
"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 Analytics can store the text chunks in a graph format with relationship entities."
|
||||
"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. Neptune Analytics can store the text chunks in a graph format with relationship entities."
|
||||
]
|
||||
}
|
||||
],
|
||||
"metadata": {
|
||||
"kernelspec": {
|
||||
"display_name": ".venv",
|
||||
"display_name": "Python 3 (ipykernel)",
|
||||
"language": "python",
|
||||
"name": "python3"
|
||||
},
|
||||
@@ -451,9 +430,9 @@
|
||||
"name": "python",
|
||||
"nbconvert_exporter": "python",
|
||||
"pygments_lexer": "ipython3",
|
||||
"version": "3.13.2"
|
||||
"version": "3.13.5"
|
||||
}
|
||||
},
|
||||
"nbformat": 4,
|
||||
"nbformat_minor": 2
|
||||
"nbformat_minor": 4
|
||||
}
|
||||
|
||||
@@ -0,0 +1,27 @@
|
||||
"""
|
||||
Configuration for Amazon Neptune Analytics vector store.
|
||||
|
||||
This module provides configuration settings for integrating with Amazon Neptune Analytics
|
||||
as a vector store backend for Mem0's memory layer.
|
||||
"""
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
|
||||
class NeptuneAnalyticsConfig(BaseModel):
|
||||
"""
|
||||
Configuration class for Amazon Neptune Analytics vector store.
|
||||
|
||||
Amazon Neptune Analytics is a graph analytics engine that can be used as a vector store
|
||||
for storing and retrieving memory embeddings in Mem0.
|
||||
|
||||
Attributes:
|
||||
collection_name (str): Name of the collection to store vectors. Defaults to "mem0".
|
||||
endpoint (str): Neptune Analytics graph endpoint URL or Graph ID for the runtime.
|
||||
"""
|
||||
collection_name: str = Field("mem0", description="Default name for the collection")
|
||||
endpoint: str = Field("endpoint", description="Graph ID for the runtime")
|
||||
|
||||
model_config = {
|
||||
"arbitrary_types_allowed": False,
|
||||
}
|
||||
@@ -176,6 +176,7 @@ class VectorStoreFactory:
|
||||
"langchain": "mem0.vector_stores.langchain.Langchain",
|
||||
"s3_vectors": "mem0.vector_stores.s3_vectors.S3Vectors",
|
||||
"baidu": "mem0.vector_stores.baidu.BaiduDB",
|
||||
"neptune": "mem0.vector_stores.neptune_analytics.NeptuneAnalyticsVector",
|
||||
}
|
||||
|
||||
@classmethod
|
||||
|
||||
@@ -18,6 +18,7 @@ class VectorStoreConfig(BaseModel):
|
||||
"mongodb": "MongoDBConfig",
|
||||
"milvus": "MilvusDBConfig",
|
||||
"baidu": "BaiduDBConfig",
|
||||
"neptune": "NeptuneAnalyticsConfig",
|
||||
"upstash_vector": "UpstashVectorConfig",
|
||||
"azure_ai_search": "AzureAISearchConfig",
|
||||
"redis": "RedisDBConfig",
|
||||
|
||||
@@ -0,0 +1,467 @@
|
||||
import logging
|
||||
import time
|
||||
import uuid
|
||||
from typing import Dict, List, Optional
|
||||
|
||||
from pydantic import BaseModel
|
||||
|
||||
try:
|
||||
from langchain_aws import NeptuneAnalyticsGraph
|
||||
except ImportError:
|
||||
raise ImportError("langchain_aws is not installed. Please install it using pip install langchain_aws")
|
||||
|
||||
from mem0.vector_stores.base import VectorStoreBase
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
class OutputData(BaseModel):
|
||||
id: Optional[str] # memory id
|
||||
score: Optional[float] # distance
|
||||
payload: Optional[Dict] # metadata
|
||||
|
||||
|
||||
class NeptuneAnalyticsVector(VectorStoreBase):
|
||||
"""
|
||||
Neptune Analytics vector store implementation for Mem0.
|
||||
|
||||
Provides vector storage and similarity search capabilities using Amazon Neptune Analytics,
|
||||
a serverless graph analytics service that supports vector operations.
|
||||
"""
|
||||
|
||||
_COLLECTION_PREFIX = "MEM0_VECTOR_"
|
||||
_FIELD_N = 'n'
|
||||
_FIELD_ID = '~id'
|
||||
_FIELD_PROP = '~properties'
|
||||
_FIELD_SCORE = 'score'
|
||||
_FIELD_LABEL = 'label'
|
||||
_TIMEZONE = "UTC"
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
endpoint: str,
|
||||
collection_name: str,
|
||||
):
|
||||
"""
|
||||
Initialize the Neptune Analytics vector store.
|
||||
|
||||
Args:
|
||||
endpoint (str): Neptune Analytics endpoint in format 'neptune-graph://<graphid>'.
|
||||
collection_name (str): Name of the collection to store vectors.
|
||||
|
||||
Raises:
|
||||
ValueError: If endpoint format is invalid.
|
||||
ImportError: If langchain_aws is not installed.
|
||||
"""
|
||||
|
||||
if not endpoint.startswith("neptune-graph://"):
|
||||
raise ValueError("Please provide 'endpoint' with the format as 'neptune-graph://<graphid>'.")
|
||||
|
||||
graph_id = endpoint.replace("neptune-graph://", "")
|
||||
self.graph = NeptuneAnalyticsGraph(graph_id)
|
||||
self.collection_name = self._COLLECTION_PREFIX + collection_name
|
||||
|
||||
|
||||
def create_col(self, name, vector_size, distance):
|
||||
"""
|
||||
Create a collection (no-op for Neptune Analytics).
|
||||
|
||||
Neptune Analytics supports dynamic indices that are created implicitly
|
||||
when vectors are inserted, so this method performs no operation.
|
||||
|
||||
Args:
|
||||
name: Collection name (unused).
|
||||
vector_size: Vector dimension (unused).
|
||||
distance: Distance metric (unused).
|
||||
"""
|
||||
pass
|
||||
|
||||
|
||||
def insert(self, vectors: List[list],
|
||||
payloads: Optional[List[Dict]] = None,
|
||||
ids: Optional[List[str]] = None):
|
||||
"""
|
||||
Insert vectors into the collection.
|
||||
|
||||
Creates or updates nodes in Neptune Analytics with vector embeddings and metadata.
|
||||
Uses MERGE operation to handle both creation and updates.
|
||||
|
||||
Args:
|
||||
vectors (List[list]): List of embedding vectors to insert.
|
||||
payloads (Optional[List[Dict]]): Optional metadata for each vector.
|
||||
ids (Optional[List[str]]): Optional IDs for vectors. Generated if not provided.
|
||||
"""
|
||||
|
||||
para_list = []
|
||||
for index, data_vector in enumerate(vectors):
|
||||
if payloads:
|
||||
payload = payloads[index]
|
||||
payload[self._FIELD_LABEL] = self.collection_name
|
||||
payload["updated_at"] = str(int(time.time()))
|
||||
else:
|
||||
payload = {}
|
||||
para_list.append(dict(
|
||||
node_id=ids[index] if ids else str(uuid.uuid4()),
|
||||
properties=payload,
|
||||
embedding=data_vector,
|
||||
))
|
||||
|
||||
para_map_to_insert = {"rows": para_list}
|
||||
|
||||
query_string = (f"""
|
||||
UNWIND $rows AS row
|
||||
MERGE (n :{self.collection_name} {{`~id`: row.node_id}})
|
||||
ON CREATE SET n = row.properties
|
||||
ON MATCH SET n += row.properties
|
||||
"""
|
||||
)
|
||||
self.execute_query(query_string, para_map_to_insert)
|
||||
|
||||
|
||||
query_string_vector = (f"""
|
||||
UNWIND $rows AS row
|
||||
MATCH (n
|
||||
:{self.collection_name}
|
||||
{{`~id`: row.node_id}})
|
||||
WITH n, row.embedding AS embedding
|
||||
CALL neptune.algo.vectors.upsert(n, embedding)
|
||||
YIELD success
|
||||
RETURN success
|
||||
"""
|
||||
)
|
||||
result = self.execute_query(query_string_vector, para_map_to_insert)
|
||||
self._process_success_message(result, "Vector store - Insert")
|
||||
|
||||
|
||||
def search(
|
||||
self, query: str, vectors: List[float], limit: int = 5, filters: Optional[Dict] = None
|
||||
) -> List[OutputData]:
|
||||
"""
|
||||
Search for similar vectors using embedding similarity.
|
||||
|
||||
Performs vector similarity search using Neptune Analytics' topKByEmbeddingWithFiltering
|
||||
algorithm to find the most similar vectors.
|
||||
|
||||
Args:
|
||||
query (str): Search query text (unused in vector search).
|
||||
vectors (List[float]): Query embedding vector.
|
||||
limit (int, optional): Maximum number of results to return. Defaults to 5.
|
||||
filters (Optional[Dict]): Optional filters to apply to search results.
|
||||
|
||||
Returns:
|
||||
List[OutputData]: List of similar vectors with scores and metadata.
|
||||
"""
|
||||
|
||||
if not filters:
|
||||
filters = {}
|
||||
filters[self._FIELD_LABEL] = self.collection_name
|
||||
|
||||
filter_clause = self._get_node_filter_clause(filters)
|
||||
|
||||
query_string = f"""
|
||||
CALL neptune.algo.vectors.topKByEmbeddingWithFiltering({{
|
||||
topK: {limit},
|
||||
embedding: {vectors}
|
||||
{filter_clause}
|
||||
}}
|
||||
)
|
||||
YIELD node, score
|
||||
RETURN node as n, score
|
||||
"""
|
||||
query_response = self.execute_query(query_string)
|
||||
if len(query_response) > 0:
|
||||
return self._parse_query_responses(query_response, with_score=True)
|
||||
else :
|
||||
return []
|
||||
|
||||
|
||||
def delete(self, vector_id: str):
|
||||
"""
|
||||
Delete a vector by its ID.
|
||||
|
||||
Removes the node and all its relationships from the Neptune Analytics graph.
|
||||
|
||||
Args:
|
||||
vector_id (str): ID of the vector to delete.
|
||||
"""
|
||||
params = dict(node_id=vector_id)
|
||||
query_string = f"""
|
||||
MATCH (n :{self.collection_name})
|
||||
WHERE id(n) = $node_id
|
||||
DETACH DELETE n
|
||||
"""
|
||||
self.execute_query(query_string, params)
|
||||
|
||||
def update(
|
||||
self,
|
||||
vector_id: str,
|
||||
vector: Optional[List[float]] = None,
|
||||
payload: Optional[Dict] = None,
|
||||
):
|
||||
"""
|
||||
Update a vector's embedding and/or metadata.
|
||||
|
||||
Updates the node properties and/or vector embedding for an existing vector.
|
||||
Can update either the payload, the vector, or both.
|
||||
|
||||
Args:
|
||||
vector_id (str): ID of the vector to update.
|
||||
vector (Optional[List[float]]): New embedding vector.
|
||||
payload (Optional[Dict]): New metadata to replace existing payload.
|
||||
"""
|
||||
|
||||
if payload:
|
||||
# Replace payload
|
||||
payload[self._FIELD_LABEL] = self.collection_name
|
||||
payload["updated_at"] = str(int(time.time()))
|
||||
para_payload = {
|
||||
"properties": payload,
|
||||
"vector_id": vector_id
|
||||
}
|
||||
query_string_embedding = f"""
|
||||
MATCH (n :{self.collection_name})
|
||||
WHERE id(n) = $vector_id
|
||||
SET n = $properties
|
||||
"""
|
||||
self.execute_query(query_string_embedding, para_payload)
|
||||
|
||||
if vector:
|
||||
para_embedding = {
|
||||
"embedding": vector,
|
||||
"vector_id": vector_id
|
||||
}
|
||||
query_string_embedding = f"""
|
||||
MATCH (n :{self.collection_name})
|
||||
WHERE id(n) = $vector_id
|
||||
WITH $embedding as embedding, n as n
|
||||
CALL neptune.algo.vectors.upsert(n, embedding)
|
||||
YIELD success
|
||||
RETURN success
|
||||
"""
|
||||
self.execute_query(query_string_embedding, para_embedding)
|
||||
|
||||
|
||||
|
||||
def get(self, vector_id: str):
|
||||
"""
|
||||
Retrieve a vector by its ID.
|
||||
|
||||
Fetches the node data including metadata for the specified vector ID.
|
||||
|
||||
Args:
|
||||
vector_id (str): ID of the vector to retrieve.
|
||||
|
||||
Returns:
|
||||
OutputData: Vector data with metadata, or None if not found.
|
||||
"""
|
||||
params = dict(node_id=vector_id)
|
||||
query_string = f"""
|
||||
MATCH (n :{self.collection_name})
|
||||
WHERE id(n) = $node_id
|
||||
RETURN n
|
||||
"""
|
||||
|
||||
# Composite the query
|
||||
result = self.execute_query(query_string, params)
|
||||
|
||||
if len(result) != 0:
|
||||
return self._parse_query_responses(result)[0]
|
||||
|
||||
|
||||
def list_cols(self):
|
||||
"""
|
||||
List all collections with the Mem0 prefix.
|
||||
|
||||
Queries the Neptune Analytics schema to find all node labels that start
|
||||
with the Mem0 collection prefix.
|
||||
|
||||
Returns:
|
||||
List[str]: List of collection names.
|
||||
"""
|
||||
query_string = f"""
|
||||
CALL neptune.graph.pg_schema()
|
||||
YIELD schema
|
||||
RETURN [ label IN schema.nodeLabels WHERE label STARTS WITH '{self.collection_name}'] AS result
|
||||
"""
|
||||
result = self.execute_query(query_string)
|
||||
if len(result) == 1 and "result" in result[0]:
|
||||
return result[0]["result"]
|
||||
else:
|
||||
return []
|
||||
|
||||
|
||||
def delete_col(self):
|
||||
"""
|
||||
Delete the entire collection.
|
||||
|
||||
Removes all nodes with the collection label and their relationships
|
||||
from the Neptune Analytics graph.
|
||||
"""
|
||||
self.execute_query(f"MATCH (n :{self.collection_name}) DETACH DELETE n")
|
||||
|
||||
|
||||
def col_info(self):
|
||||
"""
|
||||
Get collection information (no-op for Neptune Analytics).
|
||||
|
||||
Collections are created dynamically in Neptune Analytics, so no
|
||||
collection-specific metadata is available.
|
||||
"""
|
||||
pass
|
||||
|
||||
|
||||
def list(self, filters: Optional[Dict] = None, limit: int = 100) -> List[OutputData]:
|
||||
"""
|
||||
List all vectors in the collection with optional filtering.
|
||||
|
||||
Retrieves vectors from the collection, optionally filtered by metadata properties.
|
||||
|
||||
Args:
|
||||
filters (Optional[Dict]): Optional filters to apply based on metadata.
|
||||
limit (int, optional): Maximum number of vectors to return. Defaults to 100.
|
||||
|
||||
Returns:
|
||||
List[OutputData]: List of vectors with their metadata.
|
||||
"""
|
||||
where_clause = self._get_where_clause(filters) if filters else ""
|
||||
|
||||
para = {
|
||||
"limit": limit,
|
||||
}
|
||||
query_string = f"""
|
||||
MATCH (n :{self.collection_name})
|
||||
{where_clause}
|
||||
RETURN n
|
||||
LIMIT $limit
|
||||
"""
|
||||
query_response = self.execute_query(query_string, para)
|
||||
|
||||
if len(query_response) > 0:
|
||||
# Handle if there is no match.
|
||||
return [self._parse_query_responses(query_response)]
|
||||
return [[]]
|
||||
|
||||
|
||||
def reset(self):
|
||||
"""
|
||||
Reset the collection by deleting all vectors.
|
||||
|
||||
Removes all vectors from the collection, effectively resetting it to empty state.
|
||||
"""
|
||||
self.delete_col()
|
||||
|
||||
|
||||
def _parse_query_responses(self, response: dict, with_score: bool = False):
|
||||
"""
|
||||
Parse Neptune Analytics query responses into OutputData objects.
|
||||
|
||||
Args:
|
||||
response (dict): Raw query response from Neptune Analytics.
|
||||
with_score (bool, optional): Whether to include similarity scores. Defaults to False.
|
||||
|
||||
Returns:
|
||||
List[OutputData]: Parsed response data.
|
||||
"""
|
||||
result = []
|
||||
# Handle if there is no match.
|
||||
for item in response:
|
||||
id = item[self._FIELD_N][self._FIELD_ID]
|
||||
properties = item[self._FIELD_N][self._FIELD_PROP]
|
||||
properties.pop("label", None)
|
||||
if with_score:
|
||||
score = item[self._FIELD_SCORE]
|
||||
else:
|
||||
score = None
|
||||
result.append(OutputData(
|
||||
id=id,
|
||||
score=score,
|
||||
payload=properties,
|
||||
))
|
||||
return result
|
||||
|
||||
|
||||
def execute_query(self, query_string: str, params=None):
|
||||
"""
|
||||
Execute an openCypher query on Neptune Analytics.
|
||||
|
||||
This is a wrapper method around the Neptune Analytics graph query execution
|
||||
that provides debug logging for query monitoring and troubleshooting.
|
||||
|
||||
Args:
|
||||
query_string (str): The openCypher query string to execute.
|
||||
params (dict): Parameters to bind to the query.
|
||||
|
||||
Returns:
|
||||
Query result from Neptune Analytics graph execution.
|
||||
"""
|
||||
if params is None:
|
||||
params = {}
|
||||
logger.debug(f"Executing openCypher query:[{query_string}], with parameters:[{params}].")
|
||||
return self.graph.query(query_string, params)
|
||||
|
||||
|
||||
@staticmethod
|
||||
def _get_where_clause(filters: dict):
|
||||
"""
|
||||
Build WHERE clause for Cypher queries from filters.
|
||||
|
||||
Args:
|
||||
filters (dict): Filter conditions as key-value pairs.
|
||||
|
||||
Returns:
|
||||
str: Formatted WHERE clause for Cypher query.
|
||||
"""
|
||||
where_clause = ""
|
||||
for i, (k, v) in enumerate(filters.items()):
|
||||
if i == 0:
|
||||
where_clause += f"WHERE n.{k} = '{v}' "
|
||||
else:
|
||||
where_clause += f"AND n.{k} = '{v}' "
|
||||
return where_clause
|
||||
|
||||
@staticmethod
|
||||
def _get_node_filter_clause(filters: dict):
|
||||
"""
|
||||
Build node filter clause for vector search operations.
|
||||
|
||||
Creates filter conditions for Neptune Analytics vector search operations
|
||||
using the nodeFilter parameter format.
|
||||
|
||||
Args:
|
||||
filters (dict): Filter conditions as key-value pairs.
|
||||
|
||||
Returns:
|
||||
str: Formatted node filter clause for vector search.
|
||||
"""
|
||||
conditions = []
|
||||
for k, v in filters.items():
|
||||
conditions.append(f"{{equals:{{property: '{k}', value: '{v}'}}}}")
|
||||
|
||||
if len(conditions) == 1:
|
||||
filter_clause = f", nodeFilter: {conditions[0]}"
|
||||
else:
|
||||
filter_clause = f"""
|
||||
, nodeFilter: {{andAll: [ {", ".join(conditions)} ]}}
|
||||
"""
|
||||
return filter_clause
|
||||
|
||||
|
||||
@staticmethod
|
||||
def _process_success_message(response, context):
|
||||
"""
|
||||
Process and validate success messages from Neptune Analytics operations.
|
||||
|
||||
Checks the response from vector operations (insert/update) to ensure they
|
||||
completed successfully. Logs errors if operations fail.
|
||||
|
||||
Args:
|
||||
response: Response from Neptune Analytics vector operation.
|
||||
context (str): Context description for logging (e.g., "Vector store - Insert").
|
||||
"""
|
||||
for success_message in response:
|
||||
if "success" not in success_message:
|
||||
logger.error(f"Query execution status is absent on action: [{context}]")
|
||||
break
|
||||
|
||||
if success_message["success"] is not True:
|
||||
logger.error(f"Abnormal response status on action: [{context}] with message: [{success_message['success']}] ")
|
||||
break
|
||||
@@ -49,6 +49,7 @@ vector_stores = [
|
||||
"redisvl>=0.1.0,<1.0.0",
|
||||
"elasticsearch>=8.0.0,<9.0.0",
|
||||
"pymilvus>=2.4.0,<2.6.0",
|
||||
"langchain-aws>=0.2.23",
|
||||
]
|
||||
llms = [
|
||||
"groq>=0.3.0",
|
||||
|
||||
@@ -0,0 +1,186 @@
|
||||
import logging
|
||||
import os
|
||||
import sys
|
||||
|
||||
import pytest
|
||||
from dotenv import load_dotenv
|
||||
|
||||
from mem0.utils.factory import VectorStoreFactory
|
||||
|
||||
load_dotenv()
|
||||
|
||||
# Configure logging
|
||||
logging.getLogger("mem0.vector.neptune.main").setLevel(logging.INFO)
|
||||
logging.getLogger("mem0.vector.neptune.base").setLevel(logging.INFO)
|
||||
logger = logging.getLogger(__name__)
|
||||
logger.setLevel(logging.DEBUG)
|
||||
|
||||
logging.basicConfig(
|
||||
format="%(levelname)s - %(message)s",
|
||||
datefmt="%Y-%m-%d %H:%M:%S",
|
||||
stream=sys.stdout,
|
||||
)
|
||||
|
||||
# Test constants
|
||||
EMBEDDING_MODEL_DIMS = 1024
|
||||
VECTOR_1 = [-0.1] * EMBEDDING_MODEL_DIMS
|
||||
VECTOR_2 = [-0.2] * EMBEDDING_MODEL_DIMS
|
||||
VECTOR_3 = [-0.3] * EMBEDDING_MODEL_DIMS
|
||||
|
||||
SAMPLE_PAYLOADS = [
|
||||
{"test_text": "text_value", "another_field": "field_2_value"},
|
||||
{"test_text": "text_value_BBBB"},
|
||||
{"test_text": "text_value_CCCC"}
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.skipif(not os.getenv("RUN_TEST_NEPTUNE_ANALYTICS"), reason="Only run with RUN_TEST_NEPTUNE_ANALYTICS is true")
|
||||
class TestNeptuneAnalyticsOperations:
|
||||
"""Test basic CRUD operations."""
|
||||
|
||||
@pytest.fixture
|
||||
def na_instance(self):
|
||||
"""Create Neptune Analytics vector store instance for testing."""
|
||||
config = {
|
||||
"endpoint": f"neptune-graph://{os.getenv('GRAPH_ID')}",
|
||||
"collection_name": "test",
|
||||
}
|
||||
return VectorStoreFactory.create("neptune", config)
|
||||
|
||||
|
||||
def test_insert_and_list(self, na_instance):
|
||||
"""Test vector insertion and listing."""
|
||||
na_instance.reset()
|
||||
na_instance.insert(
|
||||
vectors=[VECTOR_1, VECTOR_2, VECTOR_3],
|
||||
ids=["A", "B", "C"],
|
||||
payloads=SAMPLE_PAYLOADS
|
||||
)
|
||||
|
||||
list_result = na_instance.list()[0]
|
||||
assert len(list_result) == 3
|
||||
assert "label" not in list_result[0].payload
|
||||
|
||||
|
||||
def test_get(self, na_instance):
|
||||
"""Test retrieving a specific vector."""
|
||||
na_instance.reset()
|
||||
na_instance.insert(
|
||||
vectors=[VECTOR_1],
|
||||
ids=["A"],
|
||||
payloads=[SAMPLE_PAYLOADS[0]]
|
||||
)
|
||||
|
||||
vector_a = na_instance.get("A")
|
||||
assert vector_a.id == "A"
|
||||
assert vector_a.score is None
|
||||
assert vector_a.payload["test_text"] == "text_value"
|
||||
assert vector_a.payload["another_field"] == "field_2_value"
|
||||
assert "label" not in vector_a.payload
|
||||
|
||||
|
||||
def test_update(self, na_instance):
|
||||
"""Test updating vector payload."""
|
||||
na_instance.reset()
|
||||
na_instance.insert(
|
||||
vectors=[VECTOR_1],
|
||||
ids=["A"],
|
||||
payloads=[SAMPLE_PAYLOADS[0]]
|
||||
)
|
||||
|
||||
na_instance.update(vector_id="A", payload={"updated_payload_str": "update_str"})
|
||||
vector_a = na_instance.get("A")
|
||||
|
||||
assert vector_a.id == "A"
|
||||
assert vector_a.score is None
|
||||
assert vector_a.payload["updated_payload_str"] == "update_str"
|
||||
assert "label" not in vector_a.payload
|
||||
|
||||
|
||||
def test_delete(self, na_instance):
|
||||
"""Test deleting a specific vector."""
|
||||
na_instance.reset()
|
||||
na_instance.insert(
|
||||
vectors=[VECTOR_1],
|
||||
ids=["A"],
|
||||
payloads=[SAMPLE_PAYLOADS[0]]
|
||||
)
|
||||
|
||||
size_before = na_instance.list()[0]
|
||||
assert len(size_before) == 1
|
||||
|
||||
na_instance.delete("A")
|
||||
size_after = na_instance.list()[0]
|
||||
assert len(size_after) == 0
|
||||
|
||||
|
||||
def test_search(self, na_instance):
|
||||
"""Test vector similarity search."""
|
||||
na_instance.reset()
|
||||
na_instance.insert(
|
||||
vectors=[VECTOR_1, VECTOR_2, VECTOR_3],
|
||||
ids=["A", "B", "C"],
|
||||
payloads=SAMPLE_PAYLOADS
|
||||
)
|
||||
|
||||
result = na_instance.search(query="", vectors=VECTOR_1, limit=1)
|
||||
assert len(result) == 1
|
||||
assert "label" not in result[0].payload
|
||||
|
||||
|
||||
def test_reset(self, na_instance):
|
||||
"""Test resetting the collection."""
|
||||
na_instance.reset()
|
||||
na_instance.insert(
|
||||
vectors=[VECTOR_1, VECTOR_2, VECTOR_3],
|
||||
ids=["A", "B", "C"],
|
||||
payloads=SAMPLE_PAYLOADS
|
||||
)
|
||||
|
||||
list_result = na_instance.list()[0]
|
||||
assert len(list_result) == 3
|
||||
|
||||
na_instance.reset()
|
||||
list_result = na_instance.list()[0]
|
||||
assert len(list_result) == 0
|
||||
|
||||
|
||||
def test_delete_col(self, na_instance):
|
||||
"""Test deleting the entire collection."""
|
||||
na_instance.reset()
|
||||
na_instance.insert(
|
||||
vectors=[VECTOR_1, VECTOR_2, VECTOR_3],
|
||||
ids=["A", "B", "C"],
|
||||
payloads=SAMPLE_PAYLOADS
|
||||
)
|
||||
|
||||
list_result = na_instance.list()[0]
|
||||
assert len(list_result) == 3
|
||||
|
||||
na_instance.delete_col()
|
||||
list_result = na_instance.list()[0]
|
||||
assert len(list_result) == 0
|
||||
|
||||
|
||||
def test_list_cols(self, na_instance):
|
||||
"""Test listing collections."""
|
||||
na_instance.reset()
|
||||
na_instance.insert(
|
||||
vectors=[VECTOR_1, VECTOR_2, VECTOR_3],
|
||||
ids=["A", "B", "C"],
|
||||
payloads=SAMPLE_PAYLOADS
|
||||
)
|
||||
|
||||
result = na_instance.list_cols()
|
||||
assert result == ["MEM0_VECTOR_test"]
|
||||
|
||||
|
||||
def test_invalid_endpoint_format(self):
|
||||
"""Test that invalid endpoint format raises ValueError."""
|
||||
config = {
|
||||
"endpoint": f"xxx://{os.getenv('GRAPH_ID')}",
|
||||
"collection_name": "test",
|
||||
}
|
||||
|
||||
with pytest.raises(ValueError):
|
||||
VectorStoreFactory.create("neptune", config)
|
||||
Reference in New Issue
Block a user