From 148bbf0a5cc6123496eebc20b14f2ecfe5eb7125 Mon Sep 17 00:00:00 2001 From: Parshva Daftari <89991302+parshvadaftari@users.noreply.github.com> Date: Wed, 6 Aug 2025 22:58:02 +0530 Subject: [PATCH] Add sslmode pgvector (#3265) --- docs/components/vectordbs/dbs/pgvector.mdx | 12 +- mem0/configs/vector_stores/pgvector.py | 17 +- mem0/vector_stores/pgvector.py | 58 ++- tests/vector_stores/test_pgvector.py | 434 +++++++++++++++++++++ 4 files changed, 509 insertions(+), 12 deletions(-) diff --git a/docs/components/vectordbs/dbs/pgvector.mdx b/docs/components/vectordbs/dbs/pgvector.mdx index 011a3f982..6a083af5c 100644 --- a/docs/components/vectordbs/dbs/pgvector.mdx +++ b/docs/components/vectordbs/dbs/pgvector.mdx @@ -36,7 +36,7 @@ Here's the parameters available for configuring pgvector: | Parameter | Description | Default Value | | --- | --- | --- | -| `dbname` | The name of the | `postgres` | +| `dbname` | The name of the database | `postgres` | | `collection_name` | The name of the collection | `mem0` | | `embedding_model_dims` | Dimensions of the embedding model | `1536` | | `user` | User name to connect to the database | `None` | @@ -44,4 +44,12 @@ Here's the parameters available for configuring pgvector: | `host` | The host where the Postgres server is running | `None` | | `port` | The port where the Postgres server is running | `None` | | `diskann` | Whether to use diskann for vector similarity search (requires pgvectorscale) | `True` | -| `hnsw` | Whether to use hnsw for vector similarity search | `False` | \ No newline at end of file +| `hnsw` | Whether to use hnsw for vector similarity search | `False` | +| `sslmode` | SSL mode for PostgreSQL connection (e.g., 'require', 'prefer', 'disable') | `None` | +| `connection_string` | PostgreSQL connection string (overrides individual connection parameters) | `None` | +| `connection_pool` | psycopg2 connection pool object (overrides connection string and individual parameters) | `None` | + +**Note**: The connection parameters have the following priority: +1. `connection_pool` (highest priority) +2. `connection_string` +3. Individual connection parameters (`user`, `password`, `host`, `port`, `sslmode`) \ No newline at end of file diff --git a/mem0/configs/vector_stores/pgvector.py b/mem0/configs/vector_stores/pgvector.py index c9047d5f4..b0cca30a9 100644 --- a/mem0/configs/vector_stores/pgvector.py +++ b/mem0/configs/vector_stores/pgvector.py @@ -13,15 +13,28 @@ class PGVectorConfig(BaseModel): port: Optional[int] = Field(None, description="Database port. Default is 1536") diskann: Optional[bool] = Field(True, description="Use diskann for approximate nearest neighbors search") hnsw: Optional[bool] = Field(False, description="Use hnsw for faster search") + # New SSL and connection options + sslmode: Optional[str] = Field(None, description="SSL mode for PostgreSQL connection (e.g., 'require', 'prefer', 'disable')") + connection_string: Optional[str] = Field(None, description="PostgreSQL connection string (overrides individual connection parameters)") + connection_pool: Optional[Any] = Field(None, description="psycopg2 connection pool object (overrides connection string and individual parameters)") @model_validator(mode="before") def check_auth_and_connection(cls, values): + # If connection_pool is provided, skip validation of individual connection parameters + if values.get("connection_pool") is not None: + return values + + # If connection_string is provided, skip validation of individual connection parameters + if values.get("connection_string") is not None: + return values + + # Otherwise, validate individual connection parameters user, password = values.get("user"), values.get("password") host, port = values.get("host"), values.get("port") if not user and not password: - raise ValueError("Both 'user' and 'password' must be provided.") + raise ValueError("Both 'user' and 'password' must be provided when not using connection_string or connection_pool.") if not host and not port: - raise ValueError("Both 'host' and 'port' must be provided.") + raise ValueError("Both 'host' and 'port' must be provided when not using connection_string or connection_pool.") return values @model_validator(mode="before") diff --git a/mem0/vector_stores/pgvector.py b/mem0/vector_stores/pgvector.py index d88d60ec2..062b4ecd8 100644 --- a/mem0/vector_stores/pgvector.py +++ b/mem0/vector_stores/pgvector.py @@ -48,6 +48,9 @@ class PGVector(VectorStoreBase): port, diskann, hnsw, + sslmode=None, + connection_string=None, + connection_pool=None, ): """ Initialize the PGVector database. @@ -62,20 +65,54 @@ class PGVector(VectorStoreBase): port (int, optional): Database port diskann (bool, optional): Use DiskANN for faster search hnsw (bool, optional): Use HNSW for faster search + sslmode (str, optional): SSL mode for PostgreSQL connection (e.g., 'require', 'prefer', 'disable') + connection_string (str, optional): PostgreSQL connection string (overrides individual connection parameters) + connection_pool (Any, optional): psycopg2 connection pool object (overrides connection string and individual parameters) """ self.collection_name = collection_name self.use_diskann = diskann self.use_hnsw = hnsw self.embedding_model_dims = embedding_model_dims - if PSYCOPG_VERSION == 3: - self.conn = psycopg.connect( - dbname=dbname, user=user, password=password, host=host, port=port - ) + # Connection setup with priority: connection_pool > connection_string > individual parameters + if connection_pool is not None: + # Use provided connection pool + self.conn = connection_pool.getconn() + self.connection_pool = connection_pool + elif connection_string is not None: + # Use connection string + if sslmode: + # Append sslmode to connection string if provided + if 'sslmode=' in connection_string: + # Replace existing sslmode + import re + connection_string = re.sub(r'sslmode=[^ ]*', f'sslmode={sslmode}', connection_string) + else: + # Add sslmode to connection string + connection_string = f"{connection_string} sslmode={sslmode}" + + if PSYCOPG_VERSION == 3: + self.conn = psycopg.connect(connection_string) + else: + self.conn = psycopg2.connect(connection_string) + self.connection_pool = None else: - self.conn = psycopg2.connect( - dbname=dbname, user=user, password=password, host=host, port=port - ) + # Use individual connection parameters + conn_params = { + 'dbname': dbname, + 'user': user, + 'password': password, + 'host': host, + 'port': port + } + if sslmode: + conn_params['sslmode'] = sslmode + + if PSYCOPG_VERSION == 3: + self.conn = psycopg.connect(**conn_params) + else: + self.conn = psycopg2.connect(**conn_params) + self.connection_pool = None self.cur = self.conn.cursor() @@ -317,7 +354,12 @@ class PGVector(VectorStoreBase): if hasattr(self, "cur"): self.cur.close() if hasattr(self, "conn"): - self.conn.close() + if hasattr(self, "connection_pool") and self.connection_pool is not None: + # Return connection to pool instead of closing it + self.connection_pool.putconn(self.conn) + else: + # Close the connection directly + self.conn.close() def reset(self): """Reset the index by deleting and recreating it.""" diff --git a/tests/vector_stores/test_pgvector.py b/tests/vector_stores/test_pgvector.py index ac18a3002..2a66559b5 100644 --- a/tests/vector_stores/test_pgvector.py +++ b/tests/vector_stores/test_pgvector.py @@ -706,6 +706,440 @@ class TestPGVector(unittest.TestCase): self.assertEqual(results[0][0].id, self.test_ids[0]) self.assertEqual(results[0][1].id, self.test_ids[1]) + @patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 3) + @patch('mem0.vector_stores.pgvector.psycopg.connect') + @patch('mem0.vector_stores.pgvector.psycopg2.connect') + def test_search_with_filters_psycopg3(self, mock_psycopg2_connect, mock_psycopg_connect): + """Test search with filters using psycopg3.""" + mock_psycopg_connect.return_value = self.mock_conn + self.mock_cursor.fetchall.return_value = [ + (self.test_ids[0], 0.1, {"user_id": "alice", "agent_id": "agent1", "run_id": "run1"}), + ] + + pgvector = PGVector( + dbname="test_db", + collection_name="test_collection", + embedding_model_dims=3, + user="test_user", + password="test_pass", + host="localhost", + port=5432, + diskann=False, + hnsw=False + ) + + filters = {"user_id": "alice", "agent_id": "agent1", "run_id": "run1"} + results = pgvector.search("test query", [0.1, 0.2, 0.3], limit=2, filters=filters) + + # Verify search query was executed with filters + search_calls = [call for call in self.mock_cursor.execute.call_args_list + if "SELECT id, vector <=" in str(call) and "WHERE" in str(call)] + self.assertTrue(len(search_calls) > 0) + + # Verify results + self.assertEqual(len(results), 1) + self.assertEqual(results[0].id, self.test_ids[0]) + self.assertEqual(results[0].score, 0.1) + self.assertEqual(results[0].payload["user_id"], "alice") + self.assertEqual(results[0].payload["agent_id"], "agent1") + self.assertEqual(results[0].payload["run_id"], "run1") + + @patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 2) + @patch('mem0.vector_stores.pgvector.psycopg2.connect') + def test_search_with_filters_psycopg2(self, mock_connect): + """Test search with filters using psycopg2.""" + mock_connect.return_value = self.mock_conn + self.mock_cursor.fetchall.return_value = [ + (self.test_ids[0], 0.1, {"user_id": "alice", "agent_id": "agent1", "run_id": "run1"}), + ] + + pgvector = PGVector( + dbname="test_db", + collection_name="test_collection", + embedding_model_dims=3, + user="test_user", + password="test_pass", + host="localhost", + port=5432, + diskann=False, + hnsw=False + ) + + filters = {"user_id": "alice", "agent_id": "agent1", "run_id": "run1"} + results = pgvector.search("test query", [0.1, 0.2, 0.3], limit=2, filters=filters) + + # Verify search query was executed with filters + search_calls = [call for call in self.mock_cursor.execute.call_args_list + if "SELECT id, vector <=" in str(call) and "WHERE" in str(call)] + self.assertTrue(len(search_calls) > 0) + + # Verify results + self.assertEqual(len(results), 1) + self.assertEqual(results[0].id, self.test_ids[0]) + self.assertEqual(results[0].score, 0.1) + self.assertEqual(results[0].payload["user_id"], "alice") + self.assertEqual(results[0].payload["agent_id"], "agent1") + self.assertEqual(results[0].payload["run_id"], "run1") + + @patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 3) + @patch('mem0.vector_stores.pgvector.psycopg.connect') + @patch('mem0.vector_stores.pgvector.psycopg2.connect') + def test_search_with_single_filter_psycopg3(self, mock_psycopg2_connect, mock_psycopg_connect): + """Test search with single filter using psycopg3.""" + mock_psycopg_connect.return_value = self.mock_conn + self.mock_cursor.fetchall.return_value = [ + (self.test_ids[0], 0.1, {"user_id": "alice"}), + ] + + pgvector = PGVector( + dbname="test_db", + collection_name="test_collection", + embedding_model_dims=3, + user="test_user", + password="test_pass", + host="localhost", + port=5432, + diskann=False, + hnsw=False + ) + + filters = {"user_id": "alice"} + results = pgvector.search("test query", [0.1, 0.2, 0.3], limit=2, filters=filters) + + # Verify search query was executed with single filter + search_calls = [call for call in self.mock_cursor.execute.call_args_list + if "SELECT id, vector <=" in str(call) and "WHERE" in str(call)] + self.assertTrue(len(search_calls) > 0) + + # Verify results + self.assertEqual(len(results), 1) + self.assertEqual(results[0].id, self.test_ids[0]) + self.assertEqual(results[0].score, 0.1) + self.assertEqual(results[0].payload["user_id"], "alice") + + @patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 2) + @patch('mem0.vector_stores.pgvector.psycopg2.connect') + def test_search_with_single_filter_psycopg2(self, mock_connect): + """Test search with single filter using psycopg2.""" + mock_connect.return_value = self.mock_conn + self.mock_cursor.fetchall.return_value = [ + (self.test_ids[0], 0.1, {"user_id": "alice"}), + ] + + pgvector = PGVector( + dbname="test_db", + collection_name="test_collection", + embedding_model_dims=3, + user="test_user", + password="test_pass", + host="localhost", + port=5432, + diskann=False, + hnsw=False + ) + + filters = {"user_id": "alice"} + results = pgvector.search("test query", [0.1, 0.2, 0.3], limit=2, filters=filters) + + # Verify search query was executed with single filter + search_calls = [call for call in self.mock_cursor.execute.call_args_list + if "SELECT id, vector <=" in str(call) and "WHERE" in str(call)] + self.assertTrue(len(search_calls) > 0) + + # Verify results + self.assertEqual(len(results), 1) + self.assertEqual(results[0].id, self.test_ids[0]) + self.assertEqual(results[0].score, 0.1) + self.assertEqual(results[0].payload["user_id"], "alice") + + @patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 3) + @patch('mem0.vector_stores.pgvector.psycopg.connect') + @patch('mem0.vector_stores.pgvector.psycopg2.connect') + def test_search_with_no_filters_psycopg3(self, mock_psycopg2_connect, mock_psycopg_connect): + """Test search with no filters using psycopg3.""" + mock_psycopg_connect.return_value = self.mock_conn + self.mock_cursor.fetchall.return_value = [ + (self.test_ids[0], 0.1, {"key": "value1"}), + (self.test_ids[1], 0.2, {"key": "value2"}), + ] + + pgvector = PGVector( + dbname="test_db", + collection_name="test_collection", + embedding_model_dims=3, + user="test_user", + password="test_pass", + host="localhost", + port=5432, + diskann=False, + hnsw=False + ) + + results = pgvector.search("test query", [0.1, 0.2, 0.3], limit=2, filters=None) + + # Verify search query was executed without WHERE clause + search_calls = [call for call in self.mock_cursor.execute.call_args_list + if "SELECT id, vector <=" in str(call) and "WHERE" not in str(call)] + self.assertTrue(len(search_calls) > 0) + + # Verify results + self.assertEqual(len(results), 2) + self.assertEqual(results[0].id, self.test_ids[0]) + self.assertEqual(results[0].score, 0.1) + self.assertEqual(results[1].id, self.test_ids[1]) + self.assertEqual(results[1].score, 0.2) + + @patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 2) + @patch('mem0.vector_stores.pgvector.psycopg2.connect') + def test_search_with_no_filters_psycopg2(self, mock_connect): + """Test search with no filters using psycopg2.""" + mock_connect.return_value = self.mock_conn + self.mock_cursor.fetchall.return_value = [ + (self.test_ids[0], 0.1, {"key": "value1"}), + (self.test_ids[1], 0.2, {"key": "value2"}), + ] + + pgvector = PGVector( + dbname="test_db", + collection_name="test_collection", + embedding_model_dims=3, + user="test_user", + password="test_pass", + host="localhost", + port=5432, + diskann=False, + hnsw=False + ) + + results = pgvector.search("test query", [0.1, 0.2, 0.3], limit=2, filters=None) + + # Verify search query was executed without WHERE clause + search_calls = [call for call in self.mock_cursor.execute.call_args_list + if "SELECT id, vector <=" in str(call) and "WHERE" not in str(call)] + self.assertTrue(len(search_calls) > 0) + + # Verify results + self.assertEqual(len(results), 2) + self.assertEqual(results[0].id, self.test_ids[0]) + self.assertEqual(results[0].score, 0.1) + self.assertEqual(results[1].id, self.test_ids[1]) + self.assertEqual(results[1].score, 0.2) + + @patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 3) + @patch('mem0.vector_stores.pgvector.psycopg.connect') + @patch('mem0.vector_stores.pgvector.psycopg2.connect') + def test_list_with_filters_psycopg3(self, mock_psycopg2_connect, mock_psycopg_connect): + """Test list with filters using psycopg3.""" + mock_psycopg_connect.return_value = self.mock_conn + self.mock_cursor.fetchall.return_value = [ + (self.test_ids[0], [0.1, 0.2, 0.3], {"user_id": "alice", "agent_id": "agent1"}), + ] + + pgvector = PGVector( + dbname="test_db", + collection_name="test_collection", + embedding_model_dims=3, + user="test_user", + password="test_pass", + host="localhost", + port=5432, + diskann=False, + hnsw=False + ) + + filters = {"user_id": "alice", "agent_id": "agent1"} + results = pgvector.list(filters=filters, limit=2) + + # Verify list query was executed with filters + list_calls = [call for call in self.mock_cursor.execute.call_args_list + if "SELECT id, vector, payload" in str(call) and "WHERE" in str(call)] + self.assertTrue(len(list_calls) > 0) + + # Verify results + self.assertEqual(len(results), 1) # Returns list of lists + self.assertEqual(len(results[0]), 1) + self.assertEqual(results[0][0].id, self.test_ids[0]) + self.assertEqual(results[0][0].payload["user_id"], "alice") + self.assertEqual(results[0][0].payload["agent_id"], "agent1") + + @patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 2) + @patch('mem0.vector_stores.pgvector.psycopg2.connect') + def test_list_with_filters_psycopg2(self, mock_connect): + """Test list with filters using psycopg2.""" + mock_connect.return_value = self.mock_conn + self.mock_cursor.fetchall.return_value = [ + (self.test_ids[0], [0.1, 0.2, 0.3], {"user_id": "alice", "agent_id": "agent1"}), + ] + + pgvector = PGVector( + dbname="test_db", + collection_name="test_collection", + embedding_model_dims=3, + user="test_user", + password="test_pass", + host="localhost", + port=5432, + diskann=False, + hnsw=False + ) + + filters = {"user_id": "alice", "agent_id": "agent1"} + results = pgvector.list(filters=filters, limit=2) + + # Verify list query was executed with filters + list_calls = [call for call in self.mock_cursor.execute.call_args_list + if "SELECT id, vector, payload" in str(call) and "WHERE" in str(call)] + self.assertTrue(len(list_calls) > 0) + + # Verify results + self.assertEqual(len(results), 1) # Returns list of lists + self.assertEqual(len(results[0]), 1) + self.assertEqual(results[0][0].id, self.test_ids[0]) + self.assertEqual(results[0][0].payload["user_id"], "alice") + self.assertEqual(results[0][0].payload["agent_id"], "agent1") + + @patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 3) + @patch('mem0.vector_stores.pgvector.psycopg.connect') + @patch('mem0.vector_stores.pgvector.psycopg2.connect') + def test_list_with_single_filter_psycopg3(self, mock_psycopg2_connect, mock_psycopg_connect): + """Test list with single filter using psycopg3.""" + mock_psycopg_connect.return_value = self.mock_conn + self.mock_cursor.fetchall.return_value = [ + (self.test_ids[0], [0.1, 0.2, 0.3], {"user_id": "alice"}), + ] + + pgvector = PGVector( + dbname="test_db", + collection_name="test_collection", + embedding_model_dims=3, + user="test_user", + password="test_pass", + host="localhost", + port=5432, + diskann=False, + hnsw=False + ) + + filters = {"user_id": "alice"} + results = pgvector.list(filters=filters, limit=2) + + # Verify list query was executed with single filter + list_calls = [call for call in self.mock_cursor.execute.call_args_list + if "SELECT id, vector, payload" in str(call) and "WHERE" in str(call)] + self.assertTrue(len(list_calls) > 0) + + # Verify results + self.assertEqual(len(results), 1) # Returns list of lists + self.assertEqual(len(results[0]), 1) + self.assertEqual(results[0][0].id, self.test_ids[0]) + self.assertEqual(results[0][0].payload["user_id"], "alice") + + @patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 2) + @patch('mem0.vector_stores.pgvector.psycopg2.connect') + def test_list_with_single_filter_psycopg2(self, mock_connect): + """Test list with single filter using psycopg2.""" + mock_connect.return_value = self.mock_conn + self.mock_cursor.fetchall.return_value = [ + (self.test_ids[0], [0.1, 0.2, 0.3], {"user_id": "alice"}), + ] + + pgvector = PGVector( + dbname="test_db", + collection_name="test_collection", + embedding_model_dims=3, + user="test_user", + password="test_pass", + host="localhost", + port=5432, + diskann=False, + hnsw=False + ) + + filters = {"user_id": "alice"} + results = pgvector.list(filters=filters, limit=2) + + # Verify list query was executed with single filter + list_calls = [call for call in self.mock_cursor.execute.call_args_list + if "SELECT id, vector, payload" in str(call) and "WHERE" in str(call)] + self.assertTrue(len(list_calls) > 0) + + # Verify results + self.assertEqual(len(results), 1) # Returns list of lists + self.assertEqual(len(results[0]), 1) + self.assertEqual(results[0][0].id, self.test_ids[0]) + self.assertEqual(results[0][0].payload["user_id"], "alice") + + @patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 3) + @patch('mem0.vector_stores.pgvector.psycopg.connect') + @patch('mem0.vector_stores.pgvector.psycopg2.connect') + def test_list_with_no_filters_psycopg3(self, mock_psycopg2_connect, mock_psycopg_connect): + """Test list with no filters using psycopg3.""" + mock_psycopg_connect.return_value = self.mock_conn + self.mock_cursor.fetchall.return_value = [ + (self.test_ids[0], [0.1, 0.2, 0.3], {"key": "value1"}), + (self.test_ids[1], [0.4, 0.5, 0.6], {"key": "value2"}), + ] + + pgvector = PGVector( + dbname="test_db", + collection_name="test_collection", + embedding_model_dims=3, + user="test_user", + password="test_pass", + host="localhost", + port=5432, + diskann=False, + hnsw=False + ) + + results = pgvector.list(filters=None, limit=2) + + # Verify list query was executed without WHERE clause + list_calls = [call for call in self.mock_cursor.execute.call_args_list + if "SELECT id, vector, payload" in str(call) and "WHERE" not in str(call)] + self.assertTrue(len(list_calls) > 0) + + # Verify results + self.assertEqual(len(results), 1) # Returns list of lists + self.assertEqual(len(results[0]), 2) + self.assertEqual(results[0][0].id, self.test_ids[0]) + self.assertEqual(results[0][1].id, self.test_ids[1]) + + @patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 2) + @patch('mem0.vector_stores.pgvector.psycopg2.connect') + def test_list_with_no_filters_psycopg2(self, mock_connect): + """Test list with no filters using psycopg2.""" + mock_connect.return_value = self.mock_conn + self.mock_cursor.fetchall.return_value = [ + (self.test_ids[0], [0.1, 0.2, 0.3], {"key": "value1"}), + (self.test_ids[1], [0.4, 0.5, 0.6], {"key": "value2"}), + ] + + pgvector = PGVector( + dbname="test_db", + collection_name="test_collection", + embedding_model_dims=3, + user="test_user", + password="test_pass", + host="localhost", + port=5432, + diskann=False, + hnsw=False + ) + + results = pgvector.list(filters=None, limit=2) + + # Verify list query was executed without WHERE clause + list_calls = [call for call in self.mock_cursor.execute.call_args_list + if "SELECT id, vector, payload" in str(call) and "WHERE" not in str(call)] + self.assertTrue(len(list_calls) > 0) + + # Verify results + self.assertEqual(len(results), 1) # Returns list of lists + self.assertEqual(len(results[0]), 2) + self.assertEqual(results[0][0].id, self.test_ids[0]) + self.assertEqual(results[0][1].id, self.test_ids[1]) + @patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 3) @patch('mem0.vector_stores.pgvector.psycopg.connect') @patch('mem0.vector_stores.pgvector.psycopg2.connect')