diff --git a/mem0/vector_stores/pgvector.py b/mem0/vector_stores/pgvector.py index a2e4e725e..4b1606a1d 100644 --- a/mem0/vector_stores/pgvector.py +++ b/mem0/vector_stores/pgvector.py @@ -181,6 +181,7 @@ class PGVector(VectorStoreBase): self.use_hnsw = hnsw self.embedding_model_dims = embedding_model_dims self.connection_pool = None + self._collection_ensured = False # Connection setup with priority: connection_pool > connection_string > individual parameters if connection_pool is not None: @@ -196,15 +197,25 @@ class PGVector(VectorStoreBase): 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) + # open=False avoids blocking when DB DNS is not yet resolvable (e.g. Docker startup) + self.connection_pool = ConnectionPool( + conninfo=connection_string, + min_size=minconn, + max_size=maxconn, + open=False, + ) + self.connection_pool.open(wait=False) else: # psycopg2 ThreadedConnectionPool self.connection_pool = ConnectionPool(minconn=minconn, maxconn=maxconn, dsn=connection_string) + def _ensure_collection(self): + if self._collection_ensured: + return collections = self.list_cols() - if collection_name not in collections: + if self.collection_name not in collections: self.create_col() + self._collection_ensured = True @contextmanager def _get_cursor(self, commit: bool = False): @@ -294,6 +305,7 @@ class PGVector(VectorStoreBase): ) def insert(self, vectors: list[list[float]], payloads=None, ids=None) -> None: + self._ensure_collection() logger.info(f"Inserting {len(vectors)} vectors into collection {self.collection_name}") json_payloads = [json.dumps(payload) for payload in payloads] @@ -331,6 +343,7 @@ class PGVector(VectorStoreBase): Returns: list: Search results. """ + self._ensure_collection() filter_conditions, filter_params = _build_filter_conditions(filters) filter_clause = sql.SQL("WHERE " + " AND ".join(filter_conditions)) if filter_conditions else sql.SQL("") @@ -361,6 +374,7 @@ class PGVector(VectorStoreBase): Returns: List[OutputData]: Search results ranked by text relevance. """ + self._ensure_collection() filter_conditions, filter_params = _build_filter_conditions(filters) filter_clause = sql.SQL("AND " + " AND ".join(filter_conditions)) if filter_conditions else sql.SQL("") @@ -391,6 +405,7 @@ class PGVector(VectorStoreBase): Args: vector_id (str): ID of the vector to delete. """ + self._ensure_collection() with self._get_cursor(commit=True) as cur: cur.execute(sql.SQL("DELETE FROM {} WHERE id = %s").format(self._col()), (vector_id,)) @@ -408,6 +423,7 @@ class PGVector(VectorStoreBase): vector (List[float], optional): Updated vector. payload (Dict, optional): Updated payload. """ + self._ensure_collection() with self._get_cursor(commit=True) as cur: if vector: cur.execute( @@ -440,6 +456,7 @@ class PGVector(VectorStoreBase): Returns: OutputData: Retrieved vector. """ + self._ensure_collection() with self._get_cursor() as cur: cur.execute( sql.SQL("SELECT id, vector, payload FROM {} WHERE id = %s").format(self._col()), @@ -473,6 +490,7 @@ class PGVector(VectorStoreBase): Returns: Dict[str, Any]: Collection information. """ + self._ensure_collection() with self._get_cursor() as cur: cur.execute( sql.SQL(""" @@ -503,6 +521,7 @@ class PGVector(VectorStoreBase): Returns: List[OutputData]: List of vectors. """ + self._ensure_collection() filter_conditions, filter_params = _build_filter_conditions(filters) filter_clause = sql.SQL("WHERE " + " AND ".join(filter_conditions)) if filter_conditions else sql.SQL("") @@ -534,6 +553,7 @@ class PGVector(VectorStoreBase): def reset(self) -> None: """Reset the index by deleting and recreating it.""" + self._ensure_collection() logger.warning(f"Resetting index {self.collection_name}...") self.delete_col() self.create_col() diff --git a/tests/vector_stores/test_pgvector.py b/tests/vector_stores/test_pgvector.py index 56632bb18..7a6e9f7c4 100644 --- a/tests/vector_stores/test_pgvector.py +++ b/tests/vector_stores/test_pgvector.py @@ -40,10 +40,9 @@ class TestPGVector(unittest.TestCase): @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_pool.return_value = self.mock_pool_psycopg - self.mock_cursor.fetchall.return_value = [] # No existing collections - + mock_pool_instance = MagicMock() + mock_psycopg_pool.return_value = mock_pool_instance + pgvector = PGVector( dbname="test_db", collection_name="test_collection", @@ -62,11 +61,56 @@ class TestPGVector(unittest.TestCase): conninfo="postgresql://test_user:test_pass@localhost:5432/test_db", min_size=1, max_size=4, - open=True, + open=False, ) + mock_pool_instance.open.assert_called_once_with(wait=False) + # No DB calls during __init__ — collection setup is deferred. + mock_pool_instance.connection.assert_not_called() + 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') + def test_init_does_not_block_on_unreachable_host_psycopg3(self, mock_psycopg_pool): + """Regression test for issue #3950. + + PGVector.__init__ must NOT block waiting for connections when the DB + host is temporarily unreachable (e.g. Docker Compose startup race). + The pool is created with open=False and collection setup is deferred + to first use, so the constructor returns immediately. + """ + import time + + mock_pool_instance = MagicMock() + mock_psycopg_pool.return_value = mock_pool_instance + + start = time.monotonic() + pv = PGVector( + dbname="test_db", + collection_name="test_collection", + embedding_model_dims=3, + user="test_user", + password="test_pass", + host="unreachable-docker-host", + port=5432, + diskann=False, + hnsw=False, + ) + elapsed = time.monotonic() - start + + self.assertLess(elapsed, 1.0, "PGVector.__init__ blocked waiting for connections") + + # Pool created with open=False, then opened non-blocking. + mock_psycopg_pool.assert_called_once() + call_kwargs = mock_psycopg_pool.call_args.kwargs + self.assertFalse(call_kwargs.get("open", True), "Pool must be created with open=False") + mock_pool_instance.open.assert_called_once_with(wait=False) + + # No DB calls during __init__ — collection setup is deferred. + mock_pool_instance.connection.assert_not_called() + self.assertFalse(pv._collection_ensured) + @patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 2) @patch('mem0.vector_stores.pgvector.ConnectionPool') def test_init_with_individual_params_psycopg2(self, mock_pcycopg2_pool): @@ -125,8 +169,10 @@ class TestPGVector(unittest.TestCase): minconn=1, maxconn=4 ) - - # Verify the _get_cursor context manager was called + + # Collection setup is deferred — trigger it explicitly. + pgvector._ensure_collection() + mock_get_cursor.assert_called() # Verify vector extension and table creation @@ -135,7 +181,6 @@ class TestPGVector(unittest.TestCase): if "CREATE TABLE IF NOT EXISTS" in str(call) and "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) @@ -147,19 +192,14 @@ class TestPGVector(unittest.TestCase): 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", @@ -175,22 +215,19 @@ class TestPGVector(unittest.TestCase): connection_pool=explicit_pool ) - # Verify the _get_cursor context manager was called - mock_get_cursor.assert_called() + # Collection setup is deferred — trigger it explicitly. + pgvector._ensure_collection() + 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" in str(call) and "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) @@ -201,19 +238,14 @@ class TestPGVector(unittest.TestCase): 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", @@ -229,21 +261,19 @@ class TestPGVector(unittest.TestCase): connection_pool=explicit_pool ) - # Verify the _get_cursor context manager was called - mock_get_cursor.assert_called() + # Collection setup is deferred — trigger it explicitly. + pgvector._ensure_collection() + 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 + table_creation_calls = [call for call in self.mock_cursor.execute.call_args_list if "CREATE TABLE IF NOT EXISTS" in str(call) and "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) @@ -251,16 +281,14 @@ class TestPGVector(unittest.TestCase): @patch.object(PGVector, '_get_cursor') def test_create_col_psycopg2(self, mock_get_cursor, mock_connection_pool): """Test collection creation with psycopg2.""" - # 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.fetchall.return_value = [] + pgvector = PGVector( dbname="test_db", collection_name="test_collection", @@ -274,17 +302,17 @@ class TestPGVector(unittest.TestCase): minconn=1, maxconn=4 ) - - # Verify the _get_cursor context manager was called + + # Collection setup is deferred — trigger it explicitly. + pgvector._ensure_collection() + 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 + table_creation_calls = [call for call in self.mock_cursor.execute.call_args_list if "CREATE TABLE IF NOT EXISTS" in str(call) and "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) @@ -2068,10 +2096,9 @@ class TestPGVector(unittest.TestCase): """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", @@ -2087,15 +2114,17 @@ class TestPGVector(unittest.TestCase): sslmode="require", connection_string=connection_string ) - + # Verify ConnectionPool was called with sslmode as a URI query parameter + # and open=False to avoid blocking __init__ in Docker (issue #3950). 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 + open=False ) + mock_pool.open.assert_called_once_with(wait=False) self.assertEqual(pgvector.collection_name, "test_collection") self.assertEqual(pgvector.embedding_model_dims, 3) @@ -2160,9 +2189,11 @@ class TestPGVector(unittest.TestCase): minconn=1, maxconn=4 ) - - # Verify DiskANN index creation query was executed - diskann_calls = [call for call in self.mock_cursor.execute.call_args_list + + # Collection setup is deferred — trigger it explicitly. + pgvector._ensure_collection() + + 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") @@ -2197,9 +2228,11 @@ class TestPGVector(unittest.TestCase): minconn=1, maxconn=4 ) - - # Verify HNSW index creation query was executed - hnsw_calls = [call for call in self.mock_cursor.execute.call_args_list + + # Collection setup is deferred — trigger it explicitly. + pgvector._ensure_collection() + + 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")