diff --git a/docs/platform/advanced-memory-operations.mdx b/docs/platform/advanced-memory-operations.mdx index 8e5c5683d..15061cbc6 100644 --- a/docs/platform/advanced-memory-operations.mdx +++ b/docs/platform/advanced-memory-operations.mdx @@ -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. diff --git a/mem0/configs/embeddings/base.py b/mem0/configs/embeddings/base.py index a5a174b7b..0737088fc 100644 --- a/mem0/configs/embeddings/base.py +++ b/mem0/configs/embeddings/base.py @@ -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" + diff --git a/mem0/configs/llms/aws_bedrock.py b/mem0/configs/llms/aws_bedrock.py index bbdebef0a..a285f9074 100644 --- a/mem0/configs/llms/aws_bedrock.py +++ b/mem0/configs/llms/aws_bedrock.py @@ -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 {} diff --git a/mem0/configs/prompts.py b/mem0/configs/prompts.py index b8daecfd6..fbfbe7f6f 100644 --- a/mem0/configs/prompts.py +++ b/mem0/configs/prompts.py @@ -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. ``` diff --git a/mem0/configs/vector_stores/pgvector.py b/mem0/configs/vector_stores/pgvector.py index b0cca30a9..66c331d3d 100644 --- a/mem0/configs/vector_stores/pgvector.py +++ b/mem0/configs/vector_stores/pgvector.py @@ -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") diff --git a/mem0/embeddings/aws_bedrock.py b/mem0/embeddings/aws_bedrock.py index d5ce21db6..5c3c1acd1 100644 --- a/mem0/embeddings/aws_bedrock.py +++ b/mem0/embeddings/aws_bedrock.py @@ -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", diff --git a/mem0/memory/kuzu_memory.py b/mem0/memory/kuzu_memory.py index acb5bb8dc..828e2e174 100644 --- a/mem0/memory/kuzu_memory.py +++ b/mem0/memory/kuzu_memory.py @@ -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 diff --git a/mem0/vector_stores/mongodb.py b/mem0/vector_stores/mongodb.py index 381bfde4e..01cb17bb8 100644 --- a/mem0/vector_stores/mongodb.py +++ b/mem0/vector_stores/mongodb.py @@ -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}'.") diff --git a/mem0/vector_stores/pgvector.py b/mem0/vector_stores/pgvector.py index 9906089bc..e2d020a66 100644 --- a/mem0/vector_stores/pgvector.py +++ b/mem0/vector_stores/pgvector.py @@ -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() diff --git a/mem0/vector_stores/weaviate.py b/mem0/vector_stores/weaviate.py index d5b2fa417..989cc49b1 100644 --- a/mem0/vector_stores/weaviate.py +++ b/mem0/vector_stores/weaviate.py @@ -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 diff --git a/openmemory/api/app/mcp_server.py b/openmemory/api/app/mcp_server.py index 914eb15b3..0911a0dac 100644 --- a/openmemory/api/app/mcp_server.py +++ b/openmemory/api/app/mcp_server.py @@ -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: diff --git a/openmemory/api/app/routers/memories.py b/openmemory/api/app/routers/memories.py index f73aa702b..a19309953 100644 --- a/openmemory/api/app/routers/memories.py +++ b/openmemory/api/app/routers/memories.py @@ -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 diff --git a/openmemory/api/app/utils/memory.py b/openmemory/api/app/utils/memory.py index 921e4dffb..a4f557fe6 100644 --- a/openmemory/api/app/utils/memory.py +++ b/openmemory/api/app/utils/memory.py @@ -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") diff --git a/openmemory/compose/chroma.yml b/openmemory/compose/chroma.yml new file mode 100644 index 000000000..3961fff2c --- /dev/null +++ b/openmemory/compose/chroma.yml @@ -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 \ No newline at end of file diff --git a/openmemory/compose/elasticsearch.yml b/openmemory/compose/elasticsearch.yml new file mode 100644 index 000000000..22e909e95 --- /dev/null +++ b/openmemory/compose/elasticsearch.yml @@ -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 \ No newline at end of file diff --git a/openmemory/compose/faiss.yml b/openmemory/compose/faiss.yml new file mode 100644 index 000000000..3b2c94173 --- /dev/null +++ b/openmemory/compose/faiss.yml @@ -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 diff --git a/openmemory/compose/milvus.yml b/openmemory/compose/milvus.yml new file mode 100644 index 000000000..013ba1e32 --- /dev/null +++ b/openmemory/compose/milvus.yml @@ -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 \ No newline at end of file diff --git a/openmemory/compose/opensearch.yml b/openmemory/compose/opensearch.yml new file mode 100644 index 000000000..8b2e7100e --- /dev/null +++ b/openmemory/compose/opensearch.yml @@ -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 \ No newline at end of file diff --git a/openmemory/compose/pgvector.yml b/openmemory/compose/pgvector.yml new file mode 100644 index 000000000..8cb956008 --- /dev/null +++ b/openmemory/compose/pgvector.yml @@ -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 \ No newline at end of file diff --git a/openmemory/compose/qdrant.yml b/openmemory/compose/qdrant.yml new file mode 100644 index 000000000..422f4a4d8 --- /dev/null +++ b/openmemory/compose/qdrant.yml @@ -0,0 +1,8 @@ +services: + mem0_store: + image: qdrant/qdrant:latest + restart: unless-stopped + ports: + - "6333:6333" + volumes: + - mem0_storage:/mem0/storage \ No newline at end of file diff --git a/openmemory/compose/redis.yml b/openmemory/compose/redis.yml new file mode 100644 index 000000000..1ca2ef803 --- /dev/null +++ b/openmemory/compose/redis.yml @@ -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 \ No newline at end of file diff --git a/openmemory/compose/weaviate.yml b/openmemory/compose/weaviate.yml new file mode 100644 index 000000000..6eab1b8bc --- /dev/null +++ b/openmemory/compose/weaviate.yml @@ -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 \ No newline at end of file diff --git a/openmemory/run.sh b/openmemory/run.sh index 029b7d4b9..ca15322e7 100644 --- a/openmemory/run.sh +++ b/openmemory/run.sh @@ -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 < 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 <> docker-compose.yml <> docker-compose.yml <> docker-compose.yml <> docker-compose.yml <> docker-compose.yml <> docker-compose.yml <> docker-compose.yml <> docker-compose.yml <> docker-compose.yml <> docker-compose.yml <> docker-compose.yml <=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 \ No newline at end of file diff --git a/pyproject.toml b/pyproject.toml index 8c07ac64f..f47525658 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -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", ] diff --git a/tests/configs/test_prompts.py b/tests/configs/test_prompts.py index e978f8c9f..1376ab181 100644 --- a/tests/configs/test_prompts.py +++ b/tests/configs/test_prompts.py @@ -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 diff --git a/tests/memory/test_kuzu.py b/tests/memory/test_kuzu.py index f2e429527..334912add 100644 --- a/tests/memory/test_kuzu.py +++ b/tests/memory/test_kuzu.py @@ -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 diff --git a/tests/vector_stores/test_pgvector.py b/tests/vector_stores/test_pgvector.py index 2a66559b5..436c9708c 100644 --- a/tests/vector_stores/test_pgvector.py +++ b/tests/vector_stores/test_pgvector.py @@ -1,3 +1,5 @@ +import importlib +import sys import unittest import uuid from unittest.mock import MagicMock, patch @@ -13,9 +15,15 @@ class TestPGVector(unittest.TestCase): self.mock_conn.cursor.return_value = self.mock_cursor # Mock connection pool - self.mock_pool = MagicMock() - self.mock_pool.getconn.return_value = self.mock_conn + self.mock_pool_psycopg2 = MagicMock() + self.mock_pool_psycopg2.getconn.return_value = self.mock_conn + + self.mock_pool_psycopg = MagicMock() + self.mock_pool_psycopg.connection.return_value = self.mock_conn + self.mock_get_cursor = MagicMock() + self.mock_get_cursor.return_value = self.mock_cursor + # Mock connection string self.connection_string = "postgresql://user:pass@host:5432/db" @@ -25,12 +33,11 @@ class TestPGVector(unittest.TestCase): self.test_ids = [str(uuid.uuid4()), str(uuid.uuid4())] @patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 3) - @patch('mem0.vector_stores.pgvector.psycopg.connect') - @patch('mem0.vector_stores.pgvector.psycopg2.connect') - def test_init_with_individual_params_psycopg3(self, mock_psycopg2_connect, mock_psycopg_connect): + @patch('mem0.vector_stores.pgvector.ConnectionPool') + def test_init_with_individual_params_psycopg3(self, mock_psycopg_pool): """Test initialization with individual parameters using psycopg3.""" # Mock psycopg3 to be available - mock_psycopg_connect.return_value = self.mock_conn + mock_psycopg_pool.return_value = self.mock_pool_psycopg self.mock_cursor.fetchall.return_value = [] # No existing collections pgvector = PGVector( @@ -42,24 +49,25 @@ class TestPGVector(unittest.TestCase): host="localhost", port=5432, diskann=False, - hnsw=False + hnsw=False, + minconn=1, + maxconn=4, ) - - mock_psycopg_connect.assert_called_once_with( - dbname="test_db", - user="test_user", - password="test_pass", - host="localhost", - port=5432 + + mock_psycopg_pool.assert_called_once_with( + conninfo="postgresql://test_user:test_pass@localhost:5432/test_db", + min_size=1, + max_size=4, + open=True, ) self.assertEqual(pgvector.collection_name, "test_collection") self.assertEqual(pgvector.embedding_model_dims, 3) @patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 2) - @patch('mem0.vector_stores.pgvector.psycopg2.connect') - def test_init_with_individual_params_psycopg2(self, mock_connect): + @patch('mem0.vector_stores.pgvector.ConnectionPool') + def test_init_with_individual_params_psycopg2(self, mock_pcycopg2_pool): """Test initialization with individual parameters using psycopg2.""" - mock_connect.return_value = self.mock_conn + mock_pcycopg2_pool.return_value = self.mock_pool_psycopg2 self.mock_cursor.fetchall.return_value = [] # No existing collections pgvector = PGVector( @@ -71,26 +79,34 @@ class TestPGVector(unittest.TestCase): host="localhost", port=5432, diskann=False, - hnsw=False + hnsw=False, + minconn=1, + maxconn=4, ) - mock_connect.assert_called_once_with( - dbname="test_db", - user="test_user", - password="test_pass", - host="localhost", - port=5432 + mock_pcycopg2_pool.assert_called_once_with( + minconn=1, + maxconn=4, + dsn="postgresql://test_user:test_pass@localhost:5432/test_db", ) + self.assertEqual(pgvector.collection_name, "test_collection") self.assertEqual(pgvector.embedding_model_dims, 3) @patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 3) - @patch('mem0.vector_stores.pgvector.psycopg.connect') - @patch('mem0.vector_stores.pgvector.psycopg2.connect') - def test_create_col_psycopg3(self, mock_psycopg2_connect, mock_psycopg_connect): + @patch('mem0.vector_stores.pgvector.ConnectionPool') + @patch.object(PGVector, '_get_cursor') + def test_create_col_psycopg3(self, mock_get_cursor, mock_connection_pool): """Test collection creation with psycopg3.""" - mock_psycopg_connect.return_value = self.mock_conn - self.mock_cursor.fetchall.return_value = [] + # Set up mock pool and cursor + mock_pool = MagicMock() + mock_connection_pool.return_value = mock_pool + + # Configure the _get_cursor mock to return our mock cursor + mock_get_cursor.return_value.__enter__.return_value = self.mock_cursor + mock_get_cursor.return_value.__exit__.return_value = None + + self.mock_cursor.fetchall.return_value = [] # No existing collections pgvector = PGVector( dbname="test_db", @@ -101,26 +117,145 @@ class TestPGVector(unittest.TestCase): host="localhost", port=5432, diskann=False, - hnsw=False + hnsw=False, + minconn=1, + maxconn=4 ) + # Verify the _get_cursor context manager was called + mock_get_cursor.assert_called() + # Verify vector extension and table creation self.mock_cursor.execute.assert_any_call("CREATE EXTENSION IF NOT EXISTS vector") table_creation_calls = [call for call in self.mock_cursor.execute.call_args_list if "CREATE TABLE IF NOT EXISTS test_collection" in str(call)] self.assertTrue(len(table_creation_calls) > 0) - self.mock_conn.commit.assert_called() # Verify pgvector instance properties self.assertEqual(pgvector.collection_name, "test_collection") self.assertEqual(pgvector.embedding_model_dims, 3) + @patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 3) + @patch('mem0.vector_stores.pgvector.ConnectionPool') + @patch.object(PGVector, '_get_cursor') + def test_create_col_psycopg3_with_explicit_pool(self, mock_get_cursor, mock_connection_pool): + """ + Test collection creation with psycopg3 when an explicit psycopg_pool.ConnectionPool is provided. + This ensures that PGVector uses the provided pool and still performs collection creation logic. + """ + # Set up a real (mocked) psycopg_pool.ConnectionPool instance + explicit_pool = MagicMock(name="ExplicitPsycopgPool") + # The patch for ConnectionPool should not be used in this case, but we patch it for isolation + mock_connection_pool.return_value = MagicMock(name="ShouldNotBeUsed") + + # Configure the _get_cursor mock to return our mock cursor as a context manager + mock_get_cursor.return_value.__enter__.return_value = self.mock_cursor + mock_get_cursor.return_value.__exit__.return_value = None + + # Simulate no existing collections in the database + self.mock_cursor.fetchall.return_value = [] + + # Pass the explicit pool to PGVector + pgvector = PGVector( + dbname="test_db", + collection_name="test_collection", + embedding_model_dims=3, + user="test_user", + password="test_pass", + host="localhost", + port=5432, + diskann=False, + hnsw=False, + minconn=1, + maxconn=4, + connection_pool=explicit_pool + ) + + # Verify the _get_cursor context manager was called + mock_get_cursor.assert_called() + + mock_connection_pool.assert_not_called() + + + # Verify vector extension and table creation + self.mock_cursor.execute.assert_any_call("CREATE EXTENSION IF NOT EXISTS vector") + table_creation_calls = [call for call in self.mock_cursor.execute.call_args_list + if "CREATE TABLE IF NOT EXISTS test_collection" in str(call)] + self.assertTrue(len(table_creation_calls) > 0) + + # Verify pgvector instance properties + self.assertEqual(pgvector.collection_name, "test_collection") + self.assertEqual(pgvector.embedding_model_dims, 3) + # Ensure the pool used is the explicit one + self.assertIs(pgvector.connection_pool, explicit_pool) + @patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 2) - @patch('mem0.vector_stores.pgvector.psycopg2.connect') - def test_create_col_psycopg2(self, mock_connect): + @patch('mem0.vector_stores.pgvector.ConnectionPool') + @patch.object(PGVector, '_get_cursor') + def test_create_col_psycopg2_with_explicit_pool(self, mock_get_cursor, mock_connection_pool): + """ + Test collection creation with psycopg2 when an explicit psycopg2 ThreadedConnectionPool is provided. + This ensures that PGVector uses the provided pool and still performs collection creation logic. + """ + # Set up a real (mocked) psycopg2 ThreadedConnectionPool instance + explicit_pool = MagicMock(name="ExplicitPsycopg2Pool") + # The patch for ConnectionPool should not be used in this case, but we patch it for isolation + mock_connection_pool.return_value = MagicMock(name="ShouldNotBeUsed") + + # Configure the _get_cursor mock to return our mock cursor as a context manager + mock_get_cursor.return_value.__enter__.return_value = self.mock_cursor + mock_get_cursor.return_value.__exit__.return_value = None + + # Simulate no existing collections in the database + self.mock_cursor.fetchall.return_value = [] + + # Pass the explicit pool to PGVector + pgvector = PGVector( + dbname="test_db", + collection_name="test_collection", + embedding_model_dims=3, + user="test_user", + password="test_pass", + host="localhost", + port=5432, + diskann=False, + hnsw=False, + minconn=1, + maxconn=4, + connection_pool=explicit_pool + ) + + # Verify the _get_cursor context manager was called + mock_get_cursor.assert_called() + + mock_connection_pool.assert_not_called() + + # Verify vector extension and table creation + self.mock_cursor.execute.assert_any_call("CREATE EXTENSION IF NOT EXISTS vector") + table_creation_calls = [call for call in self.mock_cursor.execute.call_args_list + if "CREATE TABLE IF NOT EXISTS test_collection" in str(call)] + self.assertTrue(len(table_creation_calls) > 0) + + # Verify pgvector instance properties + self.assertEqual(pgvector.collection_name, "test_collection") + self.assertEqual(pgvector.embedding_model_dims, 3) + # Ensure the pool used is the explicit one + self.assertIs(pgvector.connection_pool, explicit_pool) + + @patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 2) + @patch('mem0.vector_stores.pgvector.ConnectionPool') + @patch.object(PGVector, '_get_cursor') + def test_create_col_psycopg2(self, mock_get_cursor, mock_connection_pool): """Test collection creation with psycopg2.""" - mock_connect.return_value = self.mock_conn - self.mock_cursor.fetchall.return_value = [] + # Set up mock pool and cursor + mock_pool = MagicMock() + mock_connection_pool.return_value = mock_pool + + # Configure the _get_cursor mock to return our mock cursor + mock_get_cursor.return_value.__enter__.return_value = self.mock_cursor + mock_get_cursor.return_value.__exit__.return_value = None + + self.mock_cursor.fetchall.return_value = [] # No existing collections pgvector = PGVector( dbname="test_db", @@ -131,28 +266,37 @@ class TestPGVector(unittest.TestCase): host="localhost", port=5432, diskann=False, - hnsw=False + hnsw=False, + minconn=1, + maxconn=4 ) + # Verify the _get_cursor context manager was called + mock_get_cursor.assert_called() + # Verify vector extension and table creation self.mock_cursor.execute.assert_any_call("CREATE EXTENSION IF NOT EXISTS vector") table_creation_calls = [call for call in self.mock_cursor.execute.call_args_list if "CREATE TABLE IF NOT EXISTS test_collection" in str(call)] self.assertTrue(len(table_creation_calls) > 0) - self.mock_conn.commit.assert_called() # Verify pgvector instance properties self.assertEqual(pgvector.collection_name, "test_collection") self.assertEqual(pgvector.embedding_model_dims, 3) @patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 3) - @patch('mem0.vector_stores.pgvector.psycopg.connect') - @patch('mem0.vector_stores.pgvector.psycopg2.connect') - @patch('mem0.vector_stores.pgvector.execute_values') - def test_insert_psycopg3(self, mock_execute_values, mock_psycopg2_connect, mock_psycopg_connect): + @patch('mem0.vector_stores.pgvector.ConnectionPool') + @patch.object(PGVector, '_get_cursor') + def test_insert_psycopg3(self, mock_get_cursor, mock_connection_pool): """Test vector insertion with psycopg3.""" - mock_psycopg_connect.return_value = self.mock_conn - self.mock_cursor.fetchall.return_value = [] + # Set up mock pool and cursor + mock_connection_pool.return_value = self.mock_pool_psycopg + + # Configure the _get_cursor mock to return our mock cursor + mock_get_cursor.return_value.__enter__.return_value = self.mock_cursor + mock_get_cursor.return_value.__exit__.return_value = None + + self.mock_cursor.fetchall.return_value = [] # No existing collections pgvector = PGVector( dbname="test_db", @@ -163,61 +307,112 @@ class TestPGVector(unittest.TestCase): host="localhost", port=5432, diskann=False, - hnsw=False + hnsw=False, + minconn=1, + maxconn=4 ) pgvector.insert(self.test_vectors, self.test_payloads, self.test_ids) - # Verify execute_values was called - mock_execute_values.assert_called_once() - call_args = mock_execute_values.call_args - self.assertIn("INSERT INTO test_collection", call_args[0][1]) + # Verify the _get_cursor context manager was called + mock_get_cursor.assert_called() + + # Verify insert query was executed (psycopg3 uses executemany) + insert_calls = [call for call in self.mock_cursor.executemany.call_args_list + if "INSERT INTO test_collection" in str(call)] + self.assertTrue(len(insert_calls) > 0) # Verify data format - data_arg = call_args[0][2] + call_args = self.mock_cursor.executemany.call_args + data_arg = call_args[0][1] self.assertEqual(len(data_arg), 2) self.assertEqual(data_arg[0][0], self.test_ids[0]) self.assertEqual(data_arg[1][0], self.test_ids[1]) @patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 2) - @patch('mem0.vector_stores.pgvector.psycopg2.connect') - @patch('mem0.vector_stores.pgvector.execute_values') - def test_insert_psycopg2(self, mock_execute_values, mock_connect): - """Test vector insertion with psycopg2.""" - mock_connect.return_value = self.mock_conn - self.mock_cursor.fetchall.return_value = [] - - pgvector = PGVector( - dbname="test_db", - collection_name="test_collection", - embedding_model_dims=3, - user="test_user", - password="test_pass", - host="localhost", - port=5432, - diskann=False, - hnsw=False - ) - - pgvector.insert(self.test_vectors, self.test_payloads, self.test_ids) - - # Verify execute_values was called - mock_execute_values.assert_called_once() - call_args = mock_execute_values.call_args - self.assertIn("INSERT INTO test_collection", call_args[0][1]) - - # Verify data format - data_arg = call_args[0][2] - self.assertEqual(len(data_arg), 2) - self.assertEqual(data_arg[0][0], self.test_ids[0]) - self.assertEqual(data_arg[1][0], self.test_ids[1]) + @patch('mem0.vector_stores.pgvector.ConnectionPool') + @patch.object(PGVector, '_get_cursor') + def test_insert_psycopg2(self, mock_get_cursor, mock_connection_pool): + """ + Test vector insertion with psycopg2. + This test ensures that PGVector.insert uses psycopg2.extras.execute_values for batch inserts + and that the data passed to execute_values is correctly formatted. + """ + # --- Setup mocks for psycopg2 and its submodules --- + mock_execute_values = MagicMock() + mock_pool = MagicMock() + + # Mock psycopg2.extras with execute_values + mock_psycopg2_extras = MagicMock() + mock_psycopg2_extras.execute_values = mock_execute_values + + mock_psycopg2_pool = MagicMock() + mock_psycopg2_pool.ThreadedConnectionPool = mock_pool + + # Mock psycopg2 root module + mock_psycopg2 = MagicMock() + mock_psycopg2.extras = mock_psycopg2_extras + mock_psycopg2.pool = mock_psycopg2_pool + + # Patch sys.modules so that imports in PGVector use our mocks + with patch.dict('sys.modules', { + 'psycopg': None, # Ensure psycopg3 is not available + 'psycopg_pool': None, + 'psycopg.types.json': None, + 'psycopg2': mock_psycopg2, + 'psycopg2.extras': mock_psycopg2_extras, + 'psycopg2.pool': mock_psycopg2_pool + }): + # Force reload of PGVector to pick up the mocked modules + if 'mem0.vector_stores.pgvector' in sys.modules: + importlib.reload(sys.modules['mem0.vector_stores.pgvector']) + + mock_connection_pool.return_value = self.mock_pool_psycopg + mock_get_cursor.return_value.__enter__.return_value = self.mock_cursor + mock_get_cursor.return_value.__exit__.return_value = None + self.mock_cursor.fetchall.return_value = [] + + pgvector = PGVector( + dbname="test_db", + collection_name="test_collection", + embedding_model_dims=3, + user="test_user", + password="test_pass", + host="localhost", + port=5432, + diskann=False, + hnsw=False, + minconn=1, + maxconn=4 + ) + + pgvector.insert(self.test_vectors, self.test_payloads, self.test_ids) + + mock_get_cursor.assert_called() + mock_execute_values.assert_called_once() + call_args = mock_execute_values.call_args + + self.assertIn("INSERT INTO test_collection", call_args[0][1]) + + # The data argument should be a list of tuples, one per vector + data_arg = call_args[0][2] + self.assertEqual(len(data_arg), 2) + self.assertEqual(data_arg[0][0], self.test_ids[0]) + self.assertEqual(data_arg[1][0], self.test_ids[1]) @patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 3) - @patch('mem0.vector_stores.pgvector.psycopg.connect') - @patch('mem0.vector_stores.pgvector.psycopg2.connect') - def test_search_psycopg3(self, mock_psycopg2_connect, mock_psycopg_connect): + @patch('mem0.vector_stores.pgvector.ConnectionPool') + @patch.object(PGVector, '_get_cursor') + def test_search_psycopg3(self, mock_get_cursor, mock_connection_pool): """Test search with psycopg3.""" - mock_psycopg_connect.return_value = self.mock_conn + # Set up mock pool and cursor + mock_pool = MagicMock() + mock_connection_pool.return_value = mock_pool + + # Configure the _get_cursor mock to return our mock cursor + mock_get_cursor.return_value.__enter__.return_value = self.mock_cursor + mock_get_cursor.return_value.__exit__.return_value = None + self.mock_cursor.fetchall.return_value = [ (self.test_ids[0], 0.1, {"key": "value1"}), (self.test_ids[1], 0.2, {"key": "value2"}), @@ -232,11 +427,16 @@ class TestPGVector(unittest.TestCase): host="localhost", port=5432, diskann=False, - hnsw=False + hnsw=False, + minconn=1, + maxconn=4 ) results = pgvector.search("test query", [0.1, 0.2, 0.3], limit=2) + # Verify the _get_cursor context manager was called + mock_get_cursor.assert_called() + # Verify search query was executed search_calls = [call for call in self.mock_cursor.execute.call_args_list if "SELECT id, vector <=" in str(call)] @@ -250,10 +450,18 @@ class TestPGVector(unittest.TestCase): self.assertEqual(results[1].score, 0.2) @patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 2) - @patch('mem0.vector_stores.pgvector.psycopg2.connect') - def test_search_psycopg2(self, mock_connect): + @patch('mem0.vector_stores.pgvector.ConnectionPool') + @patch.object(PGVector, '_get_cursor') + def test_search_psycopg2(self, mock_get_cursor, mock_connection_pool): """Test search with psycopg2.""" - mock_connect.return_value = self.mock_conn + # Set up mock pool and cursor + mock_pool = MagicMock() + mock_connection_pool.return_value = mock_pool + + # Configure the _get_cursor mock to return our mock cursor + mock_get_cursor.return_value.__enter__.return_value = self.mock_cursor + mock_get_cursor.return_value.__exit__.return_value = None + self.mock_cursor.fetchall.return_value = [ (self.test_ids[0], 0.1, {"key": "value1"}), (self.test_ids[1], 0.2, {"key": "value2"}), @@ -268,11 +476,16 @@ class TestPGVector(unittest.TestCase): host="localhost", port=5432, diskann=False, - hnsw=False + hnsw=False, + minconn=1, + maxconn=4 ) results = pgvector.search("test query", [0.1, 0.2, 0.3], limit=2) + # Verify the _get_cursor context manager was called + mock_get_cursor.assert_called() + # Verify search query was executed search_calls = [call for call in self.mock_cursor.execute.call_args_list if "SELECT id, vector <=" in str(call)] @@ -286,11 +499,19 @@ class TestPGVector(unittest.TestCase): self.assertEqual(results[1].score, 0.2) @patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 3) - @patch('mem0.vector_stores.pgvector.psycopg.connect') - @patch('mem0.vector_stores.pgvector.psycopg2.connect') - def test_delete_psycopg3(self, mock_psycopg2_connect, mock_psycopg_connect): + @patch('mem0.vector_stores.pgvector.ConnectionPool') + @patch.object(PGVector, '_get_cursor') + def test_delete_psycopg3(self, mock_get_cursor, mock_connection_pool): """Test delete with psycopg3.""" - mock_psycopg_connect.return_value = self.mock_conn + # Set up mock pool and cursor + mock_pool = MagicMock() + mock_connection_pool.return_value = mock_pool + + # Configure the _get_cursor mock to return our mock cursor + mock_get_cursor.return_value.__enter__.return_value = self.mock_cursor + mock_get_cursor.return_value.__exit__.return_value = None + + self.mock_cursor.fetchall.return_value = [] # No existing collections pgvector = PGVector( dbname="test_db", @@ -301,22 +522,35 @@ class TestPGVector(unittest.TestCase): host="localhost", port=5432, diskann=False, - hnsw=False + hnsw=False, + minconn=1, + maxconn=4 ) pgvector.delete(self.test_ids[0]) + # Verify the _get_cursor context manager was called + mock_get_cursor.assert_called() + # Verify delete query was executed delete_calls = [call for call in self.mock_cursor.execute.call_args_list if "DELETE FROM test_collection" in str(call)] self.assertTrue(len(delete_calls) > 0) - self.mock_conn.commit.assert_called() @patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 2) - @patch('mem0.vector_stores.pgvector.psycopg2.connect') - def test_delete_psycopg2(self, mock_connect): + @patch('mem0.vector_stores.pgvector.ConnectionPool') + @patch.object(PGVector, '_get_cursor') + def test_delete_psycopg2(self, mock_get_cursor, mock_connection_pool): """Test delete with psycopg2.""" - mock_connect.return_value = self.mock_conn + # Set up mock pool and cursor + mock_pool = MagicMock() + mock_connection_pool.return_value = mock_pool + + # Configure the _get_cursor mock to return our mock cursor + mock_get_cursor.return_value.__enter__.return_value = self.mock_cursor + mock_get_cursor.return_value.__exit__.return_value = None + + self.mock_cursor.fetchall.return_value = [] # No existing collections pgvector = PGVector( dbname="test_db", @@ -327,23 +561,35 @@ class TestPGVector(unittest.TestCase): host="localhost", port=5432, diskann=False, - hnsw=False + hnsw=False, + minconn=1, + maxconn=4 ) pgvector.delete(self.test_ids[0]) + # Verify the _get_cursor context manager was called + mock_get_cursor.assert_called() + # Verify delete query was executed delete_calls = [call for call in self.mock_cursor.execute.call_args_list if "DELETE FROM test_collection" in str(call)] self.assertTrue(len(delete_calls) > 0) - self.mock_conn.commit.assert_called() @patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 3) - @patch('mem0.vector_stores.pgvector.psycopg.connect') - @patch('mem0.vector_stores.pgvector.psycopg2.connect') - def test_update_psycopg3(self, mock_psycopg2_connect, mock_psycopg_connect): + @patch('mem0.vector_stores.pgvector.ConnectionPool') + @patch.object(PGVector, '_get_cursor') + def test_update_psycopg3(self, mock_get_cursor, mock_connection_pool): """Test update with psycopg3.""" - mock_psycopg_connect.return_value = self.mock_conn + # Set up mock pool and cursor + mock_pool = MagicMock() + mock_connection_pool.return_value = mock_pool + + # Configure the _get_cursor mock to return our mock cursor + mock_get_cursor.return_value.__enter__.return_value = self.mock_cursor + mock_get_cursor.return_value.__exit__.return_value = None + + self.mock_cursor.fetchall.return_value = [] # No existing collections pgvector = PGVector( dbname="test_db", @@ -354,7 +600,9 @@ class TestPGVector(unittest.TestCase): host="localhost", port=5432, diskann=False, - hnsw=False + hnsw=False, + minconn=1, + maxconn=4 ) updated_vector = [0.5, 0.6, 0.7] @@ -362,17 +610,28 @@ class TestPGVector(unittest.TestCase): pgvector.update(self.test_ids[0], vector=updated_vector, payload=updated_payload) + # Verify the _get_cursor context manager was called + mock_get_cursor.assert_called() + # Verify update queries were executed update_calls = [call for call in self.mock_cursor.execute.call_args_list if "UPDATE test_collection" in str(call)] self.assertTrue(len(update_calls) > 0) - self.mock_conn.commit.assert_called() @patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 2) - @patch('mem0.vector_stores.pgvector.psycopg2.connect') - def test_update_psycopg2(self, mock_connect): + @patch('mem0.vector_stores.pgvector.ConnectionPool') + @patch.object(PGVector, '_get_cursor') + def test_update_psycopg2(self, mock_get_cursor, mock_connection_pool): """Test update with psycopg2.""" - mock_connect.return_value = self.mock_conn + # Set up mock pool and cursor + mock_pool = MagicMock() + mock_connection_pool.return_value = mock_pool + + # Configure the _get_cursor mock to return our mock cursor + mock_get_cursor.return_value.__enter__.return_value = self.mock_cursor + mock_get_cursor.return_value.__exit__.return_value = None + + self.mock_cursor.fetchall.return_value = [] # No existing collections pgvector = PGVector( dbname="test_db", @@ -383,7 +642,9 @@ class TestPGVector(unittest.TestCase): host="localhost", port=5432, diskann=False, - hnsw=False + hnsw=False, + minconn=1, + maxconn=4 ) updated_vector = [0.5, 0.6, 0.7] @@ -391,18 +652,28 @@ class TestPGVector(unittest.TestCase): pgvector.update(self.test_ids[0], vector=updated_vector, payload=updated_payload) + # Verify the _get_cursor context manager was called + mock_get_cursor.assert_called() + # Verify update queries were executed update_calls = [call for call in self.mock_cursor.execute.call_args_list if "UPDATE test_collection" in str(call)] self.assertTrue(len(update_calls) > 0) - self.mock_conn.commit.assert_called() @patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 3) - @patch('mem0.vector_stores.pgvector.psycopg.connect') - @patch('mem0.vector_stores.pgvector.psycopg2.connect') - def test_get_psycopg3(self, mock_psycopg2_connect, mock_psycopg_connect): + @patch('mem0.vector_stores.pgvector.ConnectionPool') + @patch.object(PGVector, '_get_cursor') + def test_get_psycopg3(self, mock_get_cursor, mock_connection_pool): """Test get with psycopg3.""" - mock_psycopg_connect.return_value = self.mock_conn + # Set up mock pool and cursor + mock_pool = MagicMock() + mock_connection_pool.return_value = mock_pool + + # Configure the _get_cursor mock to return our mock cursor + mock_get_cursor.return_value.__enter__.return_value = self.mock_cursor + mock_get_cursor.return_value.__exit__.return_value = None + + self.mock_cursor.fetchall.return_value = [] # No existing collections self.mock_cursor.fetchone.return_value = (self.test_ids[0], [0.1, 0.2, 0.3], {"key": "value1"}) pgvector = PGVector( @@ -414,11 +685,16 @@ class TestPGVector(unittest.TestCase): host="localhost", port=5432, diskann=False, - hnsw=False + hnsw=False, + minconn=1, + maxconn=4 ) result = pgvector.get(self.test_ids[0]) + # Verify the _get_cursor context manager was called + mock_get_cursor.assert_called() + # Verify get query was executed get_calls = [call for call in self.mock_cursor.execute.call_args_list if "SELECT id, vector, payload" in str(call)] @@ -430,10 +706,19 @@ class TestPGVector(unittest.TestCase): self.assertEqual(result.payload, {"key": "value1"}) @patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 2) - @patch('mem0.vector_stores.pgvector.psycopg2.connect') - def test_get_psycopg2(self, mock_connect): + @patch('mem0.vector_stores.pgvector.ConnectionPool') + @patch.object(PGVector, '_get_cursor') + def test_get_psycopg2(self, mock_get_cursor, mock_connection_pool): """Test get with psycopg2.""" - mock_connect.return_value = self.mock_conn + # Set up mock pool and cursor + mock_pool = MagicMock() + mock_connection_pool.return_value = mock_pool + + # Configure the _get_cursor mock to return our mock cursor + mock_get_cursor.return_value.__enter__.return_value = self.mock_cursor + mock_get_cursor.return_value.__exit__.return_value = None + + self.mock_cursor.fetchall.return_value = [] # No existing collections self.mock_cursor.fetchone.return_value = (self.test_ids[0], [0.1, 0.2, 0.3], {"key": "value1"}) pgvector = PGVector( @@ -445,11 +730,16 @@ class TestPGVector(unittest.TestCase): host="localhost", port=5432, diskann=False, - hnsw=False + hnsw=False, + minconn=1, + maxconn=4 ) result = pgvector.get(self.test_ids[0]) + # Verify the _get_cursor context manager was called + mock_get_cursor.assert_called() + # Verify get query was executed get_calls = [call for call in self.mock_cursor.execute.call_args_list if "SELECT id, vector, payload" in str(call)] @@ -461,11 +751,18 @@ class TestPGVector(unittest.TestCase): self.assertEqual(result.payload, {"key": "value1"}) @patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 3) - @patch('mem0.vector_stores.pgvector.psycopg.connect') - @patch('mem0.vector_stores.pgvector.psycopg2.connect') - def test_list_cols_psycopg3(self, mock_psycopg2_connect, mock_psycopg_connect): + @patch('mem0.vector_stores.pgvector.ConnectionPool') + @patch.object(PGVector, '_get_cursor') + def test_list_cols_psycopg3(self, mock_get_cursor, mock_connection_pool): """Test list_cols with psycopg3.""" - mock_psycopg_connect.return_value = self.mock_conn + # Set up mock pool and cursor + mock_pool = MagicMock() + mock_connection_pool.return_value = mock_pool + + # Configure the _get_cursor mock to return our mock cursor + mock_get_cursor.return_value.__enter__.return_value = self.mock_cursor + mock_get_cursor.return_value.__exit__.return_value = None + self.mock_cursor.fetchall.return_value = [("test_collection",), ("other_table",)] pgvector = PGVector( @@ -477,7 +774,9 @@ class TestPGVector(unittest.TestCase): host="localhost", port=5432, diskann=False, - hnsw=False + hnsw=False, + minconn=1, + maxconn=4 ) collections = pgvector.list_cols() @@ -491,10 +790,18 @@ class TestPGVector(unittest.TestCase): self.assertEqual(collections, ["test_collection", "other_table"]) @patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 2) - @patch('mem0.vector_stores.pgvector.psycopg2.connect') - def test_list_cols_psycopg2(self, mock_connect): + @patch('mem0.vector_stores.pgvector.ConnectionPool') + @patch.object(PGVector, '_get_cursor') + def test_list_cols_psycopg2(self, mock_get_cursor, mock_connection_pool): """Test list_cols with psycopg2.""" - mock_connect.return_value = self.mock_conn + # Set up mock pool and cursor + mock_pool = MagicMock() + mock_connection_pool.return_value = mock_pool + + # Configure the _get_cursor mock to return our mock cursor + mock_get_cursor.return_value.__enter__.return_value = self.mock_cursor + mock_get_cursor.return_value.__exit__.return_value = None + self.mock_cursor.fetchall.return_value = [("test_collection",), ("other_table",)] pgvector = PGVector( @@ -506,11 +813,16 @@ class TestPGVector(unittest.TestCase): host="localhost", port=5432, diskann=False, - hnsw=False + hnsw=False, + minconn=1, + maxconn=4 ) collections = pgvector.list_cols() + # Verify the _get_cursor context manager was called + mock_get_cursor.assert_called() + # Verify list_cols query was executed list_calls = [call for call in self.mock_cursor.execute.call_args_list if "SELECT table_name FROM information_schema.tables" in str(call)] @@ -520,11 +832,19 @@ class TestPGVector(unittest.TestCase): self.assertEqual(collections, ["test_collection", "other_table"]) @patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 3) - @patch('mem0.vector_stores.pgvector.psycopg.connect') - @patch('mem0.vector_stores.pgvector.psycopg2.connect') - def test_delete_col_psycopg3(self, mock_psycopg2_connect, mock_psycopg_connect): + @patch('mem0.vector_stores.pgvector.ConnectionPool') + @patch.object(PGVector, '_get_cursor') + def test_delete_col_psycopg3(self, mock_get_cursor, mock_connection_pool): """Test delete_col with psycopg3.""" - mock_psycopg_connect.return_value = self.mock_conn + # Set up mock pool and cursor + mock_pool = MagicMock() + mock_connection_pool.return_value = mock_pool + + # Configure the _get_cursor mock to return our mock cursor + mock_get_cursor.return_value.__enter__.return_value = self.mock_cursor + mock_get_cursor.return_value.__exit__.return_value = None + + self.mock_cursor.fetchall.return_value = [] # No existing collections pgvector = PGVector( dbname="test_db", @@ -535,22 +855,35 @@ class TestPGVector(unittest.TestCase): host="localhost", port=5432, diskann=False, - hnsw=False + hnsw=False, + minconn=1, + maxconn=4 ) pgvector.delete_col() + # Verify the _get_cursor context manager was called + mock_get_cursor.assert_called() + # Verify delete_col query was executed delete_calls = [call for call in self.mock_cursor.execute.call_args_list if "DROP TABLE IF EXISTS test_collection" in str(call)] self.assertTrue(len(delete_calls) > 0) - self.mock_conn.commit.assert_called() @patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 2) - @patch('mem0.vector_stores.pgvector.psycopg2.connect') - def test_delete_col_psycopg2(self, mock_connect): + @patch('mem0.vector_stores.pgvector.ConnectionPool') + @patch.object(PGVector, '_get_cursor') + def test_delete_col_psycopg2(self, mock_get_cursor, mock_connection_pool): """Test delete_col with psycopg2.""" - mock_connect.return_value = self.mock_conn + # Set up mock pool and cursor + mock_pool = MagicMock() + mock_connection_pool.return_value = mock_pool + + # Configure the _get_cursor mock to return our mock cursor + mock_get_cursor.return_value.__enter__.return_value = self.mock_cursor + mock_get_cursor.return_value.__exit__.return_value = None + + self.mock_cursor.fetchall.return_value = [] # No existing collections pgvector = PGVector( dbname="test_db", @@ -561,23 +894,35 @@ class TestPGVector(unittest.TestCase): host="localhost", port=5432, diskann=False, - hnsw=False + hnsw=False, + minconn=1, + maxconn=4 ) pgvector.delete_col() + # Verify the _get_cursor context manager was called + mock_get_cursor.assert_called() + # Verify delete_col query was executed delete_calls = [call for call in self.mock_cursor.execute.call_args_list if "DROP TABLE IF EXISTS test_collection" in str(call)] self.assertTrue(len(delete_calls) > 0) - self.mock_conn.commit.assert_called() @patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 3) - @patch('mem0.vector_stores.pgvector.psycopg.connect') - @patch('mem0.vector_stores.pgvector.psycopg2.connect') - def test_col_info_psycopg3(self, mock_psycopg2_connect, mock_psycopg_connect): + @patch('mem0.vector_stores.pgvector.ConnectionPool') + @patch.object(PGVector, '_get_cursor') + def test_col_info_psycopg3(self, mock_get_cursor, mock_connection_pool): """Test col_info with psycopg3.""" - mock_psycopg_connect.return_value = self.mock_conn + # Set up mock pool and cursor + mock_pool = MagicMock() + mock_connection_pool.return_value = mock_pool + + # Configure the _get_cursor mock to return our mock cursor + mock_get_cursor.return_value.__enter__.return_value = self.mock_cursor + mock_get_cursor.return_value.__exit__.return_value = None + + self.mock_cursor.fetchall.return_value = [] # No existing collections self.mock_cursor.fetchone.return_value = ("test_collection", 100, "1 MB") pgvector = PGVector( @@ -589,11 +934,16 @@ class TestPGVector(unittest.TestCase): host="localhost", port=5432, diskann=False, - hnsw=False + hnsw=False, + minconn=1, + maxconn=4 ) info = pgvector.col_info() + # Verify the _get_cursor context manager was called + mock_get_cursor.assert_called() + # Verify col_info query was executed info_calls = [call for call in self.mock_cursor.execute.call_args_list if "SELECT table_name" in str(call)] @@ -605,10 +955,19 @@ class TestPGVector(unittest.TestCase): self.assertEqual(info["size"], "1 MB") @patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 2) - @patch('mem0.vector_stores.pgvector.psycopg2.connect') - def test_col_info_psycopg2(self, mock_connect): + @patch('mem0.vector_stores.pgvector.ConnectionPool') + @patch.object(PGVector, '_get_cursor') + def test_col_info_psycopg2(self, mock_get_cursor, mock_connection_pool): """Test col_info with psycopg2.""" - mock_connect.return_value = self.mock_conn + # Set up mock pool and cursor + mock_pool = MagicMock() + mock_connection_pool.return_value = mock_pool + + # Configure the _get_cursor mock to return our mock cursor + mock_get_cursor.return_value.__enter__.return_value = self.mock_cursor + mock_get_cursor.return_value.__exit__.return_value = None + + self.mock_cursor.fetchall.return_value = [] # No existing collections self.mock_cursor.fetchone.return_value = ("test_collection", 100, "1 MB") pgvector = PGVector( @@ -620,11 +979,16 @@ class TestPGVector(unittest.TestCase): host="localhost", port=5432, diskann=False, - hnsw=False + hnsw=False, + minconn=1, + maxconn=4 ) info = pgvector.col_info() + # Verify the _get_cursor context manager was called + mock_get_cursor.assert_called() + # Verify col_info query was executed info_calls = [call for call in self.mock_cursor.execute.call_args_list if "SELECT table_name" in str(call)] @@ -636,11 +1000,18 @@ class TestPGVector(unittest.TestCase): self.assertEqual(info["size"], "1 MB") @patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 3) - @patch('mem0.vector_stores.pgvector.psycopg.connect') - @patch('mem0.vector_stores.pgvector.psycopg2.connect') - def test_list_psycopg3(self, mock_psycopg2_connect, mock_psycopg_connect): + @patch('mem0.vector_stores.pgvector.ConnectionPool') + @patch.object(PGVector, '_get_cursor') + def test_list_psycopg3(self, mock_get_cursor, mock_connection_pool): """Test list with psycopg3.""" - mock_psycopg_connect.return_value = self.mock_conn + # Set up mock pool and cursor + mock_pool = MagicMock() + mock_connection_pool.return_value = mock_pool + + # Configure the _get_cursor mock to return our mock cursor + mock_get_cursor.return_value.__enter__.return_value = self.mock_cursor + mock_get_cursor.return_value.__exit__.return_value = None + self.mock_cursor.fetchall.return_value = [ (self.test_ids[0], [0.1, 0.2, 0.3], {"key": "value1"}), (self.test_ids[1], [0.4, 0.5, 0.6], {"key": "value2"}), @@ -655,11 +1026,16 @@ class TestPGVector(unittest.TestCase): host="localhost", port=5432, diskann=False, - hnsw=False + hnsw=False, + minconn=1, + maxconn=4 ) results = pgvector.list(limit=2) + # Verify the _get_cursor context manager was called + mock_get_cursor.assert_called() + # Verify list query was executed list_calls = [call for call in self.mock_cursor.execute.call_args_list if "SELECT id, vector, payload" in str(call)] @@ -672,10 +1048,18 @@ class TestPGVector(unittest.TestCase): self.assertEqual(results[0][1].id, self.test_ids[1]) @patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 2) - @patch('mem0.vector_stores.pgvector.psycopg2.connect') - def test_list_psycopg2(self, mock_connect): + @patch('mem0.vector_stores.pgvector.ConnectionPool') + @patch.object(PGVector, '_get_cursor') + def test_list_psycopg2(self, mock_get_cursor, mock_connection_pool): """Test list with psycopg2.""" - mock_connect.return_value = self.mock_conn + # Set up mock pool and cursor + mock_pool = MagicMock() + mock_connection_pool.return_value = mock_pool + + # Configure the _get_cursor mock to return our mock cursor + mock_get_cursor.return_value.__enter__.return_value = self.mock_cursor + mock_get_cursor.return_value.__exit__.return_value = None + self.mock_cursor.fetchall.return_value = [ (self.test_ids[0], [0.1, 0.2, 0.3], {"key": "value1"}), (self.test_ids[1], [0.4, 0.5, 0.6], {"key": "value2"}), @@ -690,11 +1074,16 @@ class TestPGVector(unittest.TestCase): host="localhost", port=5432, diskann=False, - hnsw=False + hnsw=False, + minconn=1, + maxconn=4 ) results = pgvector.list(limit=2) + # Verify the _get_cursor context manager was called + mock_get_cursor.assert_called() + # Verify list query was executed list_calls = [call for call in self.mock_cursor.execute.call_args_list if "SELECT id, vector, payload" in str(call)] @@ -707,11 +1096,18 @@ class TestPGVector(unittest.TestCase): self.assertEqual(results[0][1].id, self.test_ids[1]) @patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 3) - @patch('mem0.vector_stores.pgvector.psycopg.connect') - @patch('mem0.vector_stores.pgvector.psycopg2.connect') - def test_search_with_filters_psycopg3(self, mock_psycopg2_connect, mock_psycopg_connect): + @patch('mem0.vector_stores.pgvector.ConnectionPool') + @patch.object(PGVector, '_get_cursor') + def test_search_with_filters_psycopg3(self, mock_get_cursor, mock_connection_pool): """Test search with filters using psycopg3.""" - mock_psycopg_connect.return_value = self.mock_conn + # Set up mock pool and cursor + mock_pool = MagicMock() + mock_connection_pool.return_value = mock_pool + + # Configure the _get_cursor mock to return our mock cursor + mock_get_cursor.return_value.__enter__.return_value = self.mock_cursor + mock_get_cursor.return_value.__exit__.return_value = None + self.mock_cursor.fetchall.return_value = [ (self.test_ids[0], 0.1, {"user_id": "alice", "agent_id": "agent1", "run_id": "run1"}), ] @@ -725,12 +1121,17 @@ class TestPGVector(unittest.TestCase): host="localhost", port=5432, diskann=False, - hnsw=False + hnsw=False, + minconn=1, + maxconn=4 ) filters = {"user_id": "alice", "agent_id": "agent1", "run_id": "run1"} results = pgvector.search("test query", [0.1, 0.2, 0.3], limit=2, filters=filters) + # Verify the _get_cursor context manager was called + mock_get_cursor.assert_called() + # Verify search query was executed with filters search_calls = [call for call in self.mock_cursor.execute.call_args_list if "SELECT id, vector <=" in str(call) and "WHERE" in str(call)] @@ -745,10 +1146,18 @@ class TestPGVector(unittest.TestCase): self.assertEqual(results[0].payload["run_id"], "run1") @patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 2) - @patch('mem0.vector_stores.pgvector.psycopg2.connect') - def test_search_with_filters_psycopg2(self, mock_connect): + @patch('mem0.vector_stores.pgvector.ConnectionPool') + @patch.object(PGVector, '_get_cursor') + def test_search_with_filters_psycopg2(self, mock_get_cursor, mock_connection_pool): """Test search with filters using psycopg2.""" - mock_connect.return_value = self.mock_conn + # Set up mock pool and cursor + mock_pool = MagicMock() + mock_connection_pool.return_value = mock_pool + + # Configure the _get_cursor mock to return our mock cursor + mock_get_cursor.return_value.__enter__.return_value = self.mock_cursor + mock_get_cursor.return_value.__exit__.return_value = None + self.mock_cursor.fetchall.return_value = [ (self.test_ids[0], 0.1, {"user_id": "alice", "agent_id": "agent1", "run_id": "run1"}), ] @@ -762,12 +1171,17 @@ class TestPGVector(unittest.TestCase): host="localhost", port=5432, diskann=False, - hnsw=False + hnsw=False, + minconn=1, + maxconn=4 ) filters = {"user_id": "alice", "agent_id": "agent1", "run_id": "run1"} results = pgvector.search("test query", [0.1, 0.2, 0.3], limit=2, filters=filters) + # Verify the _get_cursor context manager was called + mock_get_cursor.assert_called() + # Verify search query was executed with filters search_calls = [call for call in self.mock_cursor.execute.call_args_list if "SELECT id, vector <=" in str(call) and "WHERE" in str(call)] @@ -782,11 +1196,18 @@ class TestPGVector(unittest.TestCase): self.assertEqual(results[0].payload["run_id"], "run1") @patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 3) - @patch('mem0.vector_stores.pgvector.psycopg.connect') - @patch('mem0.vector_stores.pgvector.psycopg2.connect') - def test_search_with_single_filter_psycopg3(self, mock_psycopg2_connect, mock_psycopg_connect): + @patch('mem0.vector_stores.pgvector.ConnectionPool') + @patch.object(PGVector, '_get_cursor') + def test_search_with_single_filter_psycopg3(self, mock_get_cursor, mock_connection_pool): """Test search with single filter using psycopg3.""" - mock_psycopg_connect.return_value = self.mock_conn + # Set up mock pool and cursor + mock_pool = MagicMock() + mock_connection_pool.return_value = mock_pool + + # Configure the _get_cursor mock to return our mock cursor + mock_get_cursor.return_value.__enter__.return_value = self.mock_cursor + mock_get_cursor.return_value.__exit__.return_value = None + self.mock_cursor.fetchall.return_value = [ (self.test_ids[0], 0.1, {"user_id": "alice"}), ] @@ -800,12 +1221,17 @@ class TestPGVector(unittest.TestCase): host="localhost", port=5432, diskann=False, - hnsw=False + hnsw=False, + minconn=1, + maxconn=4 ) filters = {"user_id": "alice"} results = pgvector.search("test query", [0.1, 0.2, 0.3], limit=2, filters=filters) + # Verify the _get_cursor context manager was called + mock_get_cursor.assert_called() + # Verify search query was executed with single filter search_calls = [call for call in self.mock_cursor.execute.call_args_list if "SELECT id, vector <=" in str(call) and "WHERE" in str(call)] @@ -818,10 +1244,18 @@ class TestPGVector(unittest.TestCase): self.assertEqual(results[0].payload["user_id"], "alice") @patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 2) - @patch('mem0.vector_stores.pgvector.psycopg2.connect') - def test_search_with_single_filter_psycopg2(self, mock_connect): + @patch('mem0.vector_stores.pgvector.ConnectionPool') + @patch.object(PGVector, '_get_cursor') + def test_search_with_single_filter_psycopg2(self, mock_get_cursor, mock_connection_pool): """Test search with single filter using psycopg2.""" - mock_connect.return_value = self.mock_conn + # Set up mock pool and cursor + mock_pool = MagicMock() + mock_connection_pool.return_value = mock_pool + + # Configure the _get_cursor mock to return our mock cursor + mock_get_cursor.return_value.__enter__.return_value = self.mock_cursor + mock_get_cursor.return_value.__exit__.return_value = None + self.mock_cursor.fetchall.return_value = [ (self.test_ids[0], 0.1, {"user_id": "alice"}), ] @@ -835,12 +1269,17 @@ class TestPGVector(unittest.TestCase): host="localhost", port=5432, diskann=False, - hnsw=False + hnsw=False, + minconn=1, + maxconn=4 ) filters = {"user_id": "alice"} results = pgvector.search("test query", [0.1, 0.2, 0.3], limit=2, filters=filters) + # Verify the _get_cursor context manager was called + mock_get_cursor.assert_called() + # Verify search query was executed with single filter search_calls = [call for call in self.mock_cursor.execute.call_args_list if "SELECT id, vector <=" in str(call) and "WHERE" in str(call)] @@ -853,11 +1292,18 @@ class TestPGVector(unittest.TestCase): self.assertEqual(results[0].payload["user_id"], "alice") @patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 3) - @patch('mem0.vector_stores.pgvector.psycopg.connect') - @patch('mem0.vector_stores.pgvector.psycopg2.connect') - def test_search_with_no_filters_psycopg3(self, mock_psycopg2_connect, mock_psycopg_connect): + @patch('mem0.vector_stores.pgvector.ConnectionPool') + @patch.object(PGVector, '_get_cursor') + def test_search_with_no_filters_psycopg3(self, mock_get_cursor, mock_connection_pool): """Test search with no filters using psycopg3.""" - mock_psycopg_connect.return_value = self.mock_conn + # Set up mock pool and cursor + mock_pool = MagicMock() + mock_connection_pool.return_value = mock_pool + + # Configure the _get_cursor mock to return our mock cursor + mock_get_cursor.return_value.__enter__.return_value = self.mock_cursor + mock_get_cursor.return_value.__exit__.return_value = None + self.mock_cursor.fetchall.return_value = [ (self.test_ids[0], 0.1, {"key": "value1"}), (self.test_ids[1], 0.2, {"key": "value2"}), @@ -872,11 +1318,16 @@ class TestPGVector(unittest.TestCase): host="localhost", port=5432, diskann=False, - hnsw=False + hnsw=False, + minconn=1, + maxconn=4 ) results = pgvector.search("test query", [0.1, 0.2, 0.3], limit=2, filters=None) + # Verify the _get_cursor context manager was called + mock_get_cursor.assert_called() + # Verify search query was executed without WHERE clause search_calls = [call for call in self.mock_cursor.execute.call_args_list if "SELECT id, vector <=" in str(call) and "WHERE" not in str(call)] @@ -890,10 +1341,18 @@ class TestPGVector(unittest.TestCase): self.assertEqual(results[1].score, 0.2) @patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 2) - @patch('mem0.vector_stores.pgvector.psycopg2.connect') - def test_search_with_no_filters_psycopg2(self, mock_connect): + @patch('mem0.vector_stores.pgvector.ConnectionPool') + @patch.object(PGVector, '_get_cursor') + def test_search_with_no_filters_psycopg2(self, mock_get_cursor, mock_connection_pool): """Test search with no filters using psycopg2.""" - mock_connect.return_value = self.mock_conn + # Set up mock pool and cursor + mock_pool = MagicMock() + mock_connection_pool.return_value = mock_pool + + # Configure the _get_cursor mock to return our mock cursor + mock_get_cursor.return_value.__enter__.return_value = self.mock_cursor + mock_get_cursor.return_value.__exit__.return_value = None + self.mock_cursor.fetchall.return_value = [ (self.test_ids[0], 0.1, {"key": "value1"}), (self.test_ids[1], 0.2, {"key": "value2"}), @@ -908,11 +1367,16 @@ class TestPGVector(unittest.TestCase): host="localhost", port=5432, diskann=False, - hnsw=False + hnsw=False, + minconn=1, + maxconn=4 ) results = pgvector.search("test query", [0.1, 0.2, 0.3], limit=2, filters=None) + # Verify the _get_cursor context manager was called + mock_get_cursor.assert_called() + # Verify search query was executed without WHERE clause search_calls = [call for call in self.mock_cursor.execute.call_args_list if "SELECT id, vector <=" in str(call) and "WHERE" not in str(call)] @@ -926,11 +1390,18 @@ class TestPGVector(unittest.TestCase): self.assertEqual(results[1].score, 0.2) @patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 3) - @patch('mem0.vector_stores.pgvector.psycopg.connect') - @patch('mem0.vector_stores.pgvector.psycopg2.connect') - def test_list_with_filters_psycopg3(self, mock_psycopg2_connect, mock_psycopg_connect): + @patch('mem0.vector_stores.pgvector.ConnectionPool') + @patch.object(PGVector, '_get_cursor') + def test_list_with_filters_psycopg3(self, mock_get_cursor, mock_connection_pool): """Test list with filters using psycopg3.""" - mock_psycopg_connect.return_value = self.mock_conn + # Set up mock pool and cursor + mock_pool = MagicMock() + mock_connection_pool.return_value = mock_pool + + # Configure the _get_cursor mock to return our mock cursor + mock_get_cursor.return_value.__enter__.return_value = self.mock_cursor + mock_get_cursor.return_value.__exit__.return_value = None + self.mock_cursor.fetchall.return_value = [ (self.test_ids[0], [0.1, 0.2, 0.3], {"user_id": "alice", "agent_id": "agent1"}), ] @@ -944,12 +1415,17 @@ class TestPGVector(unittest.TestCase): host="localhost", port=5432, diskann=False, - hnsw=False + hnsw=False, + minconn=1, + maxconn=4 ) filters = {"user_id": "alice", "agent_id": "agent1"} results = pgvector.list(filters=filters, limit=2) + # Verify the _get_cursor context manager was called + mock_get_cursor.assert_called() + # Verify list query was executed with filters list_calls = [call for call in self.mock_cursor.execute.call_args_list if "SELECT id, vector, payload" in str(call) and "WHERE" in str(call)] @@ -963,10 +1439,18 @@ class TestPGVector(unittest.TestCase): self.assertEqual(results[0][0].payload["agent_id"], "agent1") @patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 2) - @patch('mem0.vector_stores.pgvector.psycopg2.connect') - def test_list_with_filters_psycopg2(self, mock_connect): + @patch('mem0.vector_stores.pgvector.ConnectionPool') + @patch.object(PGVector, '_get_cursor') + def test_list_with_filters_psycopg2(self, mock_get_cursor, mock_connection_pool): """Test list with filters using psycopg2.""" - mock_connect.return_value = self.mock_conn + # Set up mock pool and cursor + mock_pool = MagicMock() + mock_connection_pool.return_value = mock_pool + + # Configure the _get_cursor mock to return our mock cursor + mock_get_cursor.return_value.__enter__.return_value = self.mock_cursor + mock_get_cursor.return_value.__exit__.return_value = None + self.mock_cursor.fetchall.return_value = [ (self.test_ids[0], [0.1, 0.2, 0.3], {"user_id": "alice", "agent_id": "agent1"}), ] @@ -980,12 +1464,17 @@ class TestPGVector(unittest.TestCase): host="localhost", port=5432, diskann=False, - hnsw=False + hnsw=False, + minconn=1, + maxconn=4 ) filters = {"user_id": "alice", "agent_id": "agent1"} results = pgvector.list(filters=filters, limit=2) + # Verify the _get_cursor context manager was called + mock_get_cursor.assert_called() + # Verify list query was executed with filters list_calls = [call for call in self.mock_cursor.execute.call_args_list if "SELECT id, vector, payload" in str(call) and "WHERE" in str(call)] @@ -999,11 +1488,18 @@ class TestPGVector(unittest.TestCase): self.assertEqual(results[0][0].payload["agent_id"], "agent1") @patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 3) - @patch('mem0.vector_stores.pgvector.psycopg.connect') - @patch('mem0.vector_stores.pgvector.psycopg2.connect') - def test_list_with_single_filter_psycopg3(self, mock_psycopg2_connect, mock_psycopg_connect): + @patch('mem0.vector_stores.pgvector.ConnectionPool') + @patch.object(PGVector, '_get_cursor') + def test_list_with_single_filter_psycopg3(self, mock_get_cursor, mock_connection_pool): """Test list with single filter using psycopg3.""" - mock_psycopg_connect.return_value = self.mock_conn + # Set up mock pool and cursor + mock_pool = MagicMock() + mock_connection_pool.return_value = mock_pool + + # Configure the _get_cursor mock to return our mock cursor + mock_get_cursor.return_value.__enter__.return_value = self.mock_cursor + mock_get_cursor.return_value.__exit__.return_value = None + self.mock_cursor.fetchall.return_value = [ (self.test_ids[0], [0.1, 0.2, 0.3], {"user_id": "alice"}), ] @@ -1017,12 +1513,17 @@ class TestPGVector(unittest.TestCase): host="localhost", port=5432, diskann=False, - hnsw=False + hnsw=False, + minconn=1, + maxconn=4 ) filters = {"user_id": "alice"} results = pgvector.list(filters=filters, limit=2) + # Verify the _get_cursor context manager was called + mock_get_cursor.assert_called() + # Verify list query was executed with single filter list_calls = [call for call in self.mock_cursor.execute.call_args_list if "SELECT id, vector, payload" in str(call) and "WHERE" in str(call)] @@ -1035,10 +1536,18 @@ class TestPGVector(unittest.TestCase): self.assertEqual(results[0][0].payload["user_id"], "alice") @patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 2) - @patch('mem0.vector_stores.pgvector.psycopg2.connect') - def test_list_with_single_filter_psycopg2(self, mock_connect): + @patch('mem0.vector_stores.pgvector.ConnectionPool') + @patch.object(PGVector, '_get_cursor') + def test_list_with_single_filter_psycopg2(self, mock_get_cursor, mock_connection_pool): """Test list with single filter using psycopg2.""" - mock_connect.return_value = self.mock_conn + # Set up mock pool and cursor + mock_pool = MagicMock() + mock_connection_pool.return_value = mock_pool + + # Configure the _get_cursor mock to return our mock cursor + mock_get_cursor.return_value.__enter__.return_value = self.mock_cursor + mock_get_cursor.return_value.__exit__.return_value = None + self.mock_cursor.fetchall.return_value = [ (self.test_ids[0], [0.1, 0.2, 0.3], {"user_id": "alice"}), ] @@ -1052,12 +1561,17 @@ class TestPGVector(unittest.TestCase): host="localhost", port=5432, diskann=False, - hnsw=False + hnsw=False, + minconn=1, + maxconn=4 ) filters = {"user_id": "alice"} results = pgvector.list(filters=filters, limit=2) + # Verify the _get_cursor context manager was called + mock_get_cursor.assert_called() + # Verify list query was executed with single filter list_calls = [call for call in self.mock_cursor.execute.call_args_list if "SELECT id, vector, payload" in str(call) and "WHERE" in str(call)] @@ -1070,11 +1584,18 @@ class TestPGVector(unittest.TestCase): self.assertEqual(results[0][0].payload["user_id"], "alice") @patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 3) - @patch('mem0.vector_stores.pgvector.psycopg.connect') - @patch('mem0.vector_stores.pgvector.psycopg2.connect') - def test_list_with_no_filters_psycopg3(self, mock_psycopg2_connect, mock_psycopg_connect): + @patch('mem0.vector_stores.pgvector.ConnectionPool') + @patch.object(PGVector, '_get_cursor') + def test_list_with_no_filters_psycopg3(self, mock_get_cursor, mock_connection_pool): """Test list with no filters using psycopg3.""" - mock_psycopg_connect.return_value = self.mock_conn + # Set up mock pool and cursor + mock_pool = MagicMock() + mock_connection_pool.return_value = mock_pool + + # Configure the _get_cursor mock to return our mock cursor + mock_get_cursor.return_value.__enter__.return_value = self.mock_cursor + mock_get_cursor.return_value.__exit__.return_value = None + self.mock_cursor.fetchall.return_value = [ (self.test_ids[0], [0.1, 0.2, 0.3], {"key": "value1"}), (self.test_ids[1], [0.4, 0.5, 0.6], {"key": "value2"}), @@ -1089,11 +1610,16 @@ class TestPGVector(unittest.TestCase): host="localhost", port=5432, diskann=False, - hnsw=False + hnsw=False, + minconn=1, + maxconn=4 ) results = pgvector.list(filters=None, limit=2) + # Verify the _get_cursor context manager was called + mock_get_cursor.assert_called() + # Verify list query was executed without WHERE clause list_calls = [call for call in self.mock_cursor.execute.call_args_list if "SELECT id, vector, payload" in str(call) and "WHERE" not in str(call)] @@ -1106,10 +1632,18 @@ class TestPGVector(unittest.TestCase): self.assertEqual(results[0][1].id, self.test_ids[1]) @patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 2) - @patch('mem0.vector_stores.pgvector.psycopg2.connect') - def test_list_with_no_filters_psycopg2(self, mock_connect): + @patch('mem0.vector_stores.pgvector.ConnectionPool') + @patch.object(PGVector, '_get_cursor') + def test_list_with_no_filters_psycopg2(self, mock_get_cursor, mock_connection_pool): """Test list with no filters using psycopg2.""" - mock_connect.return_value = self.mock_conn + # Set up mock pool and cursor + mock_pool = MagicMock() + mock_connection_pool.return_value = mock_pool + + # Configure the _get_cursor mock to return our mock cursor + mock_get_cursor.return_value.__enter__.return_value = self.mock_cursor + mock_get_cursor.return_value.__exit__.return_value = None + self.mock_cursor.fetchall.return_value = [ (self.test_ids[0], [0.1, 0.2, 0.3], {"key": "value1"}), (self.test_ids[1], [0.4, 0.5, 0.6], {"key": "value2"}), @@ -1124,11 +1658,16 @@ class TestPGVector(unittest.TestCase): host="localhost", port=5432, diskann=False, - hnsw=False + hnsw=False, + minconn=1, + maxconn=4 ) results = pgvector.list(filters=None, limit=2) + # Verify the _get_cursor context manager was called + mock_get_cursor.assert_called() + # Verify list query was executed without WHERE clause list_calls = [call for call in self.mock_cursor.execute.call_args_list if "SELECT id, vector, payload" in str(call) and "WHERE" not in str(call)] @@ -1141,11 +1680,18 @@ class TestPGVector(unittest.TestCase): self.assertEqual(results[0][1].id, self.test_ids[1]) @patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 3) - @patch('mem0.vector_stores.pgvector.psycopg.connect') - @patch('mem0.vector_stores.pgvector.psycopg2.connect') - def test_reset_psycopg3(self, mock_psycopg2_connect, mock_psycopg_connect): + @patch('mem0.vector_stores.pgvector.ConnectionPool') + @patch.object(PGVector, '_get_cursor') + def test_reset_psycopg3(self, mock_get_cursor, mock_connection_pool): """Test reset with psycopg3.""" - mock_psycopg_connect.return_value = self.mock_conn + # Set up mock pool and cursor + mock_pool = MagicMock() + mock_connection_pool.return_value = mock_pool + + # Configure the _get_cursor mock to return our mock cursor + mock_get_cursor.return_value.__enter__.return_value = self.mock_cursor + mock_get_cursor.return_value.__exit__.return_value = None + self.mock_cursor.fetchall.return_value = [] pgvector = PGVector( @@ -1157,11 +1703,16 @@ class TestPGVector(unittest.TestCase): host="localhost", port=5432, diskann=False, - hnsw=False + hnsw=False, + minconn=1, + maxconn=4 ) pgvector.reset() + # Verify the _get_cursor context manager was called + mock_get_cursor.assert_called() + # Verify reset operations were executed drop_calls = [call for call in self.mock_cursor.execute.call_args_list if "DROP TABLE IF EXISTS" in str(call)] @@ -1171,10 +1722,18 @@ class TestPGVector(unittest.TestCase): self.assertTrue(len(create_calls) > 0) @patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 2) - @patch('mem0.vector_stores.pgvector.psycopg2.connect') - def test_reset_psycopg2(self, mock_connect): + @patch('mem0.vector_stores.pgvector.ConnectionPool') + @patch.object(PGVector, '_get_cursor') + def test_reset_psycopg2(self, mock_get_cursor, mock_connection_pool): """Test reset with psycopg2.""" - mock_connect.return_value = self.mock_conn + # Set up mock pool and cursor + mock_pool = MagicMock() + mock_connection_pool.return_value = mock_pool + + # Configure the _get_cursor mock to return our mock cursor + mock_get_cursor.return_value.__enter__.return_value = self.mock_cursor + mock_get_cursor.return_value.__exit__.return_value = None + self.mock_cursor.fetchall.return_value = [] pgvector = PGVector( @@ -1186,11 +1745,16 @@ class TestPGVector(unittest.TestCase): host="localhost", port=5432, diskann=False, - hnsw=False + hnsw=False, + minconn=1, + maxconn=4 ) pgvector.reset() + # Verify the _get_cursor context manager was called + mock_get_cursor.assert_called() + # Verify reset operations were executed drop_calls = [call for call in self.mock_cursor.execute.call_args_list if "DROP TABLE IF EXISTS" in str(call)] @@ -1199,6 +1763,468 @@ class TestPGVector(unittest.TestCase): self.assertTrue(len(drop_calls) > 0) self.assertTrue(len(create_calls) > 0) + # Enhanced Tests for JSON Serialization + @patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 3) + @patch('mem0.vector_stores.pgvector.ConnectionPool') + @patch.object(PGVector, '_get_cursor') + @patch('mem0.vector_stores.pgvector.Json') + def test_update_payload_psycopg3_json_handling(self, mock_json, mock_get_cursor, mock_connection_pool): + """Test that psycopg3 update uses Json() wrapper for payload serialization.""" + # Set up mock pool and cursor + mock_pool = MagicMock() + mock_connection_pool.return_value = mock_pool + + # Configure the _get_cursor mock to return our mock cursor + mock_get_cursor.return_value.__enter__.return_value = self.mock_cursor + mock_get_cursor.return_value.__exit__.return_value = None + + self.mock_cursor.fetchall.return_value = [] # No existing collections + + pgvector = PGVector( + dbname="test_db", + collection_name="test_collection", + embedding_model_dims=3, + user="test_user", + password="test_pass", + host="localhost", + port=5432, + diskann=False, + hnsw=False, + minconn=1, + maxconn=4 + ) + + test_payload = {"test": "data", "number": 42} + pgvector.update("test-id-123", payload=test_payload) + + # Verify Json() wrapper was used for psycopg3 + mock_json.assert_called_once_with(test_payload) + + # Verify the update query was executed + update_calls = [call for call in self.mock_cursor.execute.call_args_list + if "UPDATE test_collection SET payload" in str(call)] + self.assertTrue(len(update_calls) > 0) + + @patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 2) + @patch('mem0.vector_stores.pgvector.ConnectionPool') + @patch.object(PGVector, '_get_cursor') + @patch('mem0.vector_stores.pgvector.Json') + def test_update_payload_psycopg2_json_handling(self, mock_json, mock_get_cursor, mock_connection_pool): + """Test that psycopg2 update uses psycopg2.extras.Json() wrapper for payload serialization.""" + # Set up mock pool and cursor + mock_pool = MagicMock() + mock_connection_pool.return_value = mock_pool + + # Configure the _get_cursor mock to return our mock cursor + mock_get_cursor.return_value.__enter__.return_value = self.mock_cursor + mock_get_cursor.return_value.__exit__.return_value = None + + self.mock_cursor.fetchall.return_value = [] # No existing collections + + pgvector = PGVector( + dbname="test_db", + collection_name="test_collection", + embedding_model_dims=3, + user="test_user", + password="test_pass", + host="localhost", + port=5432, + diskann=False, + hnsw=False, + minconn=1, + maxconn=4 + ) + + test_payload = {"test": "data", "number": 42} + pgvector.update("test-id-123", payload=test_payload) + + # Verify psycopg2.extras.Json() wrapper was used + mock_json.assert_called_once_with(test_payload) + + # Verify the update query was executed + update_calls = [call for call in self.mock_cursor.execute.call_args_list + if "UPDATE test_collection SET payload" in str(call)] + self.assertTrue(len(update_calls) > 0) + + @patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 2) + @patch('mem0.vector_stores.pgvector.ConnectionPool') + def test_transaction_rollback_on_error_psycopg2(self, mock_connection_pool): + """Test that psycopg2 properly rolls back transactions on errors.""" + mock_pool = MagicMock() + mock_connection_pool.return_value = mock_pool + + # Set up mock connection that will raise an error only on delete + mock_conn = MagicMock() + mock_cursor = MagicMock() + mock_conn.cursor.return_value = mock_cursor + mock_pool.getconn.return_value = mock_conn + + # Only raise exception on the delete operation, not during setup + def execute_side_effect(*args, **kwargs): + if args and "DELETE FROM" in str(args[0]): + raise Exception("Database error") + return MagicMock() + mock_cursor.execute.side_effect = execute_side_effect + self.mock_cursor.fetchall.return_value = [] # No existing collections initially + + pgvector = PGVector( + dbname="test_db", + collection_name="test_collection", + embedding_model_dims=3, + user="test_user", + password="test_pass", + host="localhost", + port=5432, + diskann=False, + hnsw=False, + minconn=1, + maxconn=4 + ) + + # Attempt an operation that will fail + with self.assertRaises(Exception) as context: + pgvector.delete("test-id") + + self.assertIn("Database error", str(context.exception)) + # Verify rollback was called + mock_conn.rollback.assert_called() + # Verify connection was returned to pool + mock_pool.putconn.assert_called_with(mock_conn) + + @patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 2) + @patch('mem0.vector_stores.pgvector.ConnectionPool') + def test_commit_on_success_psycopg2(self, mock_connection_pool): + """Test that psycopg2 properly commits transactions on success.""" + mock_pool = MagicMock() + mock_connection_pool.return_value = mock_pool + + # Set up mock connection for successful operation + mock_conn = MagicMock() + mock_cursor = MagicMock() + mock_conn.cursor.return_value = mock_cursor + mock_pool.getconn.return_value = mock_conn + + self.mock_cursor.fetchall.return_value = [] # No existing collections initially + + pgvector = PGVector( + dbname="test_db", + collection_name="test_collection", + embedding_model_dims=3, + user="test_user", + password="test_pass", + host="localhost", + port=5432, + diskann=False, + hnsw=False, + minconn=1, + maxconn=4 + ) + + # Perform an operation that requires commit + pgvector.delete("test-id") + + # Verify commit was called + mock_conn.commit.assert_called() + # Verify connection was returned to pool + mock_pool.putconn.assert_called_with(mock_conn) + + # Enhanced Tests for Error Handling + @patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 3) + @patch('mem0.vector_stores.pgvector.ConnectionPool') + @patch.object(PGVector, '_get_cursor') + def test_pool_connection_error_handling(self, mock_get_cursor, mock_connection_pool): + """Test handling of connection pool errors.""" + mock_pool = MagicMock() + mock_connection_pool.return_value = mock_pool + + # Use a flag to only raise the exception after PGVector is initialized + raise_on_search = {'active': False} + def get_cursor_side_effect(*args, **kwargs): + if raise_on_search['active']: + raise Exception("Connection pool exhausted") + return self.mock_cursor + + mock_get_cursor.side_effect = get_cursor_side_effect + self.mock_cursor.fetchall.return_value = [] + + pgvector = PGVector( + dbname="test_db", + collection_name="test_collection", + embedding_model_dims=3, + user="test_user", + password="test_pass", + host="localhost", + port=5432, + diskann=False, + hnsw=False, + minconn=1, + maxconn=4 + ) + + # Activate the exception for search only + raise_on_search['active'] = True + with self.assertRaises(Exception) as context: + pgvector.search("test query", [0.1, 0.2, 0.3]) + + self.assertIn("Connection pool exhausted", str(context.exception)) + + # Enhanced Tests for Vector and Payload Update Combinations + @patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 3) + @patch('mem0.vector_stores.pgvector.ConnectionPool') + @patch.object(PGVector, '_get_cursor') + def test_update_vector_only_psycopg3(self, mock_get_cursor, mock_connection_pool): + """Test updating only vector without payload.""" + # Set up mock pool and cursor + mock_pool = MagicMock() + mock_connection_pool.return_value = mock_pool + + # Configure the _get_cursor mock to return our mock cursor + mock_get_cursor.return_value.__enter__.return_value = self.mock_cursor + mock_get_cursor.return_value.__exit__.return_value = None + + self.mock_cursor.fetchall.return_value = [] # No existing collections + + pgvector = PGVector( + dbname="test_db", + collection_name="test_collection", + embedding_model_dims=3, + user="test_user", + password="test_pass", + host="localhost", + port=5432, + diskann=False, + hnsw=False, + minconn=1, + maxconn=4 + ) + + test_vector = [0.1, 0.2, 0.3] + pgvector.update("test-id", vector=test_vector) + + # Verify only vector update query was executed (not payload) + vector_update_calls = [call for call in self.mock_cursor.execute.call_args_list + if "UPDATE test_collection SET vector" in str(call) and "payload" not in str(call)] + payload_update_calls = [call for call in self.mock_cursor.execute.call_args_list + if "UPDATE test_collection SET payload" in str(call)] + + self.assertTrue(len(vector_update_calls) > 0) + self.assertEqual(len(payload_update_calls), 0) + + @patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 3) + @patch('mem0.vector_stores.pgvector.ConnectionPool') + @patch.object(PGVector, '_get_cursor') + def test_update_both_vector_and_payload_psycopg3(self, mock_get_cursor, mock_connection_pool): + """Test updating both vector and payload.""" + # Set up mock pool and cursor + mock_pool = MagicMock() + mock_connection_pool.return_value = mock_pool + + # Configure the _get_cursor mock to return our mock cursor + mock_get_cursor.return_value.__enter__.return_value = self.mock_cursor + mock_get_cursor.return_value.__exit__.return_value = None + + self.mock_cursor.fetchall.return_value = [] # No existing collections + + pgvector = PGVector( + dbname="test_db", + collection_name="test_collection", + embedding_model_dims=3, + user="test_user", + password="test_pass", + host="localhost", + port=5432, + diskann=False, + hnsw=False, + minconn=1, + maxconn=4 + ) + + test_vector = [0.1, 0.2, 0.3] + test_payload = {"updated": True} + pgvector.update("test-id", vector=test_vector, payload=test_payload) + + # Verify both vector and payload update queries were executed + vector_update_calls = [call for call in self.mock_cursor.execute.call_args_list + if "UPDATE test_collection SET vector" in str(call)] + payload_update_calls = [call for call in self.mock_cursor.execute.call_args_list + if "UPDATE test_collection SET payload" in str(call)] + + self.assertTrue(len(vector_update_calls) > 0) + self.assertTrue(len(payload_update_calls) > 0) + + # Enhanced Tests for Connection String Handling + @patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 3) + @patch('mem0.vector_stores.pgvector.ConnectionPool') + def test_connection_string_with_sslmode_psycopg3(self, mock_connection_pool): + """Test connection string handling with SSL mode.""" + mock_pool = MagicMock() + mock_connection_pool.return_value = mock_pool + self.mock_cursor.fetchall.return_value = [] # No existing collections + + connection_string = "postgresql://user:pass@localhost:5432/db" + + pgvector = PGVector( + dbname="test_db", # Will be overridden by connection_string + collection_name="test_collection", + embedding_model_dims=3, + user="test_user", + password="test_pass", + host="localhost", + port=5432, + diskann=False, + hnsw=False, + minconn=1, + maxconn=4, + sslmode="require", + connection_string=connection_string + ) + + # Verify ConnectionPool was called with the connection string including sslmode + expected_conn_string = f"{connection_string} sslmode=require" + mock_connection_pool.assert_called_with( + conninfo=expected_conn_string, + min_size=1, + max_size=4, + open=True + ) + self.assertEqual(pgvector.collection_name, "test_collection") + self.assertEqual(pgvector.embedding_model_dims, 3) + + # Enhanced Test for Index Creation with DiskANN + @patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 3) + @patch('mem0.vector_stores.pgvector.ConnectionPool') + @patch.object(PGVector, '_get_cursor') + def test_create_col_with_diskann_psycopg3(self, mock_get_cursor, mock_connection_pool): + """Test collection creation with DiskANN index.""" + # Set up mock pool and cursor + mock_pool = MagicMock() + mock_connection_pool.return_value = mock_pool + + # Configure the _get_cursor mock to return our mock cursor + mock_get_cursor.return_value.__enter__.return_value = self.mock_cursor + mock_get_cursor.return_value.__exit__.return_value = None + + # Mock vectorscale extension as available + self.mock_cursor.fetchall.return_value = [] # No existing collections + self.mock_cursor.fetchone.return_value = ("vectorscale",) # Extension exists + + pgvector = PGVector( + dbname="test_db", + collection_name="test_collection", + embedding_model_dims=3, + user="test_user", + password="test_pass", + host="localhost", + port=5432, + diskann=True, # Enable DiskANN + hnsw=False, + minconn=1, + maxconn=4 + ) + + # Verify DiskANN index creation query was executed + diskann_calls = [call for call in self.mock_cursor.execute.call_args_list + if "USING diskann" in str(call)] + self.assertTrue(len(diskann_calls) > 0) + self.assertEqual(pgvector.collection_name, "test_collection") + self.assertEqual(pgvector.embedding_model_dims, 3) + + + @patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 3) + @patch('mem0.vector_stores.pgvector.ConnectionPool') + @patch.object(PGVector, '_get_cursor') + def test_create_col_with_hnsw_psycopg3(self, mock_get_cursor, mock_connection_pool): + """Test collection creation with HNSW index.""" + # Set up mock pool and cursor + mock_pool = MagicMock() + mock_connection_pool.return_value = mock_pool + + # Configure the _get_cursor mock to return our mock cursor + mock_get_cursor.return_value.__enter__.return_value = self.mock_cursor + mock_get_cursor.return_value.__exit__.return_value = None + + self.mock_cursor.fetchall.return_value = [] # No existing collections + + pgvector = PGVector( + dbname="test_db", + collection_name="test_collection", + embedding_model_dims=3, + user="test_user", + password="test_pass", + host="localhost", + port=5432, + diskann=False, + hnsw=True, # Enable HNSW + minconn=1, + maxconn=4 + ) + + # Verify HNSW index creation query was executed + hnsw_calls = [call for call in self.mock_cursor.execute.call_args_list + if "USING hnsw" in str(call)] + self.assertTrue(len(hnsw_calls) > 0) + self.assertEqual(pgvector.collection_name, "test_collection") + self.assertEqual(pgvector.embedding_model_dims, 3) + + # Enhanced Test for Pool Cleanup + def test_pool_cleanup_psycopg3(self): + """Test that psycopg3 pool is properly closed on object deletion.""" + with patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 3), \ + patch('mem0.vector_stores.pgvector.ConnectionPool') as mock_connection_pool: + + mock_pool = MagicMock() + mock_connection_pool.return_value = mock_pool + self.mock_cursor.fetchall.return_value = [] # No existing collections + + pgvector = PGVector( + dbname="test_db", + collection_name="test_collection", + embedding_model_dims=3, + user="test_user", + password="test_pass", + host="localhost", + port=5432, + diskann=False, + hnsw=False, + minconn=1, + maxconn=4 + ) + + # Trigger __del__ method + del pgvector + + # Verify pool.close() was called + mock_pool.close.assert_called() + + def test_pool_cleanup_psycopg2(self): + """Test that psycopg2 pool is properly closed on object deletion.""" + with patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 2), \ + patch('mem0.vector_stores.pgvector.ConnectionPool') as mock_connection_pool: + + mock_pool = MagicMock() + mock_connection_pool.return_value = mock_pool + self.mock_cursor.fetchall.return_value = [] # No existing collections + + pgvector = PGVector( + dbname="test_db", + collection_name="test_collection", + embedding_model_dims=3, + user="test_user", + password="test_pass", + host="localhost", + port=5432, + diskann=False, + hnsw=False, + minconn=1, + maxconn=4 + ) + + # Trigger __del__ method + del pgvector + + # Verify pool.closeall() was called + mock_pool.closeall.assert_called() + def tearDown(self): """Clean up after each test.""" pass diff --git a/vercel-ai-sdk/package.json b/vercel-ai-sdk/package.json index c457489ca..34afea7a1 100644 --- a/vercel-ai-sdk/package.json +++ b/vercel-ai-sdk/package.json @@ -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", diff --git a/vercel-ai-sdk/src/mem0-generic-language-model.ts b/vercel-ai-sdk/src/mem0-generic-language-model.ts index a7bf65a0e..7c2c8898b 100644 --- a/vercel-ai-sdk/src/mem0-generic-language-model.ts +++ b/vercel-ai-sdk/src/mem0-generic-language-model.ts @@ -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, }; diff --git a/vercel-ai-sdk/src/mem0-types.ts b/vercel-ai-sdk/src/mem0-types.ts index 70c4d4d9b..ed68aad0f 100644 --- a/vercel-ai-sdk/src/mem0-types.ts +++ b/vercel-ai-sdk/src/mem0-types.ts @@ -29,6 +29,7 @@ export interface Mem0ConfigSettings { host?: string; output_format?: string; filter_memories?: boolean; + async_mode?: boolean; } export interface Mem0ChatConfig extends Mem0ConfigSettings, Mem0ProviderSettings {} diff --git a/vercel-ai-sdk/src/mem0-utils.ts b/vercel-ai-sdk/src/mem0-utils.ts index 6484882dc..2a3906720 100644 --- a/vercel-ai-sdk/src/mem0-utils.ts +++ b/vercel-ai-sdk/src/mem0-utils.ts @@ -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, }); }