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
+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")