From 6a178a4ed04ec6f51eb3ff18a7b66bc53673d08a Mon Sep 17 00:00:00 2001 From: Cyber MacGeddon Date: Mon, 14 Jul 2025 22:36:49 +0100 Subject: [PATCH] Fixing storage and adding tests --- .../test_graph_embeddings_pinecone_query.py | 507 ++++++++++++++++++ .../test_graph_embeddings_pinecone_storage.py | 460 ++++++++++++++++ .../graph_embeddings/pinecone/service.py | 77 +-- .../graph_embeddings/pinecone/write.py | 45 +- 4 files changed, 1004 insertions(+), 85 deletions(-) create mode 100644 tests/unit/test_query/test_graph_embeddings_pinecone_query.py create mode 100644 tests/unit/test_storage/test_graph_embeddings_pinecone_storage.py diff --git a/tests/unit/test_query/test_graph_embeddings_pinecone_query.py b/tests/unit/test_query/test_graph_embeddings_pinecone_query.py new file mode 100644 index 00000000..5352e002 --- /dev/null +++ b/tests/unit/test_query/test_graph_embeddings_pinecone_query.py @@ -0,0 +1,507 @@ +""" +Tests for Pinecone graph embeddings query service +""" + +import pytest +from unittest.mock import MagicMock, patch + +from trustgraph.query.graph_embeddings.pinecone.service import Processor +from trustgraph.schema import Value + + +class TestPineconeGraphEmbeddingsQueryProcessor: + """Test cases for Pinecone graph embeddings query processor""" + + @pytest.fixture + def mock_query_message(self): + """Create a mock query message for testing""" + message = MagicMock() + message.vectors = [ + [0.1, 0.2, 0.3], + [0.4, 0.5, 0.6] + ] + message.limit = 5 + message.user = 'test_user' + message.collection = 'test_collection' + return message + + @pytest.fixture + def processor(self): + """Create a processor instance for testing""" + with patch('trustgraph.query.graph_embeddings.pinecone.service.Pinecone') as mock_pinecone_class: + mock_pinecone = MagicMock() + mock_pinecone_class.return_value = mock_pinecone + + processor = Processor( + taskgroup=MagicMock(), + id='test-pinecone-ge-query', + api_key='test-api-key' + ) + + return processor + + @patch('trustgraph.query.graph_embeddings.pinecone.service.Pinecone') + @patch('trustgraph.query.graph_embeddings.pinecone.service.default_api_key', 'env-api-key') + def test_processor_initialization_with_defaults(self, mock_pinecone_class): + """Test processor initialization with default parameters""" + mock_pinecone = MagicMock() + mock_pinecone_class.return_value = mock_pinecone + taskgroup_mock = MagicMock() + + processor = Processor(taskgroup=taskgroup_mock) + + mock_pinecone_class.assert_called_once_with(api_key='env-api-key') + assert processor.pinecone == mock_pinecone + assert processor.api_key == 'env-api-key' + + @patch('trustgraph.query.graph_embeddings.pinecone.service.Pinecone') + def test_processor_initialization_with_custom_params(self, mock_pinecone_class): + """Test processor initialization with custom parameters""" + mock_pinecone = MagicMock() + mock_pinecone_class.return_value = mock_pinecone + taskgroup_mock = MagicMock() + + processor = Processor( + taskgroup=taskgroup_mock, + api_key='custom-api-key' + ) + + mock_pinecone_class.assert_called_once_with(api_key='custom-api-key') + assert processor.api_key == 'custom-api-key' + + @patch('trustgraph.query.graph_embeddings.pinecone.service.PineconeGRPC') + def test_processor_initialization_with_url(self, mock_pinecone_grpc_class): + """Test processor initialization with custom URL (GRPC mode)""" + mock_pinecone = MagicMock() + mock_pinecone_grpc_class.return_value = mock_pinecone + taskgroup_mock = MagicMock() + + processor = Processor( + taskgroup=taskgroup_mock, + api_key='test-api-key', + url='https://custom-host.pinecone.io' + ) + + mock_pinecone_grpc_class.assert_called_once_with( + api_key='test-api-key', + host='https://custom-host.pinecone.io' + ) + assert processor.pinecone == mock_pinecone + assert processor.url == 'https://custom-host.pinecone.io' + + @patch('trustgraph.query.graph_embeddings.pinecone.service.default_api_key', 'not-specified') + def test_processor_initialization_missing_api_key(self): + """Test processor initialization fails with missing API key""" + taskgroup_mock = MagicMock() + + with pytest.raises(RuntimeError, match="Pinecone API key must be specified"): + Processor(taskgroup=taskgroup_mock) + + def test_create_value_uri(self, processor): + """Test create_value method for URI entities""" + uri_entity = "http://example.org/entity" + value = processor.create_value(uri_entity) + + assert isinstance(value, Value) + assert value.value == uri_entity + assert value.is_uri == True + + def test_create_value_https_uri(self, processor): + """Test create_value method for HTTPS URI entities""" + uri_entity = "https://example.org/entity" + value = processor.create_value(uri_entity) + + assert isinstance(value, Value) + assert value.value == uri_entity + assert value.is_uri == True + + def test_create_value_literal(self, processor): + """Test create_value method for literal entities""" + literal_entity = "literal_entity" + value = processor.create_value(literal_entity) + + assert isinstance(value, Value) + assert value.value == literal_entity + assert value.is_uri == False + + @pytest.mark.asyncio + async def test_query_graph_embeddings_single_vector(self, processor): + """Test querying graph embeddings with a single vector""" + message = MagicMock() + message.vectors = [[0.1, 0.2, 0.3]] + message.limit = 3 + message.user = 'test_user' + message.collection = 'test_collection' + + # Mock index and query results + mock_index = MagicMock() + processor.pinecone.Index.return_value = mock_index + + mock_results = MagicMock() + mock_results.matches = [ + MagicMock(metadata={'entity': 'http://example.org/entity1'}), + MagicMock(metadata={'entity': 'entity2'}), + MagicMock(metadata={'entity': 'http://example.org/entity3'}) + ] + mock_index.query.return_value = mock_results + + entities = await processor.query_graph_embeddings(message) + + # Verify index was accessed correctly + expected_index_name = "t-test_user-test_collection-3" + processor.pinecone.Index.assert_called_once_with(expected_index_name) + + # Verify query parameters + mock_index.query.assert_called_once_with( + vector=[0.1, 0.2, 0.3], + top_k=6, # 2 * limit + include_values=False, + include_metadata=True + ) + + # Verify results + assert len(entities) == 3 + assert entities[0].value == 'http://example.org/entity1' + assert entities[0].is_uri == True + assert entities[1].value == 'entity2' + assert entities[1].is_uri == False + assert entities[2].value == 'http://example.org/entity3' + assert entities[2].is_uri == True + + @pytest.mark.asyncio + async def test_query_graph_embeddings_multiple_vectors(self, processor, mock_query_message): + """Test querying graph embeddings with multiple vectors""" + # Mock index and query results + mock_index = MagicMock() + processor.pinecone.Index.return_value = mock_index + + # First query results + mock_results1 = MagicMock() + mock_results1.matches = [ + MagicMock(metadata={'entity': 'entity1'}), + MagicMock(metadata={'entity': 'entity2'}) + ] + + # Second query results + mock_results2 = MagicMock() + mock_results2.matches = [ + MagicMock(metadata={'entity': 'entity2'}), # Duplicate + MagicMock(metadata={'entity': 'entity3'}) + ] + + mock_index.query.side_effect = [mock_results1, mock_results2] + + entities = await processor.query_graph_embeddings(mock_query_message) + + # Verify both queries were made + assert mock_index.query.call_count == 2 + + # Verify deduplication occurred + entity_values = [e.value for e in entities] + assert len(entity_values) == 3 + assert 'entity1' in entity_values + assert 'entity2' in entity_values + assert 'entity3' in entity_values + + @pytest.mark.asyncio + async def test_query_graph_embeddings_limit_handling(self, processor): + """Test that query respects the limit parameter""" + message = MagicMock() + message.vectors = [[0.1, 0.2, 0.3]] + message.limit = 2 + message.user = 'test_user' + message.collection = 'test_collection' + + # Mock index with many results + mock_index = MagicMock() + processor.pinecone.Index.return_value = mock_index + + mock_results = MagicMock() + mock_results.matches = [ + MagicMock(metadata={'entity': f'entity{i}'}) for i in range(10) + ] + mock_index.query.return_value = mock_results + + entities = await processor.query_graph_embeddings(message) + + # Verify limit is respected + assert len(entities) == 2 + + @pytest.mark.asyncio + async def test_query_graph_embeddings_zero_limit(self, processor): + """Test querying with zero limit returns empty results""" + message = MagicMock() + message.vectors = [[0.1, 0.2, 0.3]] + message.limit = 0 + message.user = 'test_user' + message.collection = 'test_collection' + + mock_index = MagicMock() + processor.pinecone.Index.return_value = mock_index + + entities = await processor.query_graph_embeddings(message) + + # Verify no query was made and empty result returned + mock_index.query.assert_not_called() + assert entities == [] + + @pytest.mark.asyncio + async def test_query_graph_embeddings_negative_limit(self, processor): + """Test querying with negative limit returns empty results""" + message = MagicMock() + message.vectors = [[0.1, 0.2, 0.3]] + message.limit = -1 + message.user = 'test_user' + message.collection = 'test_collection' + + mock_index = MagicMock() + processor.pinecone.Index.return_value = mock_index + + entities = await processor.query_graph_embeddings(message) + + # Verify no query was made and empty result returned + mock_index.query.assert_not_called() + assert entities == [] + + @pytest.mark.asyncio + async def test_query_graph_embeddings_different_vector_dimensions(self, processor): + """Test querying with vectors of different dimensions""" + message = MagicMock() + message.vectors = [ + [0.1, 0.2], # 2D vector + [0.3, 0.4, 0.5, 0.6] # 4D vector + ] + message.limit = 5 + message.user = 'test_user' + message.collection = 'test_collection' + + mock_index_2d = MagicMock() + mock_index_4d = MagicMock() + + def mock_index_side_effect(name): + if name.endswith("-2"): + return mock_index_2d + elif name.endswith("-4"): + return mock_index_4d + + processor.pinecone.Index.side_effect = mock_index_side_effect + + # Mock results for different dimensions + mock_results_2d = MagicMock() + mock_results_2d.matches = [MagicMock(metadata={'entity': 'entity_2d'})] + mock_index_2d.query.return_value = mock_results_2d + + mock_results_4d = MagicMock() + mock_results_4d.matches = [MagicMock(metadata={'entity': 'entity_4d'})] + mock_index_4d.query.return_value = mock_results_4d + + entities = await processor.query_graph_embeddings(message) + + # Verify different indexes were used + assert processor.pinecone.Index.call_count == 2 + mock_index_2d.query.assert_called_once() + mock_index_4d.query.assert_called_once() + + # Verify results from both dimensions + entity_values = [e.value for e in entities] + assert 'entity_2d' in entity_values + assert 'entity_4d' in entity_values + + @pytest.mark.asyncio + async def test_query_graph_embeddings_empty_vectors_list(self, processor): + """Test querying with empty vectors list""" + message = MagicMock() + message.vectors = [] + message.limit = 5 + message.user = 'test_user' + message.collection = 'test_collection' + + mock_index = MagicMock() + processor.pinecone.Index.return_value = mock_index + + entities = await processor.query_graph_embeddings(message) + + # Verify no queries were made and empty result returned + processor.pinecone.Index.assert_not_called() + mock_index.query.assert_not_called() + assert entities == [] + + @pytest.mark.asyncio + async def test_query_graph_embeddings_no_results(self, processor): + """Test querying when index returns no results""" + message = MagicMock() + message.vectors = [[0.1, 0.2, 0.3]] + message.limit = 5 + message.user = 'test_user' + message.collection = 'test_collection' + + mock_index = MagicMock() + processor.pinecone.Index.return_value = mock_index + + mock_results = MagicMock() + mock_results.matches = [] + mock_index.query.return_value = mock_results + + entities = await processor.query_graph_embeddings(message) + + # Verify empty results + assert entities == [] + + @pytest.mark.asyncio + async def test_query_graph_embeddings_deduplication_across_vectors(self, processor): + """Test that deduplication works correctly across multiple vector queries""" + message = MagicMock() + message.vectors = [ + [0.1, 0.2, 0.3], + [0.4, 0.5, 0.6] + ] + message.limit = 3 + message.user = 'test_user' + message.collection = 'test_collection' + + mock_index = MagicMock() + processor.pinecone.Index.return_value = mock_index + + # Both queries return overlapping results + mock_results1 = MagicMock() + mock_results1.matches = [ + MagicMock(metadata={'entity': 'entity1'}), + MagicMock(metadata={'entity': 'entity2'}), + MagicMock(metadata={'entity': 'entity3'}), + MagicMock(metadata={'entity': 'entity4'}) + ] + + mock_results2 = MagicMock() + mock_results2.matches = [ + MagicMock(metadata={'entity': 'entity2'}), # Duplicate + MagicMock(metadata={'entity': 'entity3'}), # Duplicate + MagicMock(metadata={'entity': 'entity5'}) + ] + + mock_index.query.side_effect = [mock_results1, mock_results2] + + entities = await processor.query_graph_embeddings(message) + + # Should get exactly 3 unique entities (respecting limit) + assert len(entities) == 3 + entity_values = [e.value for e in entities] + assert len(set(entity_values)) == 3 # All unique + + @pytest.mark.asyncio + async def test_query_graph_embeddings_early_termination_on_limit(self, processor): + """Test that querying stops early when limit is reached""" + message = MagicMock() + message.vectors = [ + [0.1, 0.2, 0.3], + [0.4, 0.5, 0.6], + [0.7, 0.8, 0.9] + ] + message.limit = 2 + message.user = 'test_user' + message.collection = 'test_collection' + + mock_index = MagicMock() + processor.pinecone.Index.return_value = mock_index + + # First query returns enough results to meet limit + mock_results1 = MagicMock() + mock_results1.matches = [ + MagicMock(metadata={'entity': 'entity1'}), + MagicMock(metadata={'entity': 'entity2'}), + MagicMock(metadata={'entity': 'entity3'}) + ] + mock_index.query.return_value = mock_results1 + + entities = await processor.query_graph_embeddings(message) + + # Should only make one query since limit was reached + mock_index.query.assert_called_once() + assert len(entities) == 2 + + @pytest.mark.asyncio + async def test_query_graph_embeddings_exception_handling(self, processor): + """Test that exceptions are properly raised""" + message = MagicMock() + message.vectors = [[0.1, 0.2, 0.3]] + message.limit = 5 + message.user = 'test_user' + message.collection = 'test_collection' + + mock_index = MagicMock() + processor.pinecone.Index.return_value = mock_index + mock_index.query.side_effect = Exception("Query failed") + + with pytest.raises(Exception, match="Query failed"): + await processor.query_graph_embeddings(message) + + def test_add_args_method(self): + """Test that add_args properly configures argument parser""" + from argparse import ArgumentParser + from unittest.mock import patch + + parser = ArgumentParser() + + # Mock the parent class add_args method + with patch('trustgraph.query.graph_embeddings.pinecone.service.GraphEmbeddingsQueryService.add_args') as mock_parent_add_args: + Processor.add_args(parser) + + # Verify parent add_args was called + mock_parent_add_args.assert_called_once() + + # Verify our specific arguments were added + args = parser.parse_args([]) + + assert hasattr(args, 'api_key') + assert args.api_key == 'not-specified' # Default value when no env var + assert hasattr(args, 'url') + assert args.url is None + + def test_add_args_with_custom_values(self): + """Test add_args with custom command line values""" + from argparse import ArgumentParser + from unittest.mock import patch + + parser = ArgumentParser() + + with patch('trustgraph.query.graph_embeddings.pinecone.service.GraphEmbeddingsQueryService.add_args'): + Processor.add_args(parser) + + # Test parsing with custom values + args = parser.parse_args([ + '--api-key', 'custom-api-key', + '--url', 'https://custom-host.pinecone.io' + ]) + + assert args.api_key == 'custom-api-key' + assert args.url == 'https://custom-host.pinecone.io' + + def test_add_args_short_form(self): + """Test add_args with short form arguments""" + from argparse import ArgumentParser + from unittest.mock import patch + + parser = ArgumentParser() + + with patch('trustgraph.query.graph_embeddings.pinecone.service.GraphEmbeddingsQueryService.add_args'): + Processor.add_args(parser) + + # Test parsing with short form + args = parser.parse_args([ + '-a', 'short-api-key', + '-u', 'https://short-host.pinecone.io' + ]) + + assert args.api_key == 'short-api-key' + assert args.url == 'https://short-host.pinecone.io' + + @patch('trustgraph.query.graph_embeddings.pinecone.service.Processor.launch') + def test_run_function(self, mock_launch): + """Test the run function calls Processor.launch with correct parameters""" + from trustgraph.query.graph_embeddings.pinecone.service import run, default_ident + + run() + + mock_launch.assert_called_once_with( + default_ident, + "\nGraph embeddings query service. Input is vector, output is list of\nentities. Pinecone implementation.\n" + ) \ No newline at end of file diff --git a/tests/unit/test_storage/test_graph_embeddings_pinecone_storage.py b/tests/unit/test_storage/test_graph_embeddings_pinecone_storage.py new file mode 100644 index 00000000..91e60057 --- /dev/null +++ b/tests/unit/test_storage/test_graph_embeddings_pinecone_storage.py @@ -0,0 +1,460 @@ +""" +Tests for Pinecone graph embeddings storage service +""" + +import pytest +from unittest.mock import MagicMock, patch +import uuid + +from trustgraph.storage.graph_embeddings.pinecone.write import Processor +from trustgraph.schema import EntityEmbeddings, Value + + +class TestPineconeGraphEmbeddingsStorageProcessor: + """Test cases for Pinecone graph embeddings storage processor""" + + @pytest.fixture + def mock_message(self): + """Create a mock message for testing""" + message = MagicMock() + message.metadata = MagicMock() + message.metadata.user = 'test_user' + message.metadata.collection = 'test_collection' + + # Create test entity embeddings + entity1 = EntityEmbeddings( + entity=Value(value="http://example.org/entity1", is_uri=True), + vectors=[[0.1, 0.2, 0.3], [0.4, 0.5, 0.6]] + ) + entity2 = EntityEmbeddings( + entity=Value(value="entity2", is_uri=False), + vectors=[[0.7, 0.8, 0.9]] + ) + message.entities = [entity1, entity2] + + return message + + @pytest.fixture + def processor(self): + """Create a processor instance for testing""" + with patch('trustgraph.storage.graph_embeddings.pinecone.write.Pinecone') as mock_pinecone_class: + mock_pinecone = MagicMock() + mock_pinecone_class.return_value = mock_pinecone + + processor = Processor( + taskgroup=MagicMock(), + id='test-pinecone-ge-storage', + api_key='test-api-key' + ) + + return processor + + @patch('trustgraph.storage.graph_embeddings.pinecone.write.Pinecone') + @patch('trustgraph.storage.graph_embeddings.pinecone.write.default_api_key', 'env-api-key') + def test_processor_initialization_with_defaults(self, mock_pinecone_class): + """Test processor initialization with default parameters""" + mock_pinecone = MagicMock() + mock_pinecone_class.return_value = mock_pinecone + taskgroup_mock = MagicMock() + + processor = Processor(taskgroup=taskgroup_mock) + + mock_pinecone_class.assert_called_once_with(api_key='env-api-key') + assert processor.pinecone == mock_pinecone + assert processor.api_key == 'env-api-key' + assert processor.cloud == 'aws' + assert processor.region == 'us-east-1' + + @patch('trustgraph.storage.graph_embeddings.pinecone.write.Pinecone') + def test_processor_initialization_with_custom_params(self, mock_pinecone_class): + """Test processor initialization with custom parameters""" + mock_pinecone = MagicMock() + mock_pinecone_class.return_value = mock_pinecone + taskgroup_mock = MagicMock() + + processor = Processor( + taskgroup=taskgroup_mock, + api_key='custom-api-key', + cloud='gcp', + region='us-west1' + ) + + mock_pinecone_class.assert_called_once_with(api_key='custom-api-key') + assert processor.api_key == 'custom-api-key' + assert processor.cloud == 'gcp' + assert processor.region == 'us-west1' + + @patch('trustgraph.storage.graph_embeddings.pinecone.write.PineconeGRPC') + def test_processor_initialization_with_url(self, mock_pinecone_grpc_class): + """Test processor initialization with custom URL (GRPC mode)""" + mock_pinecone = MagicMock() + mock_pinecone_grpc_class.return_value = mock_pinecone + taskgroup_mock = MagicMock() + + processor = Processor( + taskgroup=taskgroup_mock, + api_key='test-api-key', + url='https://custom-host.pinecone.io' + ) + + mock_pinecone_grpc_class.assert_called_once_with( + api_key='test-api-key', + host='https://custom-host.pinecone.io' + ) + assert processor.pinecone == mock_pinecone + assert processor.url == 'https://custom-host.pinecone.io' + + @patch('trustgraph.storage.graph_embeddings.pinecone.write.default_api_key', 'not-specified') + def test_processor_initialization_missing_api_key(self): + """Test processor initialization fails with missing API key""" + taskgroup_mock = MagicMock() + + with pytest.raises(RuntimeError, match="Pinecone API key must be specified"): + Processor(taskgroup=taskgroup_mock) + + @pytest.mark.asyncio + async def test_store_graph_embeddings_single_entity(self, processor): + """Test storing graph embeddings for a single entity""" + message = MagicMock() + message.metadata = MagicMock() + message.metadata.user = 'test_user' + message.metadata.collection = 'test_collection' + + entity = EntityEmbeddings( + entity=Value(value="http://example.org/entity1", is_uri=True), + vectors=[[0.1, 0.2, 0.3], [0.4, 0.5, 0.6]] + ) + message.entities = [entity] + + # Mock index operations + mock_index = MagicMock() + processor.pinecone.Index.return_value = mock_index + processor.pinecone.has_index.return_value = True + + with patch('uuid.uuid4', side_effect=['id1', 'id2']): + await processor.store_graph_embeddings(message) + + # Verify index name and operations + expected_index_name = "t-test_user-test_collection-3" + processor.pinecone.Index.assert_called_with(expected_index_name) + + # Verify upsert was called for each vector + assert mock_index.upsert.call_count == 2 + + # Check first vector upsert + first_call = mock_index.upsert.call_args_list[0] + first_vectors = first_call[1]['vectors'] + assert len(first_vectors) == 1 + assert first_vectors[0]['id'] == 'id1' + assert first_vectors[0]['values'] == [0.1, 0.2, 0.3] + assert first_vectors[0]['metadata']['entity'] == "http://example.org/entity1" + + # Check second vector upsert + second_call = mock_index.upsert.call_args_list[1] + second_vectors = second_call[1]['vectors'] + assert len(second_vectors) == 1 + assert second_vectors[0]['id'] == 'id2' + assert second_vectors[0]['values'] == [0.4, 0.5, 0.6] + assert second_vectors[0]['metadata']['entity'] == "http://example.org/entity1" + + @pytest.mark.asyncio + async def test_store_graph_embeddings_multiple_entities(self, processor, mock_message): + """Test storing graph embeddings for multiple entities""" + # Mock index operations + mock_index = MagicMock() + processor.pinecone.Index.return_value = mock_index + processor.pinecone.has_index.return_value = True + + with patch('uuid.uuid4', side_effect=['id1', 'id2', 'id3']): + await processor.store_graph_embeddings(mock_message) + + # Verify upsert was called for each vector (3 total) + assert mock_index.upsert.call_count == 3 + + # Verify entity values in metadata + calls = mock_index.upsert.call_args_list + assert calls[0][1]['vectors'][0]['metadata']['entity'] == "http://example.org/entity1" + assert calls[1][1]['vectors'][0]['metadata']['entity'] == "http://example.org/entity1" + assert calls[2][1]['vectors'][0]['metadata']['entity'] == "entity2" + + @pytest.mark.asyncio + async def test_store_graph_embeddings_index_creation(self, processor): + """Test automatic index creation when index doesn't exist""" + message = MagicMock() + message.metadata = MagicMock() + message.metadata.user = 'test_user' + message.metadata.collection = 'test_collection' + + entity = EntityEmbeddings( + entity=Value(value="test_entity", is_uri=False), + vectors=[[0.1, 0.2, 0.3]] + ) + message.entities = [entity] + + # Mock index doesn't exist initially + processor.pinecone.has_index.return_value = False + mock_index = MagicMock() + processor.pinecone.Index.return_value = mock_index + + # Mock index creation + processor.pinecone.describe_index.return_value.status = {"ready": True} + + with patch('uuid.uuid4', return_value='test-id'): + await processor.store_graph_embeddings(message) + + # Verify index creation was called + expected_index_name = "t-test_user-test_collection-3" + processor.pinecone.create_index.assert_called_once() + create_call = processor.pinecone.create_index.call_args + assert create_call[1]['name'] == expected_index_name + assert create_call[1]['dimension'] == 3 + assert create_call[1]['metric'] == "cosine" + + @pytest.mark.asyncio + async def test_store_graph_embeddings_empty_entity_value(self, processor): + """Test storing graph embeddings with empty entity value (should be skipped)""" + message = MagicMock() + message.metadata = MagicMock() + message.metadata.user = 'test_user' + message.metadata.collection = 'test_collection' + + entity = EntityEmbeddings( + entity=Value(value="", is_uri=False), + vectors=[[0.1, 0.2, 0.3]] + ) + message.entities = [entity] + + mock_index = MagicMock() + processor.pinecone.Index.return_value = mock_index + + await processor.store_graph_embeddings(message) + + # Verify no upsert was called for empty entity + mock_index.upsert.assert_not_called() + + @pytest.mark.asyncio + async def test_store_graph_embeddings_none_entity_value(self, processor): + """Test storing graph embeddings with None entity value (should be skipped)""" + message = MagicMock() + message.metadata = MagicMock() + message.metadata.user = 'test_user' + message.metadata.collection = 'test_collection' + + entity = EntityEmbeddings( + entity=Value(value=None, is_uri=False), + vectors=[[0.1, 0.2, 0.3]] + ) + message.entities = [entity] + + mock_index = MagicMock() + processor.pinecone.Index.return_value = mock_index + + await processor.store_graph_embeddings(message) + + # Verify no upsert was called for None entity + mock_index.upsert.assert_not_called() + + @pytest.mark.asyncio + async def test_store_graph_embeddings_different_vector_dimensions(self, processor): + """Test storing graph embeddings with different vector dimensions""" + message = MagicMock() + message.metadata = MagicMock() + message.metadata.user = 'test_user' + message.metadata.collection = 'test_collection' + + entity = EntityEmbeddings( + entity=Value(value="test_entity", is_uri=False), + vectors=[ + [0.1, 0.2], # 2D vector + [0.3, 0.4, 0.5, 0.6], # 4D vector + [0.7, 0.8, 0.9] # 3D vector + ] + ) + message.entities = [entity] + + mock_index_2d = MagicMock() + mock_index_4d = MagicMock() + mock_index_3d = MagicMock() + + def mock_index_side_effect(name): + if name.endswith("-2"): + return mock_index_2d + elif name.endswith("-4"): + return mock_index_4d + elif name.endswith("-3"): + return mock_index_3d + + processor.pinecone.Index.side_effect = mock_index_side_effect + processor.pinecone.has_index.return_value = True + + with patch('uuid.uuid4', side_effect=['id1', 'id2', 'id3']): + await processor.store_graph_embeddings(message) + + # Verify different indexes were used for different dimensions + assert processor.pinecone.Index.call_count == 3 + mock_index_2d.upsert.assert_called_once() + mock_index_4d.upsert.assert_called_once() + mock_index_3d.upsert.assert_called_once() + + @pytest.mark.asyncio + async def test_store_graph_embeddings_empty_entities_list(self, processor): + """Test storing graph embeddings with empty entities list""" + message = MagicMock() + message.metadata = MagicMock() + message.metadata.user = 'test_user' + message.metadata.collection = 'test_collection' + message.entities = [] + + mock_index = MagicMock() + processor.pinecone.Index.return_value = mock_index + + await processor.store_graph_embeddings(message) + + # Verify no operations were performed + processor.pinecone.Index.assert_not_called() + mock_index.upsert.assert_not_called() + + @pytest.mark.asyncio + async def test_store_graph_embeddings_entity_with_no_vectors(self, processor): + """Test storing graph embeddings for entity with no vectors""" + message = MagicMock() + message.metadata = MagicMock() + message.metadata.user = 'test_user' + message.metadata.collection = 'test_collection' + + entity = EntityEmbeddings( + entity=Value(value="test_entity", is_uri=False), + vectors=[] + ) + message.entities = [entity] + + mock_index = MagicMock() + processor.pinecone.Index.return_value = mock_index + + await processor.store_graph_embeddings(message) + + # Verify no upsert was called (no vectors to insert) + mock_index.upsert.assert_not_called() + + @pytest.mark.asyncio + async def test_store_graph_embeddings_index_creation_failure(self, processor): + """Test handling of index creation failure""" + message = MagicMock() + message.metadata = MagicMock() + message.metadata.user = 'test_user' + message.metadata.collection = 'test_collection' + + entity = EntityEmbeddings( + entity=Value(value="test_entity", is_uri=False), + vectors=[[0.1, 0.2, 0.3]] + ) + message.entities = [entity] + + # Mock index doesn't exist and creation fails + processor.pinecone.has_index.return_value = False + processor.pinecone.create_index.side_effect = Exception("Index creation failed") + + with pytest.raises(Exception, match="Index creation failed"): + await processor.store_graph_embeddings(message) + + @pytest.mark.asyncio + async def test_store_graph_embeddings_index_creation_timeout(self, processor): + """Test handling of index creation timeout""" + message = MagicMock() + message.metadata = MagicMock() + message.metadata.user = 'test_user' + message.metadata.collection = 'test_collection' + + entity = EntityEmbeddings( + entity=Value(value="test_entity", is_uri=False), + vectors=[[0.1, 0.2, 0.3]] + ) + message.entities = [entity] + + # Mock index doesn't exist and never becomes ready + processor.pinecone.has_index.return_value = False + processor.pinecone.describe_index.return_value.status = {"ready": False} + + with patch('time.sleep'): # Speed up the test + with pytest.raises(RuntimeError, match="Gave up waiting for index creation"): + await processor.store_graph_embeddings(message) + + def test_add_args_method(self): + """Test that add_args properly configures argument parser""" + from argparse import ArgumentParser + from unittest.mock import patch + + parser = ArgumentParser() + + # Mock the parent class add_args method + with patch('trustgraph.storage.graph_embeddings.pinecone.write.GraphEmbeddingsStoreService.add_args') as mock_parent_add_args: + Processor.add_args(parser) + + # Verify parent add_args was called + mock_parent_add_args.assert_called_once() + + # Verify our specific arguments were added by parsing empty args + args = parser.parse_args([]) + + assert hasattr(args, 'api_key') + assert args.api_key == 'not-specified' # Default value when no env var + assert hasattr(args, 'url') + assert args.url is None + assert hasattr(args, 'cloud') + assert args.cloud == 'aws' + assert hasattr(args, 'region') + assert args.region == 'us-east-1' + + def test_add_args_with_custom_values(self): + """Test add_args with custom command line values""" + from argparse import ArgumentParser + from unittest.mock import patch + + parser = ArgumentParser() + + with patch('trustgraph.storage.graph_embeddings.pinecone.write.GraphEmbeddingsStoreService.add_args'): + Processor.add_args(parser) + + # Test parsing with custom values + args = parser.parse_args([ + '--api-key', 'custom-api-key', + '--url', 'https://custom-host.pinecone.io', + '--cloud', 'gcp', + '--region', 'us-west1' + ]) + + assert args.api_key == 'custom-api-key' + assert args.url == 'https://custom-host.pinecone.io' + assert args.cloud == 'gcp' + assert args.region == 'us-west1' + + def test_add_args_short_form(self): + """Test add_args with short form arguments""" + from argparse import ArgumentParser + from unittest.mock import patch + + parser = ArgumentParser() + + with patch('trustgraph.storage.graph_embeddings.pinecone.write.GraphEmbeddingsStoreService.add_args'): + Processor.add_args(parser) + + # Test parsing with short form + args = parser.parse_args([ + '-a', 'short-api-key', + '-u', 'https://short-host.pinecone.io' + ]) + + assert args.api_key == 'short-api-key' + assert args.url == 'https://short-host.pinecone.io' + + @patch('trustgraph.storage.graph_embeddings.pinecone.write.Processor.launch') + def test_run_function(self, mock_launch): + """Test the run function calls Processor.launch with correct parameters""" + from trustgraph.storage.graph_embeddings.pinecone.write import run, default_ident + + run() + + mock_launch.assert_called_once_with( + default_ident, + "\nAccepts entity/vector pairs and writes them to a Pinecone store.\n" + ) \ No newline at end of file diff --git a/trustgraph-flow/trustgraph/query/graph_embeddings/pinecone/service.py b/trustgraph-flow/trustgraph/query/graph_embeddings/pinecone/service.py index 942a1e69..94781fc1 100755 --- a/trustgraph-flow/trustgraph/query/graph_embeddings/pinecone/service.py +++ b/trustgraph-flow/trustgraph/query/graph_embeddings/pinecone/service.py @@ -10,30 +10,23 @@ from pinecone.grpc import PineconeGRPC, GRPCClientConfig import uuid import os -from .... schema import GraphEmbeddingsRequest, GraphEmbeddingsResponse +from .... schema import GraphEmbeddingsResponse from .... schema import Error, Value -from .... schema import graph_embeddings_request_queue -from .... schema import graph_embeddings_response_queue -from .... base import ConsumerProducer +from .... base import GraphEmbeddingsQueryService -module = "ge-query" - -default_input_queue = graph_embeddings_request_queue -default_output_queue = graph_embeddings_response_queue -default_subscriber = module +default_ident = "ge-query" default_api_key = os.getenv("PINECONE_API_KEY", "not-specified") -class Processor(ConsumerProducer): +class Processor(GraphEmbeddingsQueryService): def __init__(self, **params): - input_queue = params.get("input_queue", default_input_queue) - output_queue = params.get("output_queue", default_output_queue) - subscriber = params.get("subscriber", default_subscriber) - self.url = params.get("url", None) self.api_key = params.get("api_key", default_api_key) + if self.api_key is None or self.api_key == "not-specified": + raise RuntimeError("Pinecone API key must be specified") + if self.url: self.pinecone = PineconeGRPC( @@ -47,12 +40,8 @@ class Processor(ConsumerProducer): super(Processor, self).__init__( **params | { - "input_queue": input_queue, - "output_queue": output_queue, - "subscriber": subscriber, - "input_schema": GraphEmbeddingsRequest, - "output_schema": GraphEmbeddingsResponse, "url": self.url, + "api_key": self.api_key, } ) @@ -62,26 +51,23 @@ class Processor(ConsumerProducer): else: return Value(value=ent, is_uri=False) - async def handle(self, msg): + async def query_graph_embeddings(self, msg): try: - v = msg.value() - - # Sender-produced ID - id = msg.properties()["id"] - - print(f"Handling input {id}...", flush=True) + # Handle zero limit case + if msg.limit <= 0: + return [] entity_set = set() entities = [] - for vec in v.vectors: + for vec in msg.vectors: dim = len(vec) index_name = ( - "t-" + v.user + "-" + str(dim) + "t-" + msg.user + "-" + msg.collection + "-" + str(dim) ) index = self.pinecone.Index(index_name) @@ -89,9 +75,8 @@ class Processor(ConsumerProducer): # Heuristic hack, get (2*limit), so that we have more chance # of getting (limit) entities results = index.query( - namespace=v.collection, vector=vec, - top_k=v.limit * 2, + top_k=msg.limit * 2, include_values=False, include_metadata=True ) @@ -106,10 +91,10 @@ class Processor(ConsumerProducer): entities.append(ent) # Keep adding entities until limit - if len(entity_set) >= v.limit: break + if len(entity_set) >= msg.limit: break # Keep adding entities until limit - if len(entity_set) >= v.limit: break + if len(entity_set) >= msg.limit: break ents2 = [] @@ -118,37 +103,17 @@ class Processor(ConsumerProducer): entities = ents2 - print("Send response...", flush=True) - r = GraphEmbeddingsResponse(entities=entities, error=None) - await self.send(r, properties={"id": id}) - - print("Done.", flush=True) + return entities except Exception as e: print(f"Exception: {e}") - - print("Send error response...", flush=True) - - r = GraphEmbeddingsResponse( - error=Error( - type = "llm-error", - message = str(e), - ), - entities=None, - ) - - await self.send(r, properties={"id": id}) - - self.consumer.acknowledge(msg) + raise e @staticmethod def add_args(parser): - ConsumerProducer.add_args( - parser, default_input_queue, default_subscriber, - default_output_queue, - ) + GraphEmbeddingsQueryService.add_args(parser) parser.add_argument( '-a', '--api-key', @@ -163,5 +128,5 @@ class Processor(ConsumerProducer): def run(): - Processor.launch(module, __doc__) + Processor.launch(default_ident, __doc__) diff --git a/trustgraph-flow/trustgraph/storage/graph_embeddings/pinecone/write.py b/trustgraph-flow/trustgraph/storage/graph_embeddings/pinecone/write.py index 400acf26..e575d12a 100755 --- a/trustgraph-flow/trustgraph/storage/graph_embeddings/pinecone/write.py +++ b/trustgraph-flow/trustgraph/storage/graph_embeddings/pinecone/write.py @@ -10,32 +10,23 @@ import time import uuid import os -from .... schema import GraphEmbeddings -from .... schema import graph_embeddings_store_queue -from .... log_level import LogLevel -from .... base import Consumer +from .... base import GraphEmbeddingsStoreService -module = "ge-write" - -default_input_queue = graph_embeddings_store_queue -default_subscriber = module +default_ident = "ge-write" default_api_key = os.getenv("PINECONE_API_KEY", "not-specified") default_cloud = "aws" default_region = "us-east-1" -class Processor(Consumer): +class Processor(GraphEmbeddingsStoreService): def __init__(self, **params): - input_queue = params.get("input_queue", default_input_queue) - subscriber = params.get("subscriber", default_subscriber) - self.url = params.get("url", None) self.cloud = params.get("cloud", default_cloud) self.region = params.get("region", default_region) self.api_key = params.get("api_key", default_api_key) - if self.api_key is None: + if self.api_key is None or self.api_key == "not-specified": raise RuntimeError("Pinecone API key must be specified") if self.url: @@ -51,10 +42,10 @@ class Processor(Consumer): super(Processor, self).__init__( **params | { - "input_queue": input_queue, - "subscriber": subscriber, - "input_schema": GraphEmbeddings, "url": self.url, + "cloud": self.cloud, + "region": self.region, + "api_key": self.api_key, } ) @@ -88,13 +79,9 @@ class Processor(Consumer): "Gave up waiting for index creation" ) - async def handle(self, msg): + async def store_graph_embeddings(self, message): - v = msg.value() - - id = str(uuid.uuid4()) - - for entity in v.entities: + for entity in message.entities: if entity.entity.value == "" or entity.entity.value is None: continue @@ -104,7 +91,7 @@ class Processor(Consumer): dim = len(vec) index_name = ( - "t-" + v.metadata.user + "-" + str(dim) + "t-" + message.metadata.user + "-" + message.metadata.collection + "-" + str(dim) ) if index_name != self.last_index_name: @@ -125,9 +112,12 @@ class Processor(Consumer): index = self.pinecone.Index(index_name) + # Generate unique ID for each vector + vector_id = str(uuid.uuid4()) + records = [ { - "id": id, + "id": vector_id, "values": vec, "metadata": { "entity": entity.entity.value }, } @@ -135,15 +125,12 @@ class Processor(Consumer): index.upsert( vectors = records, - namespace = v.metadata.collection, ) @staticmethod def add_args(parser): - Consumer.add_args( - parser, default_input_queue, default_subscriber, - ) + GraphEmbeddingsStoreService.add_args(parser) parser.add_argument( '-a', '--api-key', @@ -170,5 +157,5 @@ class Processor(Consumer): def run(): - Processor.launch(module, __doc__) + Processor.launch(default_ident, __doc__)