fix(pgvector): use open=False to prevent ConnectionPool hang in Docker (#5155)

Co-authored-by: kartik-mem0 <kartik.labhshetwar@mem0.ai>
This commit is contained in:
shafdev
2026-06-10 17:15:53 +05:30
committed by GitHub
parent 3ac1c9452c
commit b819d95d18
2 changed files with 110 additions and 57 deletions
+23 -3
View File
@@ -181,6 +181,7 @@ class PGVector(VectorStoreBase):
self.use_hnsw = hnsw self.use_hnsw = hnsw
self.embedding_model_dims = embedding_model_dims self.embedding_model_dims = embedding_model_dims
self.connection_pool = None self.connection_pool = None
self._collection_ensured = False
# Connection setup with priority: connection_pool > connection_string > individual parameters # Connection setup with priority: connection_pool > connection_string > individual parameters
if connection_pool is not None: if connection_pool is not None:
@@ -196,15 +197,25 @@ class PGVector(VectorStoreBase):
if self.connection_pool is None: if self.connection_pool is None:
if PSYCOPG_VERSION == 3: if PSYCOPG_VERSION == 3:
# psycopg3 ConnectionPool # 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=True) self.connection_pool = ConnectionPool(
conninfo=connection_string,
min_size=minconn,
max_size=maxconn,
open=False,
)
self.connection_pool.open(wait=False)
else: else:
# psycopg2 ThreadedConnectionPool # psycopg2 ThreadedConnectionPool
self.connection_pool = ConnectionPool(minconn=minconn, maxconn=maxconn, dsn=connection_string) self.connection_pool = ConnectionPool(minconn=minconn, maxconn=maxconn, dsn=connection_string)
def _ensure_collection(self):
if self._collection_ensured:
return
collections = self.list_cols() collections = self.list_cols()
if collection_name not in collections: if self.collection_name not in collections:
self.create_col() self.create_col()
self._collection_ensured = True
@contextmanager @contextmanager
def _get_cursor(self, commit: bool = False): 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: 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}") logger.info(f"Inserting {len(vectors)} vectors into collection {self.collection_name}")
json_payloads = [json.dumps(payload) for payload in payloads] json_payloads = [json.dumps(payload) for payload in payloads]
@@ -331,6 +343,7 @@ class PGVector(VectorStoreBase):
Returns: Returns:
list: Search results. list: Search results.
""" """
self._ensure_collection()
filter_conditions, filter_params = _build_filter_conditions(filters) filter_conditions, filter_params = _build_filter_conditions(filters)
filter_clause = sql.SQL("WHERE " + " AND ".join(filter_conditions)) if filter_conditions else sql.SQL("") filter_clause = sql.SQL("WHERE " + " AND ".join(filter_conditions)) if filter_conditions else sql.SQL("")
@@ -361,6 +374,7 @@ class PGVector(VectorStoreBase):
Returns: Returns:
List[OutputData]: Search results ranked by text relevance. List[OutputData]: Search results ranked by text relevance.
""" """
self._ensure_collection()
filter_conditions, filter_params = _build_filter_conditions(filters) filter_conditions, filter_params = _build_filter_conditions(filters)
filter_clause = sql.SQL("AND " + " AND ".join(filter_conditions)) if filter_conditions else sql.SQL("") filter_clause = sql.SQL("AND " + " AND ".join(filter_conditions)) if filter_conditions else sql.SQL("")
@@ -391,6 +405,7 @@ class PGVector(VectorStoreBase):
Args: Args:
vector_id (str): ID of the vector to delete. vector_id (str): ID of the vector to delete.
""" """
self._ensure_collection()
with self._get_cursor(commit=True) as cur: with self._get_cursor(commit=True) as cur:
cur.execute(sql.SQL("DELETE FROM {} WHERE id = %s").format(self._col()), (vector_id,)) 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. vector (List[float], optional): Updated vector.
payload (Dict, optional): Updated payload. payload (Dict, optional): Updated payload.
""" """
self._ensure_collection()
with self._get_cursor(commit=True) as cur: with self._get_cursor(commit=True) as cur:
if vector: if vector:
cur.execute( cur.execute(
@@ -440,6 +456,7 @@ class PGVector(VectorStoreBase):
Returns: Returns:
OutputData: Retrieved vector. OutputData: Retrieved vector.
""" """
self._ensure_collection()
with self._get_cursor() as cur: with self._get_cursor() as cur:
cur.execute( cur.execute(
sql.SQL("SELECT id, vector, payload FROM {} WHERE id = %s").format(self._col()), sql.SQL("SELECT id, vector, payload FROM {} WHERE id = %s").format(self._col()),
@@ -473,6 +490,7 @@ class PGVector(VectorStoreBase):
Returns: Returns:
Dict[str, Any]: Collection information. Dict[str, Any]: Collection information.
""" """
self._ensure_collection()
with self._get_cursor() as cur: with self._get_cursor() as cur:
cur.execute( cur.execute(
sql.SQL(""" sql.SQL("""
@@ -503,6 +521,7 @@ class PGVector(VectorStoreBase):
Returns: Returns:
List[OutputData]: List of vectors. List[OutputData]: List of vectors.
""" """
self._ensure_collection()
filter_conditions, filter_params = _build_filter_conditions(filters) filter_conditions, filter_params = _build_filter_conditions(filters)
filter_clause = sql.SQL("WHERE " + " AND ".join(filter_conditions)) if filter_conditions else sql.SQL("") 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: def reset(self) -> None:
"""Reset the index by deleting and recreating it.""" """Reset the index by deleting and recreating it."""
self._ensure_collection()
logger.warning(f"Resetting index {self.collection_name}...") logger.warning(f"Resetting index {self.collection_name}...")
self.delete_col() self.delete_col()
self.create_col() self.create_col()
+87 -54
View File
@@ -40,10 +40,9 @@ class TestPGVector(unittest.TestCase):
@patch('mem0.vector_stores.pgvector.ConnectionPool') @patch('mem0.vector_stores.pgvector.ConnectionPool')
def test_init_with_individual_params_psycopg3(self, mock_psycopg_pool): def test_init_with_individual_params_psycopg3(self, mock_psycopg_pool):
"""Test initialization with individual parameters using psycopg3.""" """Test initialization with individual parameters using psycopg3."""
# Mock psycopg3 to be available mock_pool_instance = MagicMock()
mock_psycopg_pool.return_value = self.mock_pool_psycopg mock_psycopg_pool.return_value = mock_pool_instance
self.mock_cursor.fetchall.return_value = [] # No existing collections
pgvector = PGVector( pgvector = PGVector(
dbname="test_db", dbname="test_db",
collection_name="test_collection", collection_name="test_collection",
@@ -62,11 +61,56 @@ class TestPGVector(unittest.TestCase):
conninfo="postgresql://test_user:test_pass@localhost:5432/test_db", conninfo="postgresql://test_user:test_pass@localhost:5432/test_db",
min_size=1, min_size=1,
max_size=4, 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.collection_name, "test_collection")
self.assertEqual(pgvector.embedding_model_dims, 3) 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.PSYCOPG_VERSION', 2)
@patch('mem0.vector_stores.pgvector.ConnectionPool') @patch('mem0.vector_stores.pgvector.ConnectionPool')
def test_init_with_individual_params_psycopg2(self, mock_pcycopg2_pool): def test_init_with_individual_params_psycopg2(self, mock_pcycopg2_pool):
@@ -125,8 +169,10 @@ class TestPGVector(unittest.TestCase):
minconn=1, minconn=1,
maxconn=4 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() mock_get_cursor.assert_called()
# Verify vector extension and table creation # 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)] if "CREATE TABLE IF NOT EXISTS" in str(call) and "test_collection" in str(call)]
self.assertTrue(len(table_creation_calls) > 0) self.assertTrue(len(table_creation_calls) > 0)
# Verify pgvector instance properties
self.assertEqual(pgvector.collection_name, "test_collection") self.assertEqual(pgvector.collection_name, "test_collection")
self.assertEqual(pgvector.embedding_model_dims, 3) 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. 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. 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") 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") 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.__enter__.return_value = self.mock_cursor
mock_get_cursor.return_value.__exit__.return_value = None mock_get_cursor.return_value.__exit__.return_value = None
# Simulate no existing collections in the database
self.mock_cursor.fetchall.return_value = [] self.mock_cursor.fetchall.return_value = []
# Pass the explicit pool to PGVector
pgvector = PGVector( pgvector = PGVector(
dbname="test_db", dbname="test_db",
collection_name="test_collection", collection_name="test_collection",
@@ -175,22 +215,19 @@ class TestPGVector(unittest.TestCase):
connection_pool=explicit_pool connection_pool=explicit_pool
) )
# Verify the _get_cursor context manager was called # Collection setup is deferred — trigger it explicitly.
mock_get_cursor.assert_called() pgvector._ensure_collection()
mock_get_cursor.assert_called()
mock_connection_pool.assert_not_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") 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)] if "CREATE TABLE IF NOT EXISTS" in str(call) and "test_collection" in str(call)]
self.assertTrue(len(table_creation_calls) > 0) self.assertTrue(len(table_creation_calls) > 0)
# Verify pgvector instance properties
self.assertEqual(pgvector.collection_name, "test_collection") self.assertEqual(pgvector.collection_name, "test_collection")
self.assertEqual(pgvector.embedding_model_dims, 3) self.assertEqual(pgvector.embedding_model_dims, 3)
# Ensure the pool used is the explicit one
self.assertIs(pgvector.connection_pool, explicit_pool) self.assertIs(pgvector.connection_pool, explicit_pool)
@patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 2) @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. 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. 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") 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") 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.__enter__.return_value = self.mock_cursor
mock_get_cursor.return_value.__exit__.return_value = None mock_get_cursor.return_value.__exit__.return_value = None
# Simulate no existing collections in the database
self.mock_cursor.fetchall.return_value = [] self.mock_cursor.fetchall.return_value = []
# Pass the explicit pool to PGVector
pgvector = PGVector( pgvector = PGVector(
dbname="test_db", dbname="test_db",
collection_name="test_collection", collection_name="test_collection",
@@ -229,21 +261,19 @@ class TestPGVector(unittest.TestCase):
connection_pool=explicit_pool connection_pool=explicit_pool
) )
# Verify the _get_cursor context manager was called # Collection setup is deferred — trigger it explicitly.
mock_get_cursor.assert_called() pgvector._ensure_collection()
mock_get_cursor.assert_called()
mock_connection_pool.assert_not_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") 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)] if "CREATE TABLE IF NOT EXISTS" in str(call) and "test_collection" in str(call)]
self.assertTrue(len(table_creation_calls) > 0) self.assertTrue(len(table_creation_calls) > 0)
# Verify pgvector instance properties
self.assertEqual(pgvector.collection_name, "test_collection") self.assertEqual(pgvector.collection_name, "test_collection")
self.assertEqual(pgvector.embedding_model_dims, 3) self.assertEqual(pgvector.embedding_model_dims, 3)
# Ensure the pool used is the explicit one
self.assertIs(pgvector.connection_pool, explicit_pool) self.assertIs(pgvector.connection_pool, explicit_pool)
@patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 2) @patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 2)
@@ -251,16 +281,14 @@ class TestPGVector(unittest.TestCase):
@patch.object(PGVector, '_get_cursor') @patch.object(PGVector, '_get_cursor')
def test_create_col_psycopg2(self, mock_get_cursor, mock_connection_pool): def test_create_col_psycopg2(self, mock_get_cursor, mock_connection_pool):
"""Test collection creation with psycopg2.""" """Test collection creation with psycopg2."""
# Set up mock pool and cursor
mock_pool = MagicMock() mock_pool = MagicMock()
mock_connection_pool.return_value = mock_pool 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.__enter__.return_value = self.mock_cursor
mock_get_cursor.return_value.__exit__.return_value = None 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( pgvector = PGVector(
dbname="test_db", dbname="test_db",
collection_name="test_collection", collection_name="test_collection",
@@ -274,17 +302,17 @@ class TestPGVector(unittest.TestCase):
minconn=1, minconn=1,
maxconn=4 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() mock_get_cursor.assert_called()
# Verify vector extension and table creation
self.mock_cursor.execute.assert_any_call("CREATE EXTENSION IF NOT EXISTS vector") 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)] if "CREATE TABLE IF NOT EXISTS" in str(call) and "test_collection" in str(call)]
self.assertTrue(len(table_creation_calls) > 0) self.assertTrue(len(table_creation_calls) > 0)
# Verify pgvector instance properties
self.assertEqual(pgvector.collection_name, "test_collection") self.assertEqual(pgvector.collection_name, "test_collection")
self.assertEqual(pgvector.embedding_model_dims, 3) self.assertEqual(pgvector.embedding_model_dims, 3)
@@ -2068,10 +2096,9 @@ class TestPGVector(unittest.TestCase):
"""Test connection string handling with SSL mode.""" """Test connection string handling with SSL mode."""
mock_pool = MagicMock() mock_pool = MagicMock()
mock_connection_pool.return_value = mock_pool mock_connection_pool.return_value = mock_pool
self.mock_cursor.fetchall.return_value = [] # No existing collections
connection_string = "postgresql://user:pass@localhost:5432/db" connection_string = "postgresql://user:pass@localhost:5432/db"
pgvector = PGVector( pgvector = PGVector(
dbname="test_db", # Will be overridden by connection_string dbname="test_db", # Will be overridden by connection_string
collection_name="test_collection", collection_name="test_collection",
@@ -2087,15 +2114,17 @@ class TestPGVector(unittest.TestCase):
sslmode="require", sslmode="require",
connection_string=connection_string connection_string=connection_string
) )
# Verify ConnectionPool was called with sslmode as a URI query parameter # 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" expected_conn_string = f"{connection_string}?sslmode=require"
mock_connection_pool.assert_called_with( mock_connection_pool.assert_called_with(
conninfo=expected_conn_string, conninfo=expected_conn_string,
min_size=1, min_size=1,
max_size=4, 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.collection_name, "test_collection")
self.assertEqual(pgvector.embedding_model_dims, 3) self.assertEqual(pgvector.embedding_model_dims, 3)
@@ -2160,9 +2189,11 @@ class TestPGVector(unittest.TestCase):
minconn=1, minconn=1,
maxconn=4 maxconn=4
) )
# Verify DiskANN index creation query was executed # Collection setup is deferred — trigger it explicitly.
diskann_calls = [call for call in self.mock_cursor.execute.call_args_list pgvector._ensure_collection()
diskann_calls = [call for call in self.mock_cursor.execute.call_args_list
if "USING diskann" in str(call)] if "USING diskann" in str(call)]
self.assertTrue(len(diskann_calls) > 0) self.assertTrue(len(diskann_calls) > 0)
self.assertEqual(pgvector.collection_name, "test_collection") self.assertEqual(pgvector.collection_name, "test_collection")
@@ -2197,9 +2228,11 @@ class TestPGVector(unittest.TestCase):
minconn=1, minconn=1,
maxconn=4 maxconn=4
) )
# Verify HNSW index creation query was executed # Collection setup is deferred — trigger it explicitly.
hnsw_calls = [call for call in self.mock_cursor.execute.call_args_list pgvector._ensure_collection()
hnsw_calls = [call for call in self.mock_cursor.execute.call_args_list
if "USING hnsw" in str(call)] if "USING hnsw" in str(call)]
self.assertTrue(len(hnsw_calls) > 0) self.assertTrue(len(hnsw_calls) > 0)
self.assertEqual(pgvector.collection_name, "test_collection") self.assertEqual(pgvector.collection_name, "test_collection")