Merge branch 'main' into user/parshva/mem0_1.0.0
This commit is contained in:
@@ -283,7 +283,7 @@ curl -X POST "https://api.mem0.ai/v1/memories/" \
|
||||
|
||||
### Search with Custom Filters
|
||||
|
||||
Our advanced search allows you to set custom search filters. You can filter by user_id, agent_id, app_id, run_id, created_at, updated_at, categories, and text. The filters support logical operators (AND, OR) and comparison operators (in, gte, lte, gt, lt, ne, contains, icontains, *). The wildcard character (*) matches everything for a specific field.
|
||||
Our advanced search allows you to set custom search filters. You can filter by user_id, agent_id, app_id, run_id, created_at, updated_at, categories, and text. The filters support logical operators (AND, OR) and comparison operators (in, gte, lte, gt, lt, ne, contains, icontains, `*`). The wildcard character (`*`) matches everything for a specific field.
|
||||
|
||||
Here you need to define `version` as `v2` in the search method.
|
||||
|
||||
@@ -691,7 +691,7 @@ curl -X GET "https://api.mem0.ai/v1/memories/?user_id=alex&keywords=to play&page
|
||||
|
||||
#### Get all memories using custom filters
|
||||
|
||||
Our advanced retrieval allows you to set custom filters when fetching memories. You can filter by user_id, agent_id, app_id, run_id, created_at, updated_at, categories, and keywords. The filters support logical operators (AND, OR) and comparison operators (in, gte, lte, gt, lt, ne, contains, icontains, *). The wildcard character (*) matches everything for a specific field.
|
||||
Our advanced retrieval allows you to set custom filters when fetching memories. You can filter by user_id, agent_id, app_id, run_id, created_at, updated_at, categories, and keywords. The filters support logical operators (AND, OR) and comparison operators (in, gte, lte, gt, lt, ne, contains, icontains, `*`). The wildcard character (`*`) matches everything for a specific field.
|
||||
|
||||
Here you need to define `version` as `v2` in the get_all method.
|
||||
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
import os
|
||||
from abc import ABC
|
||||
from typing import Dict, Optional, Union
|
||||
|
||||
@@ -38,7 +39,7 @@ class BaseEmbedderConfig(ABC):
|
||||
# AWS Bedrock specific
|
||||
aws_access_key_id: Optional[str] = None,
|
||||
aws_secret_access_key: Optional[str] = None,
|
||||
aws_region: Optional[str] = "us-west-2",
|
||||
aws_region: Optional[str] = None,
|
||||
):
|
||||
"""
|
||||
Initializes a configuration class instance for the Embeddings.
|
||||
@@ -105,4 +106,5 @@ class BaseEmbedderConfig(ABC):
|
||||
# AWS Bedrock specific
|
||||
self.aws_access_key_id = aws_access_key_id
|
||||
self.aws_secret_access_key = aws_secret_access_key
|
||||
self.aws_region = aws_region
|
||||
self.aws_region = aws_region or os.environ.get("AWS_REGION") or "us-west-2"
|
||||
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
from typing import Optional, Dict, Any, List
|
||||
from mem0.configs.llms.base import BaseLlmConfig
|
||||
import os
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
from mem0.configs.llms.base import BaseLlmConfig
|
||||
|
||||
|
||||
class AWSBedrockConfig(BaseLlmConfig):
|
||||
@@ -19,7 +20,7 @@ class AWSBedrockConfig(BaseLlmConfig):
|
||||
top_k: int = 1,
|
||||
aws_access_key_id: Optional[str] = None,
|
||||
aws_secret_access_key: Optional[str] = None,
|
||||
aws_region: str = "us-west-2",
|
||||
aws_region: str = "",
|
||||
aws_session_token: Optional[str] = None,
|
||||
aws_profile: Optional[str] = None,
|
||||
model_kwargs: Optional[Dict[str, Any]] = None,
|
||||
@@ -53,7 +54,7 @@ class AWSBedrockConfig(BaseLlmConfig):
|
||||
|
||||
self.aws_access_key_id = aws_access_key_id
|
||||
self.aws_secret_access_key = aws_secret_access_key
|
||||
self.aws_region = aws_region
|
||||
self.aws_region = aws_region or os.getenv("AWS_REGION", "us-west-2")
|
||||
self.aws_session_token = aws_session_token
|
||||
self.aws_profile = aws_profile
|
||||
self.model_kwargs = model_kwargs or {}
|
||||
|
||||
+13
-1
@@ -293,14 +293,26 @@ def get_update_memory_messages(retrieved_old_memory_dict, response_content, cust
|
||||
global DEFAULT_UPDATE_MEMORY_PROMPT
|
||||
custom_update_memory_prompt = DEFAULT_UPDATE_MEMORY_PROMPT
|
||||
|
||||
return f"""{custom_update_memory_prompt}
|
||||
|
||||
if retrieved_old_memory_dict:
|
||||
current_memory_part = f"""
|
||||
Below is the current content of my memory which I have collected till now. You have to update it in the following format only:
|
||||
|
||||
```
|
||||
{retrieved_old_memory_dict}
|
||||
```
|
||||
|
||||
"""
|
||||
else:
|
||||
current_memory_part = """
|
||||
Current memory is empty.
|
||||
|
||||
"""
|
||||
|
||||
return f"""{custom_update_memory_prompt}
|
||||
|
||||
{current_memory_part}
|
||||
|
||||
The new retrieved facts are mentioned in the triple backticks. You have to analyze the new retrieved facts and determine whether these facts should be added, updated, or deleted in the memory.
|
||||
|
||||
```
|
||||
|
||||
@@ -11,19 +11,21 @@ class PGVectorConfig(BaseModel):
|
||||
password: Optional[str] = Field(None, description="Database password")
|
||||
host: Optional[str] = Field(None, description="Database host. Default is localhost")
|
||||
port: Optional[int] = Field(None, description="Database port. Default is 1536")
|
||||
diskann: Optional[bool] = Field(True, description="Use diskann for approximate nearest neighbors search")
|
||||
hnsw: Optional[bool] = Field(False, description="Use hnsw for faster search")
|
||||
diskann: Optional[bool] = Field(False, description="Use diskann for approximate nearest neighbors search")
|
||||
hnsw: Optional[bool] = Field(True, description="Use hnsw for faster search")
|
||||
minconn: Optional[int] = Field(1, description="Minimum number of connections in the pool")
|
||||
maxconn: Optional[int] = Field(5, description="Maximum number of connections in the pool")
|
||||
# New SSL and connection options
|
||||
sslmode: Optional[str] = Field(None, description="SSL mode for PostgreSQL connection (e.g., 'require', 'prefer', 'disable')")
|
||||
connection_string: Optional[str] = Field(None, description="PostgreSQL connection string (overrides individual connection parameters)")
|
||||
connection_pool: Optional[Any] = Field(None, description="psycopg2 connection pool object (overrides connection string and individual parameters)")
|
||||
connection_pool: Optional[Any] = Field(None, description="psycopg connection pool object (overrides connection string and individual parameters)")
|
||||
|
||||
@model_validator(mode="before")
|
||||
def check_auth_and_connection(cls, values):
|
||||
# If connection_pool is provided, skip validation of individual connection parameters
|
||||
if values.get("connection_pool") is not None:
|
||||
return values
|
||||
|
||||
|
||||
# If connection_string is provided, skip validation of individual connection parameters
|
||||
if values.get("connection_string") is not None:
|
||||
return values
|
||||
@@ -32,9 +34,9 @@ class PGVectorConfig(BaseModel):
|
||||
user, password = values.get("user"), values.get("password")
|
||||
host, port = values.get("host"), values.get("port")
|
||||
if not user and not password:
|
||||
raise ValueError("Both 'user' and 'password' must be provided when not using connection_string or connection_pool.")
|
||||
raise ValueError("Both 'user' and 'password' must be provided when not using connection_string.")
|
||||
if not host and not port:
|
||||
raise ValueError("Both 'host' and 'port' must be provided when not using connection_string or connection_pool.")
|
||||
raise ValueError("Both 'host' and 'port' must be provided when not using connection_string.")
|
||||
return values
|
||||
|
||||
@model_validator(mode="before")
|
||||
|
||||
@@ -28,15 +28,15 @@ class AWSBedrockEmbedding(EmbeddingBase):
|
||||
aws_access_key = os.environ.get("AWS_ACCESS_KEY_ID", "")
|
||||
aws_secret_key = os.environ.get("AWS_SECRET_ACCESS_KEY", "")
|
||||
aws_session_token = os.environ.get("AWS_SESSION_TOKEN", "")
|
||||
aws_region = os.environ.get("AWS_REGION", "us-west-2")
|
||||
|
||||
# Check if AWS config is provided in the config
|
||||
if hasattr(self.config, "aws_access_key_id"):
|
||||
aws_access_key = self.config.aws_access_key_id
|
||||
if hasattr(self.config, "aws_secret_access_key"):
|
||||
aws_secret_key = self.config.aws_secret_access_key
|
||||
if hasattr(self.config, "aws_region"):
|
||||
aws_region = self.config.aws_region
|
||||
|
||||
# AWS region is always set in config - see BaseEmbedderConfig
|
||||
aws_region = self.config.aws_region or "us-west-2"
|
||||
|
||||
self.client = boto3.client(
|
||||
"bedrock-runtime",
|
||||
|
||||
@@ -522,12 +522,12 @@ class MemoryGraph:
|
||||
WITH destination
|
||||
MERGE (source {source_label} {{{merge_props_str}}})
|
||||
ON CREATE SET
|
||||
source.created = current_timestamp(),
|
||||
source.mentions = 1
|
||||
source.embedding = CAST($source_embedding,'FLOAT[{self.embedding_dims}]')
|
||||
source.created = current_timestamp(),
|
||||
source.mentions = 1,
|
||||
source.embedding = CAST($source_embedding,'FLOAT[{self.embedding_dims}]')
|
||||
ON MATCH SET
|
||||
source.mentions = coalesce(source.mentions, 0) + 1
|
||||
source.embedding = CAST($source_embedding,'FLOAT[{self.embedding_dims}]')
|
||||
source.mentions = coalesce(source.mentions, 0) + 1,
|
||||
source.embedding = CAST($source_embedding,'FLOAT[{self.embedding_dims}]')
|
||||
WITH source, destination
|
||||
MERGE (source)-[r {relationship_label} {{name: $relationship_name}}]->(destination)
|
||||
ON CREATE SET
|
||||
|
||||
@@ -111,15 +111,13 @@ class MongoDB(VectorStoreBase):
|
||||
except PyMongoError as e:
|
||||
logger.error(f"Error inserting data: {e}")
|
||||
|
||||
def search(
|
||||
self, query: str, query_vector: List[float], limit=5, filters: Optional[Dict] = None
|
||||
) -> List[OutputData]:
|
||||
def search(self, query: str, vectors: List[float], limit=5, filters: Optional[Dict] = None) -> List[OutputData]:
|
||||
"""
|
||||
Search for similar vectors using the vector search index.
|
||||
|
||||
Args:
|
||||
query (str): Query string
|
||||
query_vector (List[float]): Query vector.
|
||||
vectors (List[float]): Query vector.
|
||||
limit (int, optional): Number of results to return. Defaults to 5.
|
||||
filters (Dict, optional): Filters to apply to the search.
|
||||
|
||||
@@ -141,24 +139,24 @@ class MongoDB(VectorStoreBase):
|
||||
"index": self.index_name,
|
||||
"limit": limit,
|
||||
"numCandidates": limit,
|
||||
"queryVector": query_vector,
|
||||
"queryVector": vectors,
|
||||
"path": "embedding",
|
||||
}
|
||||
},
|
||||
{"$set": {"score": {"$meta": "vectorSearchScore"}}},
|
||||
{"$project": {"embedding": 0}},
|
||||
]
|
||||
|
||||
|
||||
# Add filter stage if filters are provided
|
||||
if filters:
|
||||
filter_conditions = []
|
||||
for key, value in filters.items():
|
||||
filter_conditions.append({"payload." + key: value})
|
||||
|
||||
|
||||
if filter_conditions:
|
||||
# Add a $match stage after vector search to apply filters
|
||||
pipeline.insert(1, {"$match": {"$and": filter_conditions}})
|
||||
|
||||
|
||||
results = list(collection.aggregate(pipeline))
|
||||
logger.info(f"Vector search completed. Found {len(results)} documents.")
|
||||
except Exception as e:
|
||||
@@ -290,7 +288,7 @@ class MongoDB(VectorStoreBase):
|
||||
filter_conditions.append({"payload." + key: value})
|
||||
if filter_conditions:
|
||||
query = {"$and": filter_conditions}
|
||||
|
||||
|
||||
cursor = self.collection.find(query).limit(limit)
|
||||
results = [OutputData(id=str(doc["_id"]), score=None, payload=doc.get("payload")) for doc in cursor]
|
||||
logger.info(f"Retrieved {len(results)} documents from collection '{self.collection_name}'.")
|
||||
|
||||
+198
-162
@@ -1,28 +1,28 @@
|
||||
import json
|
||||
import logging
|
||||
from typing import List, Optional
|
||||
from contextlib import contextmanager
|
||||
from typing import Any, List, Optional
|
||||
|
||||
from pydantic import BaseModel
|
||||
|
||||
# Try to import psycopg (psycopg3) first, then fall back to psycopg2
|
||||
try:
|
||||
import psycopg
|
||||
from psycopg import execute_values
|
||||
from psycopg.types.json import Json
|
||||
from psycopg_pool import ConnectionPool
|
||||
PSYCOPG_VERSION = 3
|
||||
logger = logging.getLogger(__name__)
|
||||
logger.info("Using psycopg (psycopg3) for PostgreSQL connections")
|
||||
logger.info("Using psycopg (psycopg3) with ConnectionPool for PostgreSQL connections")
|
||||
except ImportError:
|
||||
try:
|
||||
import psycopg2
|
||||
from psycopg2.extras import execute_values, Json
|
||||
from psycopg2.extras import Json, execute_values
|
||||
from psycopg2.pool import ThreadedConnectionPool as ConnectionPool
|
||||
PSYCOPG_VERSION = 2
|
||||
logger = logging.getLogger(__name__)
|
||||
logger.info("Using psycopg2 for PostgreSQL connections")
|
||||
logger.info("Using psycopg2 with ThreadedConnectionPool for PostgreSQL connections")
|
||||
except ImportError:
|
||||
raise ImportError(
|
||||
"Neither 'psycopg' nor 'psycopg2' library is available. "
|
||||
"Please install one of them using 'pip install psycopg' or 'pip install psycopg2'."
|
||||
"Please install one of them using 'pip install psycopg[pool]' or 'pip install psycopg2'"
|
||||
)
|
||||
|
||||
from mem0.vector_stores.base import VectorStoreBase
|
||||
@@ -48,6 +48,8 @@ class PGVector(VectorStoreBase):
|
||||
port,
|
||||
diskann,
|
||||
hnsw,
|
||||
minconn=1,
|
||||
maxconn=5,
|
||||
sslmode=None,
|
||||
connection_string=None,
|
||||
connection_pool=None,
|
||||
@@ -65,6 +67,8 @@ class PGVector(VectorStoreBase):
|
||||
port (int, optional): Database port
|
||||
diskann (bool, optional): Use DiskANN for faster search
|
||||
hnsw (bool, optional): Use HNSW for faster search
|
||||
minconn (int): Minimum number of connections to keep in the connection pool
|
||||
maxconn (int): Maximum number of connections allowed in the connection pool
|
||||
sslmode (str, optional): SSL mode for PostgreSQL connection (e.g., 'require', 'prefer', 'disable')
|
||||
connection_string (str, optional): PostgreSQL connection string (overrides individual connection parameters)
|
||||
connection_pool (Any, optional): psycopg2 connection pool object (overrides connection string and individual parameters)
|
||||
@@ -73,14 +77,13 @@ class PGVector(VectorStoreBase):
|
||||
self.use_diskann = diskann
|
||||
self.use_hnsw = hnsw
|
||||
self.embedding_model_dims = embedding_model_dims
|
||||
self.connection_pool = None
|
||||
|
||||
# Connection setup with priority: connection_pool > connection_string > individual parameters
|
||||
if connection_pool is not None:
|
||||
# Use provided connection pool
|
||||
self.conn = connection_pool.getconn()
|
||||
self.connection_pool = connection_pool
|
||||
elif connection_string is not None:
|
||||
# Use connection string
|
||||
elif connection_string:
|
||||
if sslmode:
|
||||
# Append sslmode to connection string if provided
|
||||
if 'sslmode=' in connection_string:
|
||||
@@ -90,99 +93,119 @@ class PGVector(VectorStoreBase):
|
||||
else:
|
||||
# Add sslmode to connection string
|
||||
connection_string = f"{connection_string} sslmode={sslmode}"
|
||||
|
||||
if PSYCOPG_VERSION == 3:
|
||||
self.conn = psycopg.connect(connection_string)
|
||||
else:
|
||||
self.conn = psycopg2.connect(connection_string)
|
||||
self.connection_pool = None
|
||||
else:
|
||||
# Use individual connection parameters
|
||||
conn_params = {
|
||||
'dbname': dbname,
|
||||
'user': user,
|
||||
'password': password,
|
||||
'host': host,
|
||||
'port': port
|
||||
}
|
||||
connection_string = f"postgresql://{user}:{password}@{host}:{port}/{dbname}"
|
||||
if sslmode:
|
||||
conn_params['sslmode'] = sslmode
|
||||
|
||||
if PSYCOPG_VERSION == 3:
|
||||
self.conn = psycopg.connect(**conn_params)
|
||||
else:
|
||||
self.conn = psycopg2.connect(**conn_params)
|
||||
self.connection_pool = None
|
||||
connection_string = f"{connection_string} sslmode={sslmode}"
|
||||
|
||||
self.cur = self.conn.cursor()
|
||||
if self.connection_pool is None:
|
||||
if PSYCOPG_VERSION == 3:
|
||||
# psycopg3 ConnectionPool
|
||||
self.connection_pool = ConnectionPool(conninfo=connection_string, min_size=minconn, max_size=maxconn, open=True)
|
||||
else:
|
||||
# psycopg2 ThreadedConnectionPool
|
||||
self.connection_pool = ConnectionPool(minconn=minconn, maxconn=maxconn, dsn=connection_string)
|
||||
|
||||
collections = self.list_cols()
|
||||
if collection_name not in collections:
|
||||
self.create_col(embedding_model_dims)
|
||||
self.create_col()
|
||||
|
||||
def create_col(self, embedding_model_dims):
|
||||
@contextmanager
|
||||
def _get_cursor(self, commit: bool = False):
|
||||
"""
|
||||
Unified context manager to get a cursor from the appropriate pool.
|
||||
Auto-commits or rolls back based on exception, and returns the connection to the pool.
|
||||
"""
|
||||
if PSYCOPG_VERSION == 3:
|
||||
# psycopg3 auto-manages commit/rollback and pool return
|
||||
with self.connection_pool.connection() as conn:
|
||||
with conn.cursor() as cur:
|
||||
try:
|
||||
yield cur
|
||||
if commit:
|
||||
conn.commit()
|
||||
except Exception:
|
||||
conn.rollback()
|
||||
logger.error("Error in cursor context (psycopg3)", exc_info=True)
|
||||
raise
|
||||
else:
|
||||
# psycopg2 manual getconn/putconn
|
||||
conn = self.connection_pool.getconn()
|
||||
cur = conn.cursor()
|
||||
try:
|
||||
yield cur
|
||||
if commit:
|
||||
conn.commit()
|
||||
except Exception as exc:
|
||||
conn.rollback()
|
||||
logger.error(f"Error occurred: {exc}")
|
||||
raise exc
|
||||
finally:
|
||||
cur.close()
|
||||
self.connection_pool.putconn(conn)
|
||||
|
||||
def create_col(self) -> None:
|
||||
"""
|
||||
Create a new collection (table in PostgreSQL).
|
||||
Will also initialize vector search index if specified.
|
||||
|
||||
Args:
|
||||
embedding_model_dims (int): Dimension of the embedding vector.
|
||||
"""
|
||||
self.cur.execute("CREATE EXTENSION IF NOT EXISTS vector")
|
||||
self.cur.execute(
|
||||
f"""
|
||||
CREATE TABLE IF NOT EXISTS {self.collection_name} (
|
||||
id UUID PRIMARY KEY,
|
||||
vector vector({embedding_model_dims}),
|
||||
payload JSONB
|
||||
);
|
||||
"""
|
||||
)
|
||||
|
||||
if self.use_diskann and embedding_model_dims < 2000:
|
||||
# Check if vectorscale extension is installed
|
||||
self.cur.execute("SELECT * FROM pg_extension WHERE extname = 'vectorscale'")
|
||||
if self.cur.fetchone():
|
||||
# Create DiskANN index if extension is installed for faster search
|
||||
self.cur.execute(
|
||||
f"""
|
||||
CREATE INDEX IF NOT EXISTS {self.collection_name}_diskann_idx
|
||||
ON {self.collection_name}
|
||||
USING diskann (vector);
|
||||
"""
|
||||
)
|
||||
elif self.use_hnsw:
|
||||
self.cur.execute(
|
||||
with self._get_cursor(commit=True) as cur:
|
||||
cur.execute("CREATE EXTENSION IF NOT EXISTS vector")
|
||||
cur.execute(
|
||||
f"""
|
||||
CREATE INDEX IF NOT EXISTS {self.collection_name}_hnsw_idx
|
||||
ON {self.collection_name}
|
||||
USING hnsw (vector vector_cosine_ops)
|
||||
"""
|
||||
CREATE TABLE IF NOT EXISTS {self.collection_name} (
|
||||
id UUID PRIMARY KEY,
|
||||
vector vector({self.embedding_model_dims}),
|
||||
payload JSONB
|
||||
);
|
||||
"""
|
||||
)
|
||||
if self.use_diskann and self.embedding_model_dims < 2000:
|
||||
cur.execute("SELECT * FROM pg_extension WHERE extname = 'vectorscale'")
|
||||
if cur.fetchone():
|
||||
# Create DiskANN index if extension is installed for faster search
|
||||
cur.execute(
|
||||
f"""
|
||||
CREATE INDEX IF NOT EXISTS {self.collection_name}_diskann_idx
|
||||
ON {self.collection_name}
|
||||
USING diskann (vector);
|
||||
"""
|
||||
)
|
||||
elif self.use_hnsw:
|
||||
cur.execute(
|
||||
f"""
|
||||
CREATE INDEX IF NOT EXISTS {self.collection_name}_hnsw_idx
|
||||
ON {self.collection_name}
|
||||
USING hnsw (vector vector_cosine_ops)
|
||||
"""
|
||||
)
|
||||
|
||||
self.conn.commit()
|
||||
|
||||
def insert(self, vectors, payloads=None, ids=None):
|
||||
"""
|
||||
Insert vectors into a collection.
|
||||
|
||||
Args:
|
||||
vectors (List[List[float]]): List of vectors to insert.
|
||||
payloads (List[Dict], optional): List of payloads corresponding to vectors.
|
||||
ids (List[str], optional): List of IDs corresponding to vectors.
|
||||
"""
|
||||
def insert(self, vectors: list[list[float]], payloads=None, ids=None) -> None:
|
||||
logger.info(f"Inserting {len(vectors)} vectors into collection {self.collection_name}")
|
||||
json_payloads = [json.dumps(payload) for payload in payloads]
|
||||
|
||||
data = [(id, vector, payload) for id, vector, payload in zip(ids, vectors, json_payloads)]
|
||||
execute_values(
|
||||
self.cur,
|
||||
f"INSERT INTO {self.collection_name} (id, vector, payload) VALUES %s",
|
||||
data,
|
||||
)
|
||||
self.conn.commit()
|
||||
if PSYCOPG_VERSION == 3:
|
||||
with self._get_cursor(commit=True) as cur:
|
||||
cur.executemany(
|
||||
f"INSERT INTO {self.collection_name} (id, vector, payload) VALUES (%s, %s, %s)",
|
||||
data,
|
||||
)
|
||||
else:
|
||||
with self._get_cursor(commit=True) as cur:
|
||||
execute_values(
|
||||
cur,
|
||||
f"INSERT INTO {self.collection_name} (id, vector, payload) VALUES %s",
|
||||
data,
|
||||
)
|
||||
|
||||
def search(self, query, vectors, limit=5, filters=None):
|
||||
def search(
|
||||
self,
|
||||
query: str,
|
||||
vectors: list[float],
|
||||
limit: Optional[int] = 5,
|
||||
filters: Optional[dict] = None,
|
||||
) -> List[OutputData]:
|
||||
"""
|
||||
Search for similar vectors.
|
||||
|
||||
@@ -205,31 +228,37 @@ class PGVector(VectorStoreBase):
|
||||
|
||||
filter_clause = "WHERE " + " AND ".join(filter_conditions) if filter_conditions else ""
|
||||
|
||||
self.cur.execute(
|
||||
f"""
|
||||
SELECT id, vector <=> %s::vector AS distance, payload
|
||||
FROM {self.collection_name}
|
||||
{filter_clause}
|
||||
ORDER BY distance
|
||||
LIMIT %s
|
||||
""",
|
||||
(vectors, *filter_params, limit),
|
||||
)
|
||||
with self._get_cursor() as cur:
|
||||
cur.execute(
|
||||
f"""
|
||||
SELECT id, vector <=> %s::vector AS distance, payload
|
||||
FROM {self.collection_name}
|
||||
{filter_clause}
|
||||
ORDER BY distance
|
||||
LIMIT %s
|
||||
""",
|
||||
(vectors, *filter_params, limit),
|
||||
)
|
||||
|
||||
results = self.cur.fetchall()
|
||||
results = cur.fetchall()
|
||||
return [OutputData(id=str(r[0]), score=float(r[1]), payload=r[2]) for r in results]
|
||||
|
||||
def delete(self, vector_id):
|
||||
def delete(self, vector_id: str) -> None:
|
||||
"""
|
||||
Delete a vector by ID.
|
||||
|
||||
Args:
|
||||
vector_id (str): ID of the vector to delete.
|
||||
"""
|
||||
self.cur.execute(f"DELETE FROM {self.collection_name} WHERE id = %s", (vector_id,))
|
||||
self.conn.commit()
|
||||
with self._get_cursor(commit=True) as cur:
|
||||
cur.execute(f"DELETE FROM {self.collection_name} WHERE id = %s", (vector_id,))
|
||||
|
||||
def update(self, vector_id, vector=None, payload=None):
|
||||
def update(
|
||||
self,
|
||||
vector_id: str,
|
||||
vector: Optional[list[float]] = None,
|
||||
payload: Optional[dict] = None,
|
||||
) -> None:
|
||||
"""
|
||||
Update a vector and its payload.
|
||||
|
||||
@@ -238,28 +267,29 @@ class PGVector(VectorStoreBase):
|
||||
vector (List[float], optional): Updated vector.
|
||||
payload (Dict, optional): Updated payload.
|
||||
"""
|
||||
if vector:
|
||||
self.cur.execute(
|
||||
f"UPDATE {self.collection_name} SET vector = %s WHERE id = %s",
|
||||
(vector, vector_id),
|
||||
)
|
||||
if payload:
|
||||
# Handle JSON serialization based on psycopg version
|
||||
if PSYCOPG_VERSION == 3:
|
||||
# psycopg3 uses psycopg.types.json.Json
|
||||
self.cur.execute(
|
||||
f"UPDATE {self.collection_name} SET payload = %s WHERE id = %s",
|
||||
(Json(payload), vector_id),
|
||||
with self._get_cursor(commit=True) as cur:
|
||||
if vector:
|
||||
cur.execute(
|
||||
f"UPDATE {self.collection_name} SET vector = %s WHERE id = %s",
|
||||
(vector, vector_id),
|
||||
)
|
||||
else:
|
||||
# psycopg2 uses psycopg2.extras.Json
|
||||
self.cur.execute(
|
||||
f"UPDATE {self.collection_name} SET payload = %s WHERE id = %s",
|
||||
(psycopg2.extras.Json(payload), vector_id),
|
||||
)
|
||||
self.conn.commit()
|
||||
if payload:
|
||||
# Handle JSON serialization based on psycopg version
|
||||
if PSYCOPG_VERSION == 3:
|
||||
# psycopg3 uses psycopg.types.json.Json
|
||||
cur.execute(
|
||||
f"UPDATE {self.collection_name} SET payload = %s WHERE id = %s",
|
||||
(Json(payload), vector_id),
|
||||
)
|
||||
else:
|
||||
# psycopg2 uses psycopg2.extras.Json
|
||||
cur.execute(
|
||||
f"UPDATE {self.collection_name} SET payload = %s WHERE id = %s",
|
||||
(Json(payload), vector_id),
|
||||
)
|
||||
|
||||
def get(self, vector_id) -> OutputData:
|
||||
|
||||
def get(self, vector_id: str) -> OutputData:
|
||||
"""
|
||||
Retrieve a vector by ID.
|
||||
|
||||
@@ -269,14 +299,15 @@ class PGVector(VectorStoreBase):
|
||||
Returns:
|
||||
OutputData: Retrieved vector.
|
||||
"""
|
||||
self.cur.execute(
|
||||
f"SELECT id, vector, payload FROM {self.collection_name} WHERE id = %s",
|
||||
(vector_id,),
|
||||
)
|
||||
result = self.cur.fetchone()
|
||||
if not result:
|
||||
return None
|
||||
return OutputData(id=str(result[0]), score=None, payload=result[2])
|
||||
with self._get_cursor() as cur:
|
||||
cur.execute(
|
||||
f"SELECT id, vector, payload FROM {self.collection_name} WHERE id = %s",
|
||||
(vector_id,),
|
||||
)
|
||||
result = cur.fetchone()
|
||||
if not result:
|
||||
return None
|
||||
return OutputData(id=str(result[0]), score=None, payload=result[2])
|
||||
|
||||
def list_cols(self) -> List[str]:
|
||||
"""
|
||||
@@ -285,36 +316,42 @@ class PGVector(VectorStoreBase):
|
||||
Returns:
|
||||
List[str]: List of collection names.
|
||||
"""
|
||||
self.cur.execute("SELECT table_name FROM information_schema.tables WHERE table_schema = 'public'")
|
||||
return [row[0] for row in self.cur.fetchall()]
|
||||
with self._get_cursor() as cur:
|
||||
cur.execute("SELECT table_name FROM information_schema.tables WHERE table_schema = 'public'")
|
||||
return [row[0] for row in cur.fetchall()]
|
||||
|
||||
def delete_col(self):
|
||||
def delete_col(self) -> None:
|
||||
"""Delete a collection."""
|
||||
self.cur.execute(f"DROP TABLE IF EXISTS {self.collection_name}")
|
||||
self.conn.commit()
|
||||
with self._get_cursor(commit=True) as cur:
|
||||
cur.execute(f"DROP TABLE IF EXISTS {self.collection_name}")
|
||||
|
||||
def col_info(self):
|
||||
def col_info(self) -> dict[str, Any]:
|
||||
"""
|
||||
Get information about a collection.
|
||||
|
||||
Returns:
|
||||
Dict[str, Any]: Collection information.
|
||||
"""
|
||||
self.cur.execute(
|
||||
f"""
|
||||
SELECT
|
||||
table_name,
|
||||
(SELECT COUNT(*) FROM {self.collection_name}) as row_count,
|
||||
(SELECT pg_size_pretty(pg_total_relation_size('{self.collection_name}'))) as total_size
|
||||
FROM information_schema.tables
|
||||
WHERE table_schema = 'public' AND table_name = %s
|
||||
""",
|
||||
(self.collection_name,),
|
||||
)
|
||||
result = self.cur.fetchone()
|
||||
with self._get_cursor() as cur:
|
||||
cur.execute(
|
||||
f"""
|
||||
SELECT
|
||||
table_name,
|
||||
(SELECT COUNT(*) FROM {self.collection_name}) as row_count,
|
||||
(SELECT pg_size_pretty(pg_total_relation_size('{self.collection_name}'))) as total_size
|
||||
FROM information_schema.tables
|
||||
WHERE table_schema = 'public' AND table_name = %s
|
||||
""",
|
||||
(self.collection_name,),
|
||||
)
|
||||
result = cur.fetchone()
|
||||
return {"name": result[0], "count": result[1], "size": result[2]}
|
||||
|
||||
def list(self, filters=None, limit=100):
|
||||
def list(
|
||||
self,
|
||||
filters: Optional[dict] = None,
|
||||
limit: Optional[int] = 100
|
||||
) -> List[OutputData]:
|
||||
"""
|
||||
List all vectors in a collection.
|
||||
|
||||
@@ -342,27 +379,26 @@ class PGVector(VectorStoreBase):
|
||||
LIMIT %s
|
||||
"""
|
||||
|
||||
self.cur.execute(query, (*filter_params, limit))
|
||||
|
||||
results = self.cur.fetchall()
|
||||
with self._get_cursor() as cur:
|
||||
cur.execute(query, (*filter_params, limit))
|
||||
results = cur.fetchall()
|
||||
return [[OutputData(id=str(r[0]), score=None, payload=r[2]) for r in results]]
|
||||
|
||||
def __del__(self):
|
||||
def __del__(self) -> None:
|
||||
"""
|
||||
Close the database connection when the object is deleted.
|
||||
Close the database connection pool when the object is deleted.
|
||||
"""
|
||||
if hasattr(self, "cur"):
|
||||
self.cur.close()
|
||||
if hasattr(self, "conn"):
|
||||
if hasattr(self, "connection_pool") and self.connection_pool is not None:
|
||||
# Return connection to pool instead of closing it
|
||||
self.connection_pool.putconn(self.conn)
|
||||
try:
|
||||
# Close pool appropriately
|
||||
if PSYCOPG_VERSION == 3:
|
||||
self.connection_pool.close()
|
||||
else:
|
||||
# Close the connection directly
|
||||
self.conn.close()
|
||||
self.connection_pool.closeall()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
def reset(self):
|
||||
def reset(self) -> None:
|
||||
"""Reset the index by deleting and recreating it."""
|
||||
logger.warning(f"Resetting index {self.collection_name}...")
|
||||
self.delete_col()
|
||||
self.create_col(self.embedding_model_dims)
|
||||
self.create_col()
|
||||
|
||||
@@ -3,6 +3,7 @@ import uuid
|
||||
from typing import Dict, List, Mapping, Optional
|
||||
|
||||
from pydantic import BaseModel
|
||||
from urllib.parse import urlparse
|
||||
|
||||
try:
|
||||
import weaviate
|
||||
@@ -12,7 +13,7 @@ except ImportError:
|
||||
)
|
||||
|
||||
import weaviate.classes.config as wvcc
|
||||
from weaviate.classes.init import Auth
|
||||
from weaviate.classes.init import Auth, AdditionalConfig, Timeout
|
||||
from weaviate.classes.query import Filter, MetadataQuery
|
||||
from weaviate.util import get_valid_uuid
|
||||
|
||||
@@ -47,14 +48,36 @@ class Weaviate(VectorStoreBase):
|
||||
auth_config (dict, optional): Authentication configuration for Weaviate. Defaults to None.
|
||||
additional_headers (dict, optional): Additional headers for requests. Defaults to None.
|
||||
"""
|
||||
if "localhost" in cluster_url:
|
||||
if "localhost" in cluster_url:
|
||||
self.client = weaviate.connect_to_local(headers=additional_headers)
|
||||
else:
|
||||
elif auth_client_secret:
|
||||
self.client = weaviate.connect_to_wcs(
|
||||
cluster_url=cluster_url,
|
||||
auth_credentials=Auth.api_key(auth_client_secret),
|
||||
headers=additional_headers,
|
||||
)
|
||||
else:
|
||||
parsed = urlparse(cluster_url) # e.g., http://mem0_store:8080
|
||||
http_host = parsed.hostname or "localhost"
|
||||
http_port = parsed.port or (443 if parsed.scheme == "https" else 8080)
|
||||
http_secure = parsed.scheme == "https"
|
||||
|
||||
# Weaviate gRPC defaults (inside Docker network)
|
||||
grpc_host = http_host
|
||||
grpc_port = 50051
|
||||
grpc_secure = False
|
||||
|
||||
self.client = weaviate.connect_to_custom(
|
||||
http_host,
|
||||
http_port,
|
||||
http_secure,
|
||||
grpc_host,
|
||||
grpc_port,
|
||||
grpc_secure,
|
||||
headers=additional_headers,
|
||||
skip_init_checks=True,
|
||||
additional_config=AdditionalConfig(timeout=Timeout(init=2.0))
|
||||
)
|
||||
|
||||
self.collection_name = collection_name
|
||||
self.embedding_model_dims = embedding_model_dims
|
||||
|
||||
@@ -31,7 +31,6 @@ from fastapi import FastAPI, Request
|
||||
from fastapi.routing import APIRouter
|
||||
from mcp.server.fastmcp import FastMCP
|
||||
from mcp.server.sse import SseServerTransport
|
||||
from qdrant_client import models as qdrant_models
|
||||
|
||||
# Load environment variables
|
||||
load_dotenv()
|
||||
@@ -165,74 +164,54 @@ async def search_memory(query: str) -> str:
|
||||
# Get accessible memory IDs based on ACL
|
||||
user_memories = db.query(Memory).filter(Memory.user_id == user.id).all()
|
||||
accessible_memory_ids = [memory.id for memory in user_memories if check_memory_access_permissions(db, memory, app.id)]
|
||||
|
||||
conditions = [qdrant_models.FieldCondition(key="user_id", match=qdrant_models.MatchValue(value=uid))]
|
||||
|
||||
if accessible_memory_ids:
|
||||
# Convert UUIDs to strings for Qdrant
|
||||
accessible_memory_ids_str = [str(memory_id) for memory_id in accessible_memory_ids]
|
||||
conditions.append(qdrant_models.HasIdCondition(has_id=accessible_memory_ids_str))
|
||||
|
||||
filters = qdrant_models.Filter(must=conditions)
|
||||
filters = {
|
||||
"user_id": uid
|
||||
}
|
||||
|
||||
embeddings = memory_client.embedding_model.embed(query, "search")
|
||||
|
||||
hits = memory_client.vector_store.client.query_points(
|
||||
collection_name=memory_client.vector_store.collection_name,
|
||||
query=embeddings,
|
||||
query_filter=filters,
|
||||
limit=10,
|
||||
|
||||
hits = memory_client.vector_store.search(
|
||||
query=query,
|
||||
vectors=embeddings,
|
||||
limit=10,
|
||||
filters=filters,
|
||||
)
|
||||
|
||||
# Process search results
|
||||
memories = hits.points
|
||||
memories = [
|
||||
{
|
||||
"id": memory.id,
|
||||
"memory": memory.payload["data"],
|
||||
"hash": memory.payload.get("hash"),
|
||||
"created_at": memory.payload.get("created_at"),
|
||||
"updated_at": memory.payload.get("updated_at"),
|
||||
"score": memory.score,
|
||||
}
|
||||
for memory in memories
|
||||
]
|
||||
allowed = set(str(mid) for mid in accessible_memory_ids) if accessible_memory_ids else None
|
||||
|
||||
# Log memory access for each memory found
|
||||
if isinstance(memories, dict) and 'results' in memories:
|
||||
print(f"Memories: {memories}")
|
||||
for memory_data in memories['results']:
|
||||
if 'id' in memory_data:
|
||||
memory_id = uuid.UUID(memory_data['id'])
|
||||
# Create access log entry
|
||||
access_log = MemoryAccessLog(
|
||||
memory_id=memory_id,
|
||||
app_id=app.id,
|
||||
access_type="search",
|
||||
metadata_={
|
||||
"query": query,
|
||||
"score": memory_data.get('score'),
|
||||
"hash": memory_data.get('hash')
|
||||
}
|
||||
)
|
||||
db.add(access_log)
|
||||
db.commit()
|
||||
else:
|
||||
for memory in memories:
|
||||
memory_id = uuid.UUID(memory['id'])
|
||||
# Create access log entry
|
||||
results = []
|
||||
for h in hits:
|
||||
# All vector db search functions return OutputData class
|
||||
id, score, payload = h.id, h.score, h.payload
|
||||
if allowed and h.id is None or h.id not in allowed:
|
||||
continue
|
||||
|
||||
results.append({
|
||||
"id": id,
|
||||
"memory": payload.get("data"),
|
||||
"hash": payload.get("hash"),
|
||||
"created_at": payload.get("created_at"),
|
||||
"updated_at": payload.get("updated_at"),
|
||||
"score": score,
|
||||
})
|
||||
|
||||
for r in results:
|
||||
if r.get("id"):
|
||||
access_log = MemoryAccessLog(
|
||||
memory_id=memory_id,
|
||||
memory_id=uuid.UUID(r["id"]),
|
||||
app_id=app.id,
|
||||
access_type="search",
|
||||
metadata_={
|
||||
"query": query,
|
||||
"score": memory.get('score'),
|
||||
"hash": memory.get('hash')
|
||||
}
|
||||
"score": r.get("score"),
|
||||
"hash": r.get("hash"),
|
||||
},
|
||||
)
|
||||
db.add(access_log)
|
||||
db.commit()
|
||||
return json.dumps(memories, indent=2)
|
||||
db.commit()
|
||||
|
||||
return json.dumps({"results": results}, indent=2)
|
||||
finally:
|
||||
db.close()
|
||||
except Exception as e:
|
||||
|
||||
@@ -260,6 +260,8 @@ async def create_memory(
|
||||
|
||||
# Process Qdrant response
|
||||
if isinstance(qdrant_response, dict) and 'results' in qdrant_response:
|
||||
created_memories = []
|
||||
|
||||
for result in qdrant_response['results']:
|
||||
if result['event'] == 'ADD':
|
||||
# Get the Qdrant-generated ID
|
||||
@@ -294,9 +296,17 @@ async def create_memory(
|
||||
)
|
||||
db.add(history)
|
||||
|
||||
db.commit()
|
||||
created_memories.append(memory)
|
||||
|
||||
# Commit all changes at once
|
||||
if created_memories:
|
||||
db.commit()
|
||||
for memory in created_memories:
|
||||
db.refresh(memory)
|
||||
return memory
|
||||
|
||||
# Return the first memory (for API compatibility)
|
||||
# but all memories are now saved to the database
|
||||
return created_memories[0]
|
||||
except Exception as qdrant_error:
|
||||
logging.warning(f"Qdrant operation failed: {qdrant_error}.")
|
||||
# Return a json response with the error
|
||||
|
||||
@@ -135,14 +135,111 @@ def reset_memory_client():
|
||||
|
||||
def get_default_memory_config():
|
||||
"""Get default memory client configuration with sensible defaults."""
|
||||
# Detect vector store based on environment variables
|
||||
vector_store_config = {
|
||||
"collection_name": "openmemory",
|
||||
"host": "mem0_store",
|
||||
}
|
||||
|
||||
# Check for different vector store configurations based on environment variables
|
||||
if os.environ.get('CHROMA_HOST') and os.environ.get('CHROMA_PORT'):
|
||||
vector_store_provider = "chroma"
|
||||
vector_store_config.update({
|
||||
"host": os.environ.get('CHROMA_HOST'),
|
||||
"port": int(os.environ.get('CHROMA_PORT'))
|
||||
})
|
||||
elif os.environ.get('QDRANT_HOST') and os.environ.get('QDRANT_PORT'):
|
||||
vector_store_provider = "qdrant"
|
||||
vector_store_config.update({
|
||||
"host": os.environ.get('QDRANT_HOST'),
|
||||
"port": int(os.environ.get('QDRANT_PORT'))
|
||||
})
|
||||
elif os.environ.get('WEAVIATE_CLUSTER_URL') or (os.environ.get('WEAVIATE_HOST') and os.environ.get('WEAVIATE_PORT')):
|
||||
vector_store_provider = "weaviate"
|
||||
# Prefer an explicit cluster URL if provided; otherwise build from host/port
|
||||
cluster_url = os.environ.get('WEAVIATE_CLUSTER_URL')
|
||||
if not cluster_url:
|
||||
weaviate_host = os.environ.get('WEAVIATE_HOST')
|
||||
weaviate_port = int(os.environ.get('WEAVIATE_PORT'))
|
||||
cluster_url = f"http://{weaviate_host}:{weaviate_port}"
|
||||
vector_store_config = {
|
||||
"collection_name": "openmemory",
|
||||
"cluster_url": cluster_url
|
||||
}
|
||||
elif os.environ.get('REDIS_URL'):
|
||||
vector_store_provider = "redis"
|
||||
vector_store_config = {
|
||||
"collection_name": "openmemory",
|
||||
"redis_url": os.environ.get('REDIS_URL')
|
||||
}
|
||||
elif os.environ.get('PG_HOST') and os.environ.get('PG_PORT'):
|
||||
vector_store_provider = "pgvector"
|
||||
vector_store_config.update({
|
||||
"host": os.environ.get('PG_HOST'),
|
||||
"port": int(os.environ.get('PG_PORT')),
|
||||
"dbname": os.environ.get('PG_DB', 'mem0'),
|
||||
"user": os.environ.get('PG_USER', 'mem0'),
|
||||
"password": os.environ.get('PG_PASSWORD', 'mem0')
|
||||
})
|
||||
elif os.environ.get('MILVUS_HOST') and os.environ.get('MILVUS_PORT'):
|
||||
vector_store_provider = "milvus"
|
||||
# Construct the full URL as expected by MilvusDBConfig
|
||||
milvus_host = os.environ.get('MILVUS_HOST')
|
||||
milvus_port = int(os.environ.get('MILVUS_PORT'))
|
||||
milvus_url = f"http://{milvus_host}:{milvus_port}"
|
||||
|
||||
vector_store_config = {
|
||||
"collection_name": "openmemory",
|
||||
"url": milvus_url,
|
||||
"token": os.environ.get('MILVUS_TOKEN', ''), # Always include, empty string for local setup
|
||||
"db_name": os.environ.get('MILVUS_DB_NAME', ''),
|
||||
"embedding_model_dims": 1536,
|
||||
"metric_type": "COSINE" # Using COSINE for better semantic similarity
|
||||
}
|
||||
elif os.environ.get('ELASTICSEARCH_HOST') and os.environ.get('ELASTICSEARCH_PORT'):
|
||||
vector_store_provider = "elasticsearch"
|
||||
# Construct the full URL with scheme since Elasticsearch client expects it
|
||||
elasticsearch_host = os.environ.get('ELASTICSEARCH_HOST')
|
||||
elasticsearch_port = int(os.environ.get('ELASTICSEARCH_PORT'))
|
||||
# Use http:// scheme since we're not using SSL
|
||||
full_host = f"http://{elasticsearch_host}"
|
||||
|
||||
vector_store_config.update({
|
||||
"host": full_host,
|
||||
"port": elasticsearch_port,
|
||||
"user": os.environ.get('ELASTICSEARCH_USER', 'elastic'),
|
||||
"password": os.environ.get('ELASTICSEARCH_PASSWORD', 'changeme'),
|
||||
"verify_certs": False,
|
||||
"use_ssl": False,
|
||||
"embedding_model_dims": 1536
|
||||
})
|
||||
elif os.environ.get('OPENSEARCH_HOST') and os.environ.get('OPENSEARCH_PORT'):
|
||||
vector_store_provider = "opensearch"
|
||||
vector_store_config.update({
|
||||
"host": os.environ.get('OPENSEARCH_HOST'),
|
||||
"port": int(os.environ.get('OPENSEARCH_PORT'))
|
||||
})
|
||||
elif os.environ.get('FAISS_PATH'):
|
||||
vector_store_provider = "faiss"
|
||||
vector_store_config = {
|
||||
"collection_name": "openmemory",
|
||||
"path": os.environ.get('FAISS_PATH'),
|
||||
"embedding_model_dims": 1536,
|
||||
"distance_strategy": "cosine"
|
||||
}
|
||||
else:
|
||||
# Default fallback to Qdrant
|
||||
vector_store_provider = "qdrant"
|
||||
vector_store_config.update({
|
||||
"port": 6333,
|
||||
})
|
||||
|
||||
print(f"Auto-detected vector store: {vector_store_provider} with config: {vector_store_config}")
|
||||
|
||||
return {
|
||||
"vector_store": {
|
||||
"provider": "qdrant",
|
||||
"config": {
|
||||
"collection_name": "openmemory",
|
||||
"host": "mem0_store",
|
||||
"port": 6333,
|
||||
}
|
||||
"provider": vector_store_provider,
|
||||
"config": vector_store_config
|
||||
},
|
||||
"llm": {
|
||||
"provider": "openai",
|
||||
@@ -242,6 +339,9 @@ def get_memory_client(custom_instructions: str = None):
|
||||
# Fix Ollama URLs for Docker if needed
|
||||
if config["embedder"].get("provider") == "ollama":
|
||||
config["embedder"] = _fix_ollama_urls(config["embedder"])
|
||||
|
||||
if "vector_store" in mem0_config and mem0_config["vector_store"] is not None:
|
||||
config["vector_store"] = mem0_config["vector_store"]
|
||||
else:
|
||||
print("No configuration found in database, using defaults")
|
||||
|
||||
|
||||
@@ -0,0 +1,11 @@
|
||||
services:
|
||||
mem0_store:
|
||||
image: ghcr.io/chroma-core/chroma:latest
|
||||
restart: unless-stopped
|
||||
environment:
|
||||
- CHROMA_SERVER_HOST=0.0.0.0
|
||||
- CHROMA_SERVER_HTTP_PORT=8000
|
||||
ports:
|
||||
- "8000:8000"
|
||||
volumes:
|
||||
- mem0_storage:/data
|
||||
@@ -0,0 +1,15 @@
|
||||
services:
|
||||
mem0_store:
|
||||
image: docker.elastic.co/elasticsearch/elasticsearch:8.13.4
|
||||
restart: unless-stopped
|
||||
environment:
|
||||
- discovery.type=single-node
|
||||
- xpack.security.enabled=false
|
||||
- ES_JAVA_OPTS=-Xms512m -Xmx512m
|
||||
ulimits:
|
||||
memlock: { soft: -1, hard: -1 }
|
||||
nofile: { soft: 65536, hard: 65536 }
|
||||
ports:
|
||||
- "9200:9200"
|
||||
volumes:
|
||||
- mem0_storage:/usr/share/elasticsearch/data
|
||||
@@ -0,0 +1,3 @@
|
||||
services:
|
||||
# FAISS is a local file-based vector store, so no separate container is needed
|
||||
# Data will be persisted through volume mounts in the main application
|
||||
@@ -0,0 +1,43 @@
|
||||
services:
|
||||
etcd:
|
||||
image: quay.io/coreos/etcd:v3.5.5
|
||||
restart: unless-stopped
|
||||
environment:
|
||||
- ETCD_AUTO_COMPACTION_MODE=revision
|
||||
- ETCD_QUOTA_BACKEND_BYTES=4294967296
|
||||
- ETCD_SNAPSHOT_COUNT=50000
|
||||
- ETCD_LISTEN_CLIENT_URLS=http://0.0.0.0:2379
|
||||
- ETCD_ADVERTISE_CLIENT_URLS=http://etcd:2379
|
||||
- ETCD_LISTEN_PEER_URLS=http://0.0.0.0:2380
|
||||
- ETCD_INITIAL_ADVERTISE_PEER_URLS=http://etcd:2380
|
||||
- ETCD_INITIAL_CLUSTER=default=http://etcd:2380
|
||||
- ETCD_NAME=default
|
||||
- ETCD_DATA_DIR=/etcd
|
||||
volumes:
|
||||
- ./data/milvus/etcd:/etcd
|
||||
|
||||
minio:
|
||||
image: minio/minio:RELEASE.2023-10-25T06-33-25Z
|
||||
restart: unless-stopped
|
||||
command: server /minio_data
|
||||
environment:
|
||||
- MINIO_ACCESS_KEY=minioadmin
|
||||
- MINIO_SECRET_KEY=minioadmin
|
||||
volumes:
|
||||
- ./data/milvus/minio:/minio_data
|
||||
|
||||
mem0_store:
|
||||
image: milvusdb/milvus:v2.4.7
|
||||
restart: unless-stopped
|
||||
command: ["milvus", "run", "standalone"]
|
||||
depends_on:
|
||||
- etcd
|
||||
- minio
|
||||
environment:
|
||||
- ETCD_ENDPOINTS=etcd:2379
|
||||
- MINIO_ADDRESS=minio:9000
|
||||
ports:
|
||||
- "19530:19530"
|
||||
- "9091:9091"
|
||||
volumes:
|
||||
- ./data/milvus/milvus:/var/lib/milvus
|
||||
@@ -0,0 +1,19 @@
|
||||
services:
|
||||
mem0_store:
|
||||
image: opensearchproject/opensearch:2.13.0
|
||||
restart: unless-stopped
|
||||
user: "1000:1000"
|
||||
environment:
|
||||
- discovery.type=single-node
|
||||
- plugins.security.disabled=true
|
||||
- OPENSEARCH_JAVA_OPTS=-Xms512m -Xmx512m
|
||||
- OPENSEARCH_INITIAL_ADMIN_PASSWORD=Openmemory123!
|
||||
- bootstrap.memory_lock=true
|
||||
ulimits:
|
||||
memlock: { soft: -1, hard: -1 }
|
||||
nofile: { soft: 65536, hard: 65536 }
|
||||
ports:
|
||||
- "9200:9200"
|
||||
- "9600:9600"
|
||||
volumes:
|
||||
- mem0_storage:/usr/share/opensearch/data
|
||||
@@ -0,0 +1,12 @@
|
||||
services:
|
||||
mem0_store:
|
||||
image: pgvector/pgvector:pg16
|
||||
restart: unless-stopped
|
||||
environment:
|
||||
- POSTGRES_DB=mem0
|
||||
- POSTGRES_USER=mem0
|
||||
- POSTGRES_PASSWORD=mem0
|
||||
ports:
|
||||
- "5432:5432"
|
||||
volumes:
|
||||
- mem0_storage:/var/lib/postgresql/data
|
||||
@@ -0,0 +1,8 @@
|
||||
services:
|
||||
mem0_store:
|
||||
image: qdrant/qdrant:latest
|
||||
restart: unless-stopped
|
||||
ports:
|
||||
- "6333:6333"
|
||||
volumes:
|
||||
- mem0_storage:/mem0/storage
|
||||
@@ -0,0 +1,13 @@
|
||||
services:
|
||||
mem0_store:
|
||||
image: redis/redis-stack-server:latest
|
||||
restart: unless-stopped
|
||||
ports:
|
||||
- "6379:6379"
|
||||
volumes:
|
||||
- mem0_storage:/var/lib/redis-stack
|
||||
command: >
|
||||
redis-stack-server
|
||||
--appendonly yes
|
||||
--appendfsync everysec
|
||||
--save 900 1 300 10 60 10000
|
||||
@@ -0,0 +1,14 @@
|
||||
services:
|
||||
mem0_store:
|
||||
image: semitechnologies/weaviate:latest
|
||||
restart: unless-stopped
|
||||
environment:
|
||||
- QUERY_DEFAULTS_LIMIT=25
|
||||
- AUTHENTICATION_ANONYMOUS_ACCESS_ENABLED=true
|
||||
- PERSISTENCE_DATA_PATH=/var/lib/weaviate
|
||||
- CLUSTER_HOSTNAME=node1
|
||||
- WEAVIATE_CLUSTER_URL=http://mem0_store:8080
|
||||
ports:
|
||||
- "8080:8080"
|
||||
volumes:
|
||||
- mem0_storage:/var/lib/weaviate
|
||||
+301
-12
@@ -54,36 +54,325 @@ export NEXT_PUBLIC_API_URL
|
||||
export NEXT_PUBLIC_USER_ID="$USER"
|
||||
export FRONTEND_PORT
|
||||
|
||||
# Create docker-compose.yml file
|
||||
echo "📝 Creating docker-compose.yml..."
|
||||
cat > docker-compose.yml <<EOF
|
||||
services:
|
||||
mem0_store:
|
||||
image: qdrant/qdrant
|
||||
ports:
|
||||
- "6333:6333"
|
||||
volumes:
|
||||
- mem0_storage:/mem0/storage
|
||||
# Parse vector store selection (env var or flag). Default: qdrant
|
||||
VECTOR_STORE="${VECTOR_STORE:-qdrant}"
|
||||
EMBEDDING_DIMS="${EMBEDDING_DIMS:-1536}"
|
||||
|
||||
for arg in "$@"; do
|
||||
case $arg in
|
||||
--vector-store=*)
|
||||
VECTOR_STORE="${arg#*=}"
|
||||
shift
|
||||
;;
|
||||
--vector-store)
|
||||
VECTOR_STORE="$2"
|
||||
shift 2
|
||||
;;
|
||||
*)
|
||||
;;
|
||||
esac
|
||||
done
|
||||
|
||||
export VECTOR_STORE
|
||||
echo "🧰 Using vector store: $VECTOR_STORE"
|
||||
|
||||
# Function to create compose file by merging vector store config with openmemory-mcp service
|
||||
create_compose_file() {
|
||||
local vector_store=$1
|
||||
local compose_file="compose/${vector_store}.yml"
|
||||
local volume_name="${vector_store}_data" # Vector-store-specific volume name
|
||||
|
||||
# Check if the compose file exists
|
||||
if [ ! -f "$compose_file" ]; then
|
||||
echo "❌ Compose file not found: $compose_file"
|
||||
echo "Available vector stores: $(ls compose/*.yml | sed 's/compose\///g' | sed 's/\.yml//g' | tr '\n' ' ')"
|
||||
exit 1
|
||||
fi
|
||||
|
||||
echo "📝 Creating docker-compose.yml using $compose_file..."
|
||||
echo "💾 Using volume: $volume_name"
|
||||
|
||||
# Start the compose file with services section
|
||||
echo "services:" > docker-compose.yml
|
||||
|
||||
# Extract services from the compose file and replace volume name
|
||||
# First get everything except the last volumes section
|
||||
tail -n +2 "$compose_file" | sed '/^volumes:/,$d' | sed "s/mem0_storage/${volume_name}/g" >> docker-compose.yml
|
||||
|
||||
# Add a newline to ensure proper YAML formatting
|
||||
echo "" >> docker-compose.yml
|
||||
|
||||
# Add the openmemory-mcp service
|
||||
cat >> docker-compose.yml <<EOF
|
||||
openmemory-mcp:
|
||||
image: mem0/openmemory-mcp:latest
|
||||
environment:
|
||||
- OPENAI_API_KEY=${OPENAI_API_KEY}
|
||||
- USER=${USER}
|
||||
EOF
|
||||
|
||||
# Add vector store specific environment variables
|
||||
case "$vector_store" in
|
||||
weaviate)
|
||||
cat >> docker-compose.yml <<EOF
|
||||
- WEAVIATE_HOST=mem0_store
|
||||
- WEAVIATE_PORT=8080
|
||||
EOF
|
||||
;;
|
||||
redis)
|
||||
cat >> docker-compose.yml <<EOF
|
||||
- REDIS_URL=redis://mem0_store:6379
|
||||
EOF
|
||||
;;
|
||||
pgvector)
|
||||
cat >> docker-compose.yml <<EOF
|
||||
- PG_HOST=mem0_store
|
||||
- PG_PORT=5432
|
||||
- PG_DB=mem0
|
||||
- PG_USER=mem0
|
||||
- PG_PASSWORD=mem0
|
||||
EOF
|
||||
;;
|
||||
qdrant)
|
||||
cat >> docker-compose.yml <<EOF
|
||||
- QDRANT_HOST=mem0_store
|
||||
- QDRANT_PORT=6333
|
||||
EOF
|
||||
;;
|
||||
chroma)
|
||||
cat >> docker-compose.yml <<EOF
|
||||
- CHROMA_HOST=mem0_store
|
||||
- CHROMA_PORT=8000
|
||||
EOF
|
||||
;;
|
||||
milvus)
|
||||
cat >> docker-compose.yml <<EOF
|
||||
- MILVUS_HOST=mem0_store
|
||||
- MILVUS_PORT=19530
|
||||
EOF
|
||||
;;
|
||||
elasticsearch)
|
||||
cat >> docker-compose.yml <<EOF
|
||||
- ELASTICSEARCH_HOST=mem0_store
|
||||
- ELASTICSEARCH_PORT=9200
|
||||
- ELASTICSEARCH_USER=elastic
|
||||
- ELASTICSEARCH_PASSWORD=changeme
|
||||
EOF
|
||||
;;
|
||||
faiss)
|
||||
cat >> docker-compose.yml <<EOF
|
||||
- FAISS_PATH=/tmp/faiss
|
||||
EOF
|
||||
;;
|
||||
*)
|
||||
echo "⚠️ Unknown vector store: $vector_store. Using default Qdrant configuration."
|
||||
cat >> docker-compose.yml <<EOF
|
||||
- QDRANT_HOST=mem0_store
|
||||
- QDRANT_PORT=6333
|
||||
EOF
|
||||
;;
|
||||
esac
|
||||
|
||||
# Add common openmemory-mcp service configuration
|
||||
if [ "$vector_store" = "faiss" ]; then
|
||||
# FAISS doesn't need a separate service, just volume mounts
|
||||
cat >> docker-compose.yml <<EOF
|
||||
ports:
|
||||
- "8765:8765"
|
||||
volumes:
|
||||
- openmemory_db:/usr/src/openmemory
|
||||
- ${volume_name}:/tmp/faiss
|
||||
|
||||
volumes:
|
||||
${volume_name}:
|
||||
openmemory_db:
|
||||
EOF
|
||||
else
|
||||
cat >> docker-compose.yml <<EOF
|
||||
depends_on:
|
||||
- mem0_store
|
||||
ports:
|
||||
- "8765:8765"
|
||||
volumes:
|
||||
- openmemory_db:/usr/src/openmemory
|
||||
|
||||
volumes:
|
||||
mem0_storage:
|
||||
${volume_name}:
|
||||
openmemory_db:
|
||||
EOF
|
||||
fi
|
||||
}
|
||||
|
||||
# Create docker-compose.yml file based on selected vector store
|
||||
echo "📝 Creating docker-compose.yml..."
|
||||
create_compose_file "$VECTOR_STORE"
|
||||
|
||||
# Ensure local data directories exist for bind-mounted vector stores
|
||||
if [ "$VECTOR_STORE" = "milvus" ]; then
|
||||
echo "🗂️ Ensuring local data directories for Milvus exist..."
|
||||
mkdir -p ./data/milvus/etcd ./data/milvus/minio ./data/milvus/milvus
|
||||
fi
|
||||
|
||||
# Function to install vector store specific packages
|
||||
install_vector_store_packages() {
|
||||
local vector_store=$1
|
||||
echo "📦 Installing packages for vector store: $vector_store..."
|
||||
|
||||
case "$vector_store" in
|
||||
qdrant)
|
||||
docker exec openmemory-openmemory-mcp-1 pip install "qdrant-client>=1.9.1" || echo "⚠️ Failed to install qdrant packages"
|
||||
;;
|
||||
chroma)
|
||||
docker exec openmemory-openmemory-mcp-1 pip install "chromadb>=0.4.24" || echo "⚠️ Failed to install chroma packages"
|
||||
;;
|
||||
weaviate)
|
||||
docker exec openmemory-openmemory-mcp-1 pip install "weaviate-client>=4.4.0,<4.15.0" || echo "⚠️ Failed to install weaviate packages"
|
||||
;;
|
||||
faiss)
|
||||
docker exec openmemory-openmemory-mcp-1 pip install "faiss-cpu>=1.7.4" || echo "⚠️ Failed to install faiss packages"
|
||||
;;
|
||||
pgvector)
|
||||
docker exec openmemory-openmemory-mcp-1 pip install "vecs>=0.4.0" "psycopg>=3.2.8" || echo "⚠️ Failed to install pgvector packages"
|
||||
;;
|
||||
redis)
|
||||
docker exec openmemory-openmemory-mcp-1 pip install "redis>=5.0.0,<6.0.0" "redisvl>=0.1.0,<1.0.0" || echo "⚠️ Failed to install redis packages"
|
||||
;;
|
||||
elasticsearch)
|
||||
docker exec openmemory-openmemory-mcp-1 pip install "elasticsearch>=8.0.0,<9.0.0" || echo "⚠️ Failed to install elasticsearch packages"
|
||||
;;
|
||||
milvus)
|
||||
docker exec openmemory-openmemory-mcp-1 pip install "pymilvus>=2.4.0,<2.6.0" || echo "⚠️ Failed to install milvus packages"
|
||||
;;
|
||||
*)
|
||||
echo "⚠️ Unknown vector store: $vector_store. Installing default qdrant packages."
|
||||
docker exec openmemory-openmemory-mcp-1 pip install "qdrant-client>=1.9.1" || echo "⚠️ Failed to install qdrant packages"
|
||||
;;
|
||||
esac
|
||||
}
|
||||
|
||||
# Start services
|
||||
echo "🚀 Starting backend services..."
|
||||
docker compose up -d
|
||||
|
||||
# Wait for container to be ready before installing packages
|
||||
echo "⏳ Waiting for container to be ready..."
|
||||
for i in {1..30}; do
|
||||
if docker exec openmemory-openmemory-mcp-1 python -c "import sys; print('ready')" >/dev/null 2>&1; then
|
||||
break
|
||||
fi
|
||||
sleep 1
|
||||
done
|
||||
|
||||
# Install vector store specific packages
|
||||
install_vector_store_packages "$VECTOR_STORE"
|
||||
|
||||
# If a specific vector store is selected, seed the backend config accordingly
|
||||
if [ "$VECTOR_STORE" = "milvus" ]; then
|
||||
echo "⏳ Waiting for API to be ready at ${NEXT_PUBLIC_API_URL}..."
|
||||
for i in {1..60}; do
|
||||
if curl -fsS "${NEXT_PUBLIC_API_URL}/api/v1/config" >/dev/null 2>&1; then
|
||||
break
|
||||
fi
|
||||
sleep 1
|
||||
done
|
||||
|
||||
echo "🧩 Configuring vector store (milvus) in backend..."
|
||||
curl -fsS -X PUT "${NEXT_PUBLIC_API_URL}/api/v1/config/mem0/vector_store" \
|
||||
-H 'Content-Type: application/json' \
|
||||
-d "{\"provider\":\"milvus\",\"config\":{\"collection_name\":\"openmemory\",\"embedding_model_dims\":${EMBEDDING_DIMS},\"url\":\"http://mem0_store:19530\",\"token\":\"\",\"db_name\":\"\",\"metric_type\":\"COSINE\"}}" >/dev/null || true
|
||||
elif [ "$VECTOR_STORE" = "weaviate" ]; then
|
||||
echo "⏳ Waiting for API to be ready at ${NEXT_PUBLIC_API_URL}..."
|
||||
for i in {1..60}; do
|
||||
if curl -fsS "${NEXT_PUBLIC_API_URL}/api/v1/config" >/dev/null 2>&1; then
|
||||
break
|
||||
fi
|
||||
sleep 1
|
||||
done
|
||||
|
||||
echo "🧩 Configuring vector store (weaviate) in backend..."
|
||||
curl -fsS -X PUT "${NEXT_PUBLIC_API_URL}/api/v1/config/mem0/vector_store" \
|
||||
-H 'Content-Type: application/json' \
|
||||
-d "{\"provider\":\"weaviate\",\"config\":{\"collection_name\":\"openmemory\",\"embedding_model_dims\":${EMBEDDING_DIMS},\"cluster_url\":\"http://mem0_store:8080\"}}" >/dev/null || true
|
||||
elif [ "$VECTOR_STORE" = "redis" ]; then
|
||||
echo "⏳ Waiting for API to be ready at ${NEXT_PUBLIC_API_URL}..."
|
||||
for i in {1..60}; do
|
||||
if curl -fsS "${NEXT_PUBLIC_API_URL}/api/v1/config" >/dev/null 2>&1; then
|
||||
break
|
||||
fi
|
||||
sleep 1
|
||||
done
|
||||
|
||||
echo "🧩 Configuring vector store (redis) in backend..."
|
||||
curl -fsS -X PUT "${NEXT_PUBLIC_API_URL}/api/v1/config/mem0/vector_store" \
|
||||
-H 'Content-Type: application/json' \
|
||||
-d "{\"provider\":\"redis\",\"config\":{\"collection_name\":\"openmemory\",\"embedding_model_dims\":${EMBEDDING_DIMS},\"redis_url\":\"redis://mem0_store:6379\"}}" >/dev/null || true
|
||||
elif [ "$VECTOR_STORE" = "pgvector" ]; then
|
||||
echo "⏳ Waiting for API to be ready at ${NEXT_PUBLIC_API_URL}..."
|
||||
for i in {1..60}; do
|
||||
if curl -fsS "${NEXT_PUBLIC_API_URL}/api/v1/config" >/dev/null 2>&1; then
|
||||
break
|
||||
fi
|
||||
sleep 1
|
||||
done
|
||||
|
||||
echo "🧩 Configuring vector store (pgvector) in backend..."
|
||||
curl -fsS -X PUT "${NEXT_PUBLIC_API_URL}/api/v1/config/mem0/vector_store" \
|
||||
-H 'Content-Type: application/json' \
|
||||
-d "{\"provider\":\"pgvector\",\"config\":{\"collection_name\":\"openmemory\",\"embedding_model_dims\":${EMBEDDING_DIMS},\"dbname\":\"mem0\",\"user\":\"mem0\",\"password\":\"mem0\",\"host\":\"mem0_store\",\"port\":5432,\"diskann\":false,\"hnsw\":true}}" >/dev/null || true
|
||||
elif [ "$VECTOR_STORE" = "qdrant" ]; then
|
||||
echo "⏳ Waiting for API to be ready at ${NEXT_PUBLIC_API_URL}..."
|
||||
for i in {1..60}; do
|
||||
if curl -fsS "${NEXT_PUBLIC_API_URL}/api/v1/config" >/dev/null 2>&1; then
|
||||
break
|
||||
fi
|
||||
sleep 1
|
||||
done
|
||||
|
||||
echo "🧩 Configuring vector store (qdrant) in backend..."
|
||||
curl -fsS -X PUT "${NEXT_PUBLIC_API_URL}/api/v1/config/mem0/vector_store" \
|
||||
-H 'Content-Type: application/json' \
|
||||
-d "{\"provider\":\"qdrant\",\"config\":{\"collection_name\":\"openmemory\",\"embedding_model_dims\":${EMBEDDING_DIMS},\"host\":\"mem0_store\",\"port\":6333}}" >/dev/null || true
|
||||
elif [ "$VECTOR_STORE" = "chroma" ]; then
|
||||
echo "⏳ Waiting for API to be ready at ${NEXT_PUBLIC_API_URL}..."
|
||||
for i in {1..60}; do
|
||||
if curl -fsS "${NEXT_PUBLIC_API_URL}/api/v1/config" >/dev/null 2>&1; then
|
||||
break
|
||||
fi
|
||||
sleep 1
|
||||
done
|
||||
|
||||
echo "🧩 Configuring vector store (chroma) in backend..."
|
||||
curl -fsS -X PUT "${NEXT_PUBLIC_API_URL}/api/v1/config/mem0/vector_store" \
|
||||
-H 'Content-Type: application/json' \
|
||||
-d "{\"provider\":\"chroma\",\"config\":{\"collection_name\":\"openmemory\",\"host\":\"mem0_store\",\"port\":8000}}" >/dev/null || true
|
||||
elif [ "$VECTOR_STORE" = "elasticsearch" ]; then
|
||||
echo "⏳ Waiting for API to be ready at ${NEXT_PUBLIC_API_URL}..."
|
||||
for i in {1..60}; do
|
||||
if curl -fsS "${NEXT_PUBLIC_API_URL}/api/v1/config" >/dev/null 2>&1; then
|
||||
break
|
||||
fi
|
||||
sleep 1
|
||||
done
|
||||
|
||||
echo "🧩 Configuring vector store (elasticsearch) in backend..."
|
||||
curl -fsS -X PUT "${NEXT_PUBLIC_API_URL}/api/v1/config/mem0/vector_store" \
|
||||
-H 'Content-Type: application/json' \
|
||||
-d "{\"provider\":\"elasticsearch\",\"config\":{\"collection_name\":\"openmemory\",\"embedding_model_dims\":${EMBEDDING_DIMS},\"host\":\"http://mem0_store\",\"port\":9200,\"user\":\"elastic\",\"password\":\"changeme\",\"verify_certs\":false,\"use_ssl\":false}}" >/dev/null || true
|
||||
elif [ "$VECTOR_STORE" = "faiss" ]; then
|
||||
echo "⏳ Waiting for API to be ready at ${NEXT_PUBLIC_API_URL}..."
|
||||
for i in {1..60}; do
|
||||
if curl -fsS "${NEXT_PUBLIC_API_URL}/api/v1/config" >/dev/null 2>&1; then
|
||||
break
|
||||
fi
|
||||
sleep 1
|
||||
done
|
||||
|
||||
echo "🧩 Configuring vector store (faiss) in backend..."
|
||||
curl -fsS -X PUT "${NEXT_PUBLIC_API_URL}/api/v1/config/mem0/vector_store" \
|
||||
-H 'Content-Type: application/json' \
|
||||
-d "{\"provider\":\"faiss\",\"config\":{\"collection_name\":\"openmemory\",\"embedding_model_dims\":${EMBEDDING_DIMS},\"path\":\"/tmp/faiss\",\"distance_strategy\":\"cosine\"}}" >/dev/null || true
|
||||
fi
|
||||
|
||||
# Start the frontend
|
||||
echo "🚀 Starting frontend on port $FRONTEND_PORT..."
|
||||
docker run -d \
|
||||
@@ -108,4 +397,4 @@ elif command -v start > /dev/null; then
|
||||
start "$URL" # Windows (if run via Git Bash or similar)
|
||||
else
|
||||
echo "⚠️ Could not detect a method to open the browser. Please open $URL manually."
|
||||
fi
|
||||
fi
|
||||
+6
-1
@@ -39,10 +39,15 @@ vector_stores = [
|
||||
"upstash-vector>=0.1.0",
|
||||
"azure-search-documents>=11.4.0b8",
|
||||
"psycopg>=3.2.8",
|
||||
"psycopg-pool>=3.2.6,<4.0.0",
|
||||
"pymongo>=4.13.2",
|
||||
"pymochow>=2.2.9",
|
||||
"databricks-sdk>=0.63.0",
|
||||
"azure-identity>=1.24.0",
|
||||
"redis>=5.0.0,<6.0.0",
|
||||
"redisvl>=0.1.0,<1.0.0",
|
||||
"elasticsearch>=8.0.0,<9.0.0",
|
||||
"pymilvus>=2.4.0,<2.6.0",
|
||||
]
|
||||
llms = [
|
||||
"groq>=0.3.0",
|
||||
@@ -58,7 +63,7 @@ extras = [
|
||||
"boto3>=1.34.0",
|
||||
"langchain-community>=0.0.0",
|
||||
"sentence-transformers>=5.0.0",
|
||||
"elasticsearch>=8.0.0",
|
||||
"elasticsearch>=8.0.0,<9.0.0",
|
||||
"opensearch-py>=2.0.0",
|
||||
"langchain-memgraph>=0.1.0",
|
||||
]
|
||||
|
||||
@@ -17,3 +17,35 @@ def test_get_update_memory_messages():
|
||||
##
|
||||
result = prompts.get_update_memory_messages(retrieved_old_memory_dict, response_content, None)
|
||||
assert result.startswith(prompts.DEFAULT_UPDATE_MEMORY_PROMPT)
|
||||
|
||||
|
||||
def test_get_update_memory_messages_empty_memory():
|
||||
# Test with None for retrieved_old_memory_dict
|
||||
result = prompts.get_update_memory_messages(
|
||||
None,
|
||||
["new fact"],
|
||||
None
|
||||
)
|
||||
assert "Current memory is empty" in result
|
||||
|
||||
# Test with empty list for retrieved_old_memory_dict
|
||||
result = prompts.get_update_memory_messages(
|
||||
[],
|
||||
["new fact"],
|
||||
None
|
||||
)
|
||||
assert "Current memory is empty" in result
|
||||
|
||||
|
||||
def test_get_update_memory_messages_non_empty_memory():
|
||||
# Non-empty memory scenario
|
||||
memory_data = [{"id": "1", "text": "existing memory"}]
|
||||
result = prompts.get_update_memory_messages(
|
||||
memory_data,
|
||||
["new fact"],
|
||||
None
|
||||
)
|
||||
# Check that the memory data is displayed
|
||||
assert str(memory_data) in result
|
||||
# And that the non-empty memory message is present
|
||||
assert "current content of my memory" in result
|
||||
|
||||
@@ -11,6 +11,7 @@ class TestKuzu:
|
||||
"alice": np.random.uniform(0.0, 0.9, 384).tolist(),
|
||||
"bob": np.random.uniform(0.0, 0.9, 384).tolist(),
|
||||
"charlie": np.random.uniform(0.0, 0.9, 384).tolist(),
|
||||
"dave": np.random.uniform(0.0, 0.9, 384).tolist(),
|
||||
}
|
||||
|
||||
@pytest.fixture
|
||||
@@ -78,6 +79,7 @@ class TestKuzu:
|
||||
assert kuzu_memory.llm == mock_llm
|
||||
assert kuzu_memory.threshold == 0.7
|
||||
|
||||
|
||||
@patch("mem0.memory.kuzu_memory.EmbedderFactory")
|
||||
@patch("mem0.memory.kuzu_memory.LlmFactory")
|
||||
def test_kuzu(self, mock_llm_factory, mock_embedder_factory, mock_config, mock_embedding_model, mock_llm):
|
||||
@@ -109,12 +111,21 @@ class TestKuzu:
|
||||
assert get_node_count(kuzu_memory) == 3
|
||||
assert get_edge_count(kuzu_memory) == 4
|
||||
|
||||
data3 = [
|
||||
{"source": "dave", "destination": "alice", "relationship": "admires"}
|
||||
]
|
||||
result = kuzu_memory._add_entities(data3, filters, {})
|
||||
assert result[0] == [{"source": "dave", "relationship": "admires", "target": "alice"}]
|
||||
assert get_node_count(kuzu_memory) == 4 # dave is new
|
||||
assert get_edge_count(kuzu_memory) == 5
|
||||
|
||||
results = kuzu_memory.get_all(filters)
|
||||
assert set([f"{result['source']}_{result['relationship']}_{result['target']}" for result in results]) == set([
|
||||
"alice_knows_bob",
|
||||
"bob_knows_charlie",
|
||||
"charlie_likes_alice",
|
||||
"charlie_knows_alice"
|
||||
"charlie_knows_alice",
|
||||
"dave_admires_alice"
|
||||
])
|
||||
|
||||
results = kuzu_memory._search_graph_db(["bob"], filters, threshold=0.8)
|
||||
@@ -125,15 +136,15 @@ class TestKuzu:
|
||||
|
||||
result = kuzu_memory._delete_entities(data2, filters)
|
||||
assert result[0] == [{"source": "charlie", "relationship": "likes", "target": "alice"}]
|
||||
assert get_node_count(kuzu_memory) == 3
|
||||
assert get_edge_count(kuzu_memory) == 3
|
||||
assert get_node_count(kuzu_memory) == 4
|
||||
assert get_edge_count(kuzu_memory) == 4
|
||||
|
||||
result = kuzu_memory._delete_entities(data1, filters)
|
||||
assert result[0] == [{"source": "alice", "relationship": "knows", "target": "bob"}]
|
||||
assert result[1] == [{"source": "bob", "relationship": "knows", "target": "charlie"}]
|
||||
assert result[2] == [{"source": "charlie", "relationship": "knows", "target": "alice"}]
|
||||
assert get_node_count(kuzu_memory) == 3
|
||||
assert get_edge_count(kuzu_memory) == 0
|
||||
assert get_node_count(kuzu_memory) == 4
|
||||
assert get_edge_count(kuzu_memory) == 1
|
||||
|
||||
result = kuzu_memory.delete_all(filters)
|
||||
assert get_node_count(kuzu_memory) == 0
|
||||
|
||||
+1247
-221
File diff suppressed because it is too large
Load Diff
@@ -1,6 +1,6 @@
|
||||
{
|
||||
"name": "@mem0/vercel-ai-provider",
|
||||
"version": "2.0.1",
|
||||
"version": "2.0.2",
|
||||
"description": "Vercel AI Provider for providing memory to LLMs",
|
||||
"main": "./dist/index.js",
|
||||
"module": "./dist/index.mjs",
|
||||
|
||||
@@ -2,17 +2,16 @@
|
||||
import {
|
||||
LanguageModelV2CallOptions,
|
||||
LanguageModelV2Message,
|
||||
LanguageModelV2Source,
|
||||
LanguageModelV2StreamPart
|
||||
LanguageModelV2Source
|
||||
} from '@ai-sdk/provider';
|
||||
|
||||
import { LanguageModelV2 } from '@ai-sdk/provider';
|
||||
import { simulateStreamingMiddleware, wrapLanguageModel } from 'ai';
|
||||
// streaming uses provider-native doStream; no middleware needed
|
||||
|
||||
import { Mem0ChatConfig, Mem0ChatModelId, Mem0ChatSettings, Mem0ConfigSettings, Mem0StreamResponse } from "./mem0-types";
|
||||
import { Mem0ClassSelector } from "./mem0-provider-selector";
|
||||
import { Mem0ProviderSettings } from "./mem0-provider";
|
||||
import { addMemories, getMemories, retrieveMemories } from "./mem0-utils";
|
||||
import { addMemories, getMemories } from "./mem0-utils";
|
||||
|
||||
const generateRandomId = () => {
|
||||
return Math.random().toString(36).substring(2, 15) + Math.random().toString(36).substring(2, 15);
|
||||
@@ -203,13 +202,8 @@ export class Mem0GenericLanguageModel implements LanguageModelV2 {
|
||||
|
||||
const baseModel = selector.createProvider();
|
||||
|
||||
// Wrap the model with streaming middleware using the new Vercel AI SDK 5.0 approach
|
||||
const model = wrapLanguageModel({
|
||||
model: baseModel,
|
||||
middleware: simulateStreamingMiddleware(),
|
||||
});
|
||||
|
||||
const streamResponse = await model.doStream({
|
||||
// Use the provider's native streaming directly to avoid buffering
|
||||
const streamResponse = await baseModel.doStream({
|
||||
...options,
|
||||
prompt: updatedPrompts,
|
||||
});
|
||||
@@ -219,65 +213,9 @@ export class Mem0GenericLanguageModel implements LanguageModelV2 {
|
||||
return streamResponse;
|
||||
}
|
||||
|
||||
// Create a new stream that includes memory sources
|
||||
const originalStream = streamResponse.stream;
|
||||
|
||||
// Create a transform stream that adds memory sources at the beginning
|
||||
const transformStream = new TransformStream({
|
||||
start(controller) {
|
||||
// Add source chunks for each memory at the beginning
|
||||
try {
|
||||
if (Array.isArray(memories) && memories?.length > 0) {
|
||||
// Create a single source that contains all memories
|
||||
controller.enqueue({
|
||||
type: 'source',
|
||||
title: "Mem0 Memories",
|
||||
sourceType: "url",
|
||||
id: "mem0-" + generateRandomId(),
|
||||
url: "https://app.mem0.ai",
|
||||
|
||||
providerOptions: {
|
||||
mem0: {
|
||||
memories: memories,
|
||||
memoriesText: memories?.map((memory: any) => memory?.memory).join("\n\n")
|
||||
}
|
||||
}
|
||||
});
|
||||
|
||||
// Also add individual memory sources for more detailed information
|
||||
memories?.forEach((memory: any) => {
|
||||
controller.enqueue({
|
||||
type: 'source',
|
||||
title: memory?.title || "Memory",
|
||||
sourceType: "url",
|
||||
id: "mem0-memory-" + generateRandomId(),
|
||||
url: "https://app.mem0.ai",
|
||||
|
||||
providerOptions: {
|
||||
mem0: {
|
||||
memory: memory,
|
||||
memoryText: memory?.memory
|
||||
}
|
||||
}
|
||||
});
|
||||
});
|
||||
}
|
||||
} catch (error) {
|
||||
console.error("Error adding memory sources:", error);
|
||||
}
|
||||
},
|
||||
transform(chunk, controller) {
|
||||
// Pass through all chunks from the original stream
|
||||
controller.enqueue(chunk);
|
||||
}
|
||||
});
|
||||
|
||||
// Pipe the original stream through our transform stream
|
||||
const enhancedStream = originalStream.pipeThrough(transformStream);
|
||||
|
||||
// Return a new stream response with our enhanced stream
|
||||
// Return stream untouched for true streaming behavior
|
||||
return {
|
||||
stream: enhancedStream,
|
||||
stream: streamResponse.stream,
|
||||
request: streamResponse.request,
|
||||
response: streamResponse.response,
|
||||
};
|
||||
|
||||
@@ -29,6 +29,7 @@ export interface Mem0ConfigSettings {
|
||||
host?: string;
|
||||
output_format?: string;
|
||||
filter_memories?: boolean;
|
||||
async_mode?: boolean;
|
||||
}
|
||||
|
||||
export interface Mem0ChatConfig extends Mem0ConfigSettings, Mem0ProviderSettings {}
|
||||
|
||||
@@ -62,26 +62,26 @@ const convertToMem0Format = (messages: LanguageModelV2Prompt) => {
|
||||
|
||||
const searchInternalMemories = async (query: string, config?: Mem0ConfigSettings, top_k: number = 5) => {
|
||||
try {
|
||||
const filters: { AND: Array<{ [key: string]: string | undefined }> } = {
|
||||
AND: [],
|
||||
const filters: { OR: Array<{ [key: string]: string | undefined }> } = {
|
||||
OR: [],
|
||||
};
|
||||
if (config?.user_id) {
|
||||
filters.AND.push({
|
||||
filters.OR.push({
|
||||
user_id: config.user_id,
|
||||
});
|
||||
}
|
||||
if (config?.app_id) {
|
||||
filters.AND.push({
|
||||
filters.OR.push({
|
||||
app_id: config.app_id,
|
||||
});
|
||||
}
|
||||
if (config?.agent_id) {
|
||||
filters.AND.push({
|
||||
filters.OR.push({
|
||||
agent_id: config.agent_id,
|
||||
});
|
||||
}
|
||||
if (config?.run_id) {
|
||||
filters.AND.push({
|
||||
filters.OR.push({
|
||||
run_id: config.run_id,
|
||||
});
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user