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.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()
+87 -54
View File
@@ -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")