From 3d0ece1bcf63823244ef1dedd8e7e5e604ec4a4c Mon Sep 17 00:00:00 2001 From: Sheharyar Ahmad Date: Fri, 29 Aug 2025 18:36:39 +0500 Subject: [PATCH] Refactor PGVector to Use Internal Connection Pools and Context Managers (#3373) --- mem0/configs/vector_stores/pgvector.py | 14 +- mem0/vector_stores/pgvector.py | 360 +++--- pyproject.toml | 1 + tests/vector_stores/test_pgvector.py | 1468 ++++++++++++++++++++---- 4 files changed, 1454 insertions(+), 389 deletions(-) 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/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/pyproject.toml b/pyproject.toml index 9fbb96107..7227de60f 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -39,6 +39,7 @@ 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", 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