diff --git a/Makefile b/Makefile index 1afe9ec1c..a8173a39f 100644 --- a/Makefile +++ b/Makefile @@ -13,7 +13,7 @@ install: install_all: pip install ruff==0.6.9 groq together boto3 litellm ollama chromadb weaviate weaviate-client sentence_transformers vertexai \ google-generativeai elasticsearch opensearch-py vecs "pinecone<7.0.0" pinecone-text faiss-cpu langchain-community \ - upstash-vector azure-search-documents langchain-memgraph langchain-neo4j langchain-aws rank-bm25 pymochow pymongo + upstash-vector azure-search-documents langchain-memgraph langchain-neo4j langchain-aws rank-bm25 pymochow pymongo psycopg # Format code with ruff format: diff --git a/mem0/vector_stores/pgvector.py b/mem0/vector_stores/pgvector.py index 3b5d4157c..d88d60ec2 100644 --- a/mem0/vector_stores/pgvector.py +++ b/mem0/vector_stores/pgvector.py @@ -2,13 +2,28 @@ import json import logging from typing import List, Optional +from psycopg.types.json import Json from pydantic import BaseModel +# Try to import psycopg (psycopg3) first, then fall back to psycopg2 try: - import psycopg2 - from psycopg2.extras import execute_values + import psycopg + from psycopg import execute_values + PSYCOPG_VERSION = 3 + logger = logging.getLogger(__name__) + logger.info("Using psycopg (psycopg3) for PostgreSQL connections") except ImportError: - raise ImportError("The 'psycopg2' library is required. Please install it using 'pip install psycopg2'.") + try: + import psycopg2 + from psycopg2.extras import execute_values + PSYCOPG_VERSION = 2 + logger = logging.getLogger(__name__) + logger.info("Using psycopg2 for PostgreSQL connections") + except ImportError: + raise ImportError( + "Neither 'psycopg' nor 'psycopg2' library is available. " + "Please install one of them using 'pip install psycopg' or 'pip install psycopg2'." + ) from mem0.vector_stores.base import VectorStoreBase @@ -53,7 +68,15 @@ class PGVector(VectorStoreBase): self.use_hnsw = hnsw self.embedding_model_dims = embedding_model_dims - self.conn = psycopg2.connect(dbname=dbname, user=user, password=password, host=host, port=port) + if PSYCOPG_VERSION == 3: + self.conn = psycopg.connect( + dbname=dbname, user=user, password=password, host=host, port=port + ) + else: + self.conn = psycopg2.connect( + dbname=dbname, user=user, password=password, host=host, port=port + ) + self.cur = self.conn.cursor() collections = self.list_cols() @@ -184,10 +207,19 @@ class PGVector(VectorStoreBase): (vector, vector_id), ) if payload: - self.cur.execute( - f"UPDATE {self.collection_name} SET payload = %s WHERE id = %s", - (psycopg2.extras.Json(payload), vector_id), - ) + # Handle JSON serialization based on psycopg version + if PSYCOPG_VERSION == 3: + # psycopg3 uses psycopg.types.json.Json + self.cur.execute( + f"UPDATE {self.collection_name} SET payload = %s WHERE id = %s", + (Json(payload), vector_id), + ) + else: + # psycopg2 uses psycopg2.extras.Json + self.cur.execute( + f"UPDATE {self.collection_name} SET payload = %s WHERE id = %s", + (psycopg2.extras.Json(payload), vector_id), + ) self.conn.commit() def get(self, vector_id) -> OutputData: diff --git a/pyproject.toml b/pyproject.toml index 05e2bb24b..d92df0935 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -36,6 +36,7 @@ vector_stores = [ "faiss-cpu>=1.7.4", "upstash-vector>=0.1.0", "azure-search-documents>=11.4.0b8", + "psycopg>=3.2.8", "pymongo>=4.13.2", "pymochow>=2.2.9", ] diff --git a/tests/vector_stores/test_pgvector.py b/tests/vector_stores/test_pgvector.py new file mode 100644 index 000000000..ac18a3002 --- /dev/null +++ b/tests/vector_stores/test_pgvector.py @@ -0,0 +1,770 @@ +import unittest +import uuid +from unittest.mock import MagicMock, patch + +from mem0.vector_stores.pgvector import PGVector + + +class TestPGVector(unittest.TestCase): + def setUp(self): + """Set up test fixtures.""" + self.mock_conn = MagicMock() + self.mock_cursor = MagicMock() + self.mock_conn.cursor.return_value = self.mock_cursor + + # Mock connection pool + self.mock_pool = MagicMock() + self.mock_pool.getconn.return_value = self.mock_conn + + # Mock connection string + self.connection_string = "postgresql://user:pass@host:5432/db" + + # Test data + self.test_vectors = [[0.1, 0.2, 0.3], [0.4, 0.5, 0.6]] + self.test_payloads = [{"key": "value1"}, {"key": "value2"}] + self.test_ids = [str(uuid.uuid4()), str(uuid.uuid4())] + + @patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 3) + @patch('mem0.vector_stores.pgvector.psycopg.connect') + @patch('mem0.vector_stores.pgvector.psycopg2.connect') + def test_init_with_individual_params_psycopg3(self, mock_psycopg2_connect, mock_psycopg_connect): + """Test initialization with individual parameters using psycopg3.""" + # Mock psycopg3 to be available + mock_psycopg_connect.return_value = self.mock_conn + self.mock_cursor.fetchall.return_value = [] # No existing collections + + 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 + ) + + mock_psycopg_connect.assert_called_once_with( + dbname="test_db", + user="test_user", + password="test_pass", + host="localhost", + port=5432 + ) + self.assertEqual(pgvector.collection_name, "test_collection") + self.assertEqual(pgvector.embedding_model_dims, 3) + + @patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 2) + @patch('mem0.vector_stores.pgvector.psycopg2.connect') + def test_init_with_individual_params_psycopg2(self, mock_connect): + """Test initialization with individual parameters using psycopg2.""" + mock_connect.return_value = self.mock_conn + self.mock_cursor.fetchall.return_value = [] # No existing collections + + 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 + ) + + mock_connect.assert_called_once_with( + dbname="test_db", + user="test_user", + password="test_pass", + host="localhost", + port=5432 + ) + 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.psycopg.connect') + @patch('mem0.vector_stores.pgvector.psycopg2.connect') + def test_create_col_psycopg3(self, mock_psycopg2_connect, mock_psycopg_connect): + """Test collection creation with psycopg3.""" + mock_psycopg_connect.return_value = self.mock_conn + self.mock_cursor.fetchall.return_value = [] + + 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 + ) + + # 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 test_collection" in str(call)] + self.assertTrue(len(table_creation_calls) > 0) + self.mock_conn.commit.assert_called() + + # Verify pgvector instance properties + self.assertEqual(pgvector.collection_name, "test_collection") + self.assertEqual(pgvector.embedding_model_dims, 3) + + @patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 2) + @patch('mem0.vector_stores.pgvector.psycopg2.connect') + def test_create_col_psycopg2(self, mock_connect): + """Test collection creation with psycopg2.""" + mock_connect.return_value = self.mock_conn + self.mock_cursor.fetchall.return_value = [] + + 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 + ) + + # 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 test_collection" in str(call)] + self.assertTrue(len(table_creation_calls) > 0) + self.mock_conn.commit.assert_called() + + # Verify pgvector instance properties + 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.psycopg.connect') + @patch('mem0.vector_stores.pgvector.psycopg2.connect') + @patch('mem0.vector_stores.pgvector.execute_values') + def test_insert_psycopg3(self, mock_execute_values, mock_psycopg2_connect, mock_psycopg_connect): + """Test vector insertion with psycopg3.""" + mock_psycopg_connect.return_value = self.mock_conn + self.mock_cursor.fetchall.return_value = [] + + 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 + ) + + pgvector.insert(self.test_vectors, self.test_payloads, self.test_ids) + + # Verify execute_values was called + mock_execute_values.assert_called_once() + call_args = mock_execute_values.call_args + self.assertIn("INSERT INTO test_collection", call_args[0][1]) + + # Verify data format + data_arg = call_args[0][2] + self.assertEqual(len(data_arg), 2) + self.assertEqual(data_arg[0][0], self.test_ids[0]) + self.assertEqual(data_arg[1][0], self.test_ids[1]) + + @patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 2) + @patch('mem0.vector_stores.pgvector.psycopg2.connect') + @patch('mem0.vector_stores.pgvector.execute_values') + def test_insert_psycopg2(self, mock_execute_values, mock_connect): + """Test vector insertion with psycopg2.""" + mock_connect.return_value = self.mock_conn + self.mock_cursor.fetchall.return_value = [] + + 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 + ) + + pgvector.insert(self.test_vectors, self.test_payloads, self.test_ids) + + # Verify execute_values was called + mock_execute_values.assert_called_once() + call_args = mock_execute_values.call_args + self.assertIn("INSERT INTO test_collection", call_args[0][1]) + + # Verify data format + data_arg = call_args[0][2] + self.assertEqual(len(data_arg), 2) + self.assertEqual(data_arg[0][0], self.test_ids[0]) + self.assertEqual(data_arg[1][0], 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_psycopg3(self, mock_psycopg2_connect, mock_psycopg_connect): + """Test search with 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) + + # Verify search query was executed + search_calls = [call for call in self.mock_cursor.execute.call_args_list + if "SELECT id, vector <=" 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_psycopg2(self, mock_connect): + """Test search with 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) + + # Verify search query was executed + search_calls = [call for call in self.mock_cursor.execute.call_args_list + if "SELECT id, vector <=" 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_delete_psycopg3(self, mock_psycopg2_connect, mock_psycopg_connect): + """Test delete with psycopg3.""" + mock_psycopg_connect.return_value = self.mock_conn + + 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 + ) + + pgvector.delete(self.test_ids[0]) + + # Verify delete query was executed + delete_calls = [call for call in self.mock_cursor.execute.call_args_list + if "DELETE FROM test_collection" in str(call)] + self.assertTrue(len(delete_calls) > 0) + self.mock_conn.commit.assert_called() + + @patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 2) + @patch('mem0.vector_stores.pgvector.psycopg2.connect') + def test_delete_psycopg2(self, mock_connect): + """Test delete with psycopg2.""" + mock_connect.return_value = self.mock_conn + + 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 + ) + + pgvector.delete(self.test_ids[0]) + + # Verify delete query was executed + delete_calls = [call for call in self.mock_cursor.execute.call_args_list + if "DELETE FROM test_collection" in str(call)] + self.assertTrue(len(delete_calls) > 0) + self.mock_conn.commit.assert_called() + + @patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 3) + @patch('mem0.vector_stores.pgvector.psycopg.connect') + @patch('mem0.vector_stores.pgvector.psycopg2.connect') + def test_update_psycopg3(self, mock_psycopg2_connect, mock_psycopg_connect): + """Test update with psycopg3.""" + mock_psycopg_connect.return_value = self.mock_conn + + 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 + ) + + updated_vector = [0.5, 0.6, 0.7] + updated_payload = {"updated": "value"} + + pgvector.update(self.test_ids[0], vector=updated_vector, payload=updated_payload) + + # Verify update queries were executed + update_calls = [call for call in self.mock_cursor.execute.call_args_list + if "UPDATE test_collection" in str(call)] + self.assertTrue(len(update_calls) > 0) + self.mock_conn.commit.assert_called() + + @patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 2) + @patch('mem0.vector_stores.pgvector.psycopg2.connect') + def test_update_psycopg2(self, mock_connect): + """Test update with psycopg2.""" + mock_connect.return_value = self.mock_conn + + 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 + ) + + updated_vector = [0.5, 0.6, 0.7] + updated_payload = {"updated": "value"} + + pgvector.update(self.test_ids[0], vector=updated_vector, payload=updated_payload) + + # Verify update queries were executed + update_calls = [call for call in self.mock_cursor.execute.call_args_list + if "UPDATE test_collection" in str(call)] + self.assertTrue(len(update_calls) > 0) + self.mock_conn.commit.assert_called() + + @patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 3) + @patch('mem0.vector_stores.pgvector.psycopg.connect') + @patch('mem0.vector_stores.pgvector.psycopg2.connect') + def test_get_psycopg3(self, mock_psycopg2_connect, mock_psycopg_connect): + """Test get with psycopg3.""" + mock_psycopg_connect.return_value = self.mock_conn + self.mock_cursor.fetchone.return_value = (self.test_ids[0], [0.1, 0.2, 0.3], {"key": "value1"}) + + 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 + ) + + result = pgvector.get(self.test_ids[0]) + + # Verify get query was executed + get_calls = [call for call in self.mock_cursor.execute.call_args_list + if "SELECT id, vector, payload" in str(call)] + self.assertTrue(len(get_calls) > 0) + + # Verify result + self.assertIsNotNone(result) + self.assertEqual(result.id, self.test_ids[0]) + self.assertEqual(result.payload, {"key": "value1"}) + + @patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 2) + @patch('mem0.vector_stores.pgvector.psycopg2.connect') + def test_get_psycopg2(self, mock_connect): + """Test get with psycopg2.""" + mock_connect.return_value = self.mock_conn + self.mock_cursor.fetchone.return_value = (self.test_ids[0], [0.1, 0.2, 0.3], {"key": "value1"}) + + 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 + ) + + result = pgvector.get(self.test_ids[0]) + + # Verify get query was executed + get_calls = [call for call in self.mock_cursor.execute.call_args_list + if "SELECT id, vector, payload" in str(call)] + self.assertTrue(len(get_calls) > 0) + + # Verify result + self.assertIsNotNone(result) + self.assertEqual(result.id, self.test_ids[0]) + self.assertEqual(result.payload, {"key": "value1"}) + + @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_cols_psycopg3(self, mock_psycopg2_connect, mock_psycopg_connect): + """Test list_cols with psycopg3.""" + mock_psycopg_connect.return_value = self.mock_conn + self.mock_cursor.fetchall.return_value = [("test_collection",), ("other_table",)] + + 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 + ) + + collections = pgvector.list_cols() + + # Verify list_cols query was executed + list_calls = [call for call in self.mock_cursor.execute.call_args_list + if "SELECT table_name FROM information_schema.tables" in str(call)] + self.assertTrue(len(list_calls) > 0) + + # Verify result + self.assertEqual(collections, ["test_collection", "other_table"]) + + @patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 2) + @patch('mem0.vector_stores.pgvector.psycopg2.connect') + def test_list_cols_psycopg2(self, mock_connect): + """Test list_cols with psycopg2.""" + mock_connect.return_value = self.mock_conn + self.mock_cursor.fetchall.return_value = [("test_collection",), ("other_table",)] + + 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 + ) + + collections = pgvector.list_cols() + + # Verify list_cols query was executed + list_calls = [call for call in self.mock_cursor.execute.call_args_list + if "SELECT table_name FROM information_schema.tables" in str(call)] + self.assertTrue(len(list_calls) > 0) + + # Verify result + self.assertEqual(collections, ["test_collection", "other_table"]) + + @patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 3) + @patch('mem0.vector_stores.pgvector.psycopg.connect') + @patch('mem0.vector_stores.pgvector.psycopg2.connect') + def test_delete_col_psycopg3(self, mock_psycopg2_connect, mock_psycopg_connect): + """Test delete_col with psycopg3.""" + mock_psycopg_connect.return_value = self.mock_conn + + 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 + ) + + pgvector.delete_col() + + # Verify delete_col query was executed + delete_calls = [call for call in self.mock_cursor.execute.call_args_list + if "DROP TABLE IF EXISTS test_collection" in str(call)] + self.assertTrue(len(delete_calls) > 0) + self.mock_conn.commit.assert_called() + + @patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 2) + @patch('mem0.vector_stores.pgvector.psycopg2.connect') + def test_delete_col_psycopg2(self, mock_connect): + """Test delete_col with psycopg2.""" + mock_connect.return_value = self.mock_conn + + 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 + ) + + pgvector.delete_col() + + # Verify delete_col query was executed + delete_calls = [call for call in self.mock_cursor.execute.call_args_list + if "DROP TABLE IF EXISTS test_collection" in str(call)] + self.assertTrue(len(delete_calls) > 0) + self.mock_conn.commit.assert_called() + + @patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 3) + @patch('mem0.vector_stores.pgvector.psycopg.connect') + @patch('mem0.vector_stores.pgvector.psycopg2.connect') + def test_col_info_psycopg3(self, mock_psycopg2_connect, mock_psycopg_connect): + """Test col_info with psycopg3.""" + mock_psycopg_connect.return_value = self.mock_conn + self.mock_cursor.fetchone.return_value = ("test_collection", 100, "1 MB") + + 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 + ) + + info = pgvector.col_info() + + # Verify col_info query was executed + info_calls = [call for call in self.mock_cursor.execute.call_args_list + if "SELECT table_name" in str(call)] + self.assertTrue(len(info_calls) > 0) + + # Verify result + self.assertEqual(info["name"], "test_collection") + self.assertEqual(info["count"], 100) + self.assertEqual(info["size"], "1 MB") + + @patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 2) + @patch('mem0.vector_stores.pgvector.psycopg2.connect') + def test_col_info_psycopg2(self, mock_connect): + """Test col_info with psycopg2.""" + mock_connect.return_value = self.mock_conn + self.mock_cursor.fetchone.return_value = ("test_collection", 100, "1 MB") + + 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 + ) + + info = pgvector.col_info() + + # Verify col_info query was executed + info_calls = [call for call in self.mock_cursor.execute.call_args_list + if "SELECT table_name" in str(call)] + self.assertTrue(len(info_calls) > 0) + + # Verify result + self.assertEqual(info["name"], "test_collection") + self.assertEqual(info["count"], 100) + self.assertEqual(info["size"], "1 MB") + + @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_psycopg3(self, mock_psycopg2_connect, mock_psycopg_connect): + """Test list with 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(limit=2) + + # Verify list query was executed + list_calls = [call for call in self.mock_cursor.execute.call_args_list + if "SELECT id, vector, payload" in str(call)] + self.assertTrue(len(list_calls) > 0) + + # Verify result + 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_psycopg2(self, mock_connect): + """Test list with 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(limit=2) + + # Verify list query was executed + list_calls = [call for call in self.mock_cursor.execute.call_args_list + if "SELECT id, vector, payload" in str(call)] + self.assertTrue(len(list_calls) > 0) + + # Verify result + 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') + def test_reset_psycopg3(self, mock_psycopg2_connect, mock_psycopg_connect): + """Test reset with psycopg3.""" + mock_psycopg_connect.return_value = self.mock_conn + self.mock_cursor.fetchall.return_value = [] + + 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 + ) + + pgvector.reset() + + # Verify reset operations were executed + drop_calls = [call for call in self.mock_cursor.execute.call_args_list + if "DROP TABLE IF EXISTS" in str(call)] + create_calls = [call for call in self.mock_cursor.execute.call_args_list + if "CREATE TABLE IF NOT EXISTS" in str(call)] + self.assertTrue(len(drop_calls) > 0) + self.assertTrue(len(create_calls) > 0) + + @patch('mem0.vector_stores.pgvector.PSYCOPG_VERSION', 2) + @patch('mem0.vector_stores.pgvector.psycopg2.connect') + def test_reset_psycopg2(self, mock_connect): + """Test reset with psycopg2.""" + mock_connect.return_value = self.mock_conn + self.mock_cursor.fetchall.return_value = [] + + 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 + ) + + pgvector.reset() + + # Verify reset operations were executed + drop_calls = [call for call in self.mock_cursor.execute.call_args_list + if "DROP TABLE IF EXISTS" in str(call)] + create_calls = [call for call in self.mock_cursor.execute.call_args_list + if "CREATE TABLE IF NOT EXISTS" in str(call)] + self.assertTrue(len(drop_calls) > 0) + self.assertTrue(len(create_calls) > 0) + + def tearDown(self): + """Clean up after each test.""" + pass