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