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:
@@ -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")
|
||||
|
||||
Reference in New Issue
Block a user