From fcd8246d8e5fc309aa710d7c3a96cd6363bff3b5 Mon Sep 17 00:00:00 2001 From: Cyber MacGeddon Date: Mon, 14 Jul 2025 21:31:05 +0100 Subject: [PATCH] Fixing storage and adding tests --- .../test_graph_embeddings_milvus_query.py | 490 ++++++++++++++++++ .../test_graph_embeddings_milvus_storage.py | 354 +++++++++++++ .../query/graph_embeddings/milvus/service.py | 76 +-- .../storage/graph_embeddings/milvus/write.py | 29 +- 4 files changed, 874 insertions(+), 75 deletions(-) create mode 100644 tests/unit/test_query/test_graph_embeddings_milvus_query.py create mode 100644 tests/unit/test_storage/test_graph_embeddings_milvus_storage.py diff --git a/tests/unit/test_query/test_graph_embeddings_milvus_query.py b/tests/unit/test_query/test_graph_embeddings_milvus_query.py new file mode 100644 index 00000000..6d856c96 --- /dev/null +++ b/tests/unit/test_query/test_graph_embeddings_milvus_query.py @@ -0,0 +1,490 @@ +""" +Tests for Milvus graph embeddings query service +""" + +import pytest +from unittest.mock import MagicMock, patch + +from trustgraph.query.graph_embeddings.milvus.service import Processor +from trustgraph.schema import Value, GraphEmbeddingsRequest + + +class TestMilvusGraphEmbeddingsQueryProcessor: + """Test cases for Milvus graph embeddings query processor""" + + @pytest.fixture + def processor(self): + """Create a processor instance for testing""" + with patch('trustgraph.query.graph_embeddings.milvus.service.EntityVectors') as mock_entity_vectors: + mock_vecstore = MagicMock() + mock_entity_vectors.return_value = mock_vecstore + + processor = Processor( + taskgroup=MagicMock(), + id='test-milvus-ge-query', + store_uri='http://localhost:19530' + ) + + return processor + + @pytest.fixture + def mock_query_request(self): + """Create a mock query request for testing""" + query = GraphEmbeddingsRequest( + user='test_user', + collection='test_collection', + vectors=[[0.1, 0.2, 0.3], [0.4, 0.5, 0.6]], + limit=10 + ) + return query + + @patch('trustgraph.query.graph_embeddings.milvus.service.EntityVectors') + def test_processor_initialization_with_defaults(self, mock_entity_vectors): + """Test processor initialization with default parameters""" + taskgroup_mock = MagicMock() + mock_vecstore = MagicMock() + mock_entity_vectors.return_value = mock_vecstore + + processor = Processor(taskgroup=taskgroup_mock) + + mock_entity_vectors.assert_called_once_with('http://localhost:19530') + assert processor.vecstore == mock_vecstore + + @patch('trustgraph.query.graph_embeddings.milvus.service.EntityVectors') + def test_processor_initialization_with_custom_params(self, mock_entity_vectors): + """Test processor initialization with custom parameters""" + taskgroup_mock = MagicMock() + mock_vecstore = MagicMock() + mock_entity_vectors.return_value = mock_vecstore + + processor = Processor( + taskgroup=taskgroup_mock, + store_uri='http://custom-milvus:19530' + ) + + mock_entity_vectors.assert_called_once_with('http://custom-milvus:19530') + assert processor.vecstore == mock_vecstore + + def test_create_value_with_http_uri(self, processor): + """Test create_value with HTTP URI""" + result = processor.create_value("http://example.com/resource") + + assert isinstance(result, Value) + assert result.value == "http://example.com/resource" + assert result.is_uri is True + + def test_create_value_with_https_uri(self, processor): + """Test create_value with HTTPS URI""" + result = processor.create_value("https://example.com/resource") + + assert isinstance(result, Value) + assert result.value == "https://example.com/resource" + assert result.is_uri is True + + def test_create_value_with_literal(self, processor): + """Test create_value with literal value""" + result = processor.create_value("just a literal string") + + assert isinstance(result, Value) + assert result.value == "just a literal string" + assert result.is_uri is False + + def test_create_value_with_empty_string(self, processor): + """Test create_value with empty string""" + result = processor.create_value("") + + assert isinstance(result, Value) + assert result.value == "" + assert result.is_uri is False + + def test_create_value_with_partial_uri(self, processor): + """Test create_value with string that looks like URI but isn't complete""" + result = processor.create_value("http") + + assert isinstance(result, Value) + assert result.value == "http" + assert result.is_uri is False + + def test_create_value_with_ftp_uri(self, processor): + """Test create_value with FTP URI (should not be detected as URI)""" + result = processor.create_value("ftp://example.com/file") + + assert isinstance(result, Value) + assert result.value == "ftp://example.com/file" + assert result.is_uri is False + + @pytest.mark.asyncio + async def test_query_graph_embeddings_single_vector(self, processor): + """Test querying graph embeddings with a single vector""" + query = GraphEmbeddingsRequest( + user='test_user', + collection='test_collection', + vectors=[[0.1, 0.2, 0.3]], + limit=5 + ) + + # Mock search results + mock_results = [ + {"entity": {"entity": "http://example.com/entity1"}}, + {"entity": {"entity": "http://example.com/entity2"}}, + {"entity": {"entity": "literal entity"}}, + ] + processor.vecstore.search.return_value = mock_results + + result = await processor.query_graph_embeddings(query) + + # Verify search was called with correct parameters + processor.vecstore.search.assert_called_once_with([0.1, 0.2, 0.3], limit=10) + + # Verify results are converted to Value objects + assert len(result) == 3 + assert isinstance(result[0], Value) + assert result[0].value == "http://example.com/entity1" + assert result[0].is_uri is True + assert isinstance(result[1], Value) + assert result[1].value == "http://example.com/entity2" + assert result[1].is_uri is True + assert isinstance(result[2], Value) + assert result[2].value == "literal entity" + assert result[2].is_uri is False + + @pytest.mark.asyncio + async def test_query_graph_embeddings_multiple_vectors(self, processor): + """Test querying graph embeddings with multiple vectors""" + query = GraphEmbeddingsRequest( + user='test_user', + collection='test_collection', + vectors=[[0.1, 0.2, 0.3], [0.4, 0.5, 0.6]], + limit=3 + ) + + # Mock search results - different results for each vector + mock_results_1 = [ + {"entity": {"entity": "http://example.com/entity1"}}, + {"entity": {"entity": "http://example.com/entity2"}}, + ] + mock_results_2 = [ + {"entity": {"entity": "http://example.com/entity2"}}, # Duplicate + {"entity": {"entity": "http://example.com/entity3"}}, + ] + processor.vecstore.search.side_effect = [mock_results_1, mock_results_2] + + result = await processor.query_graph_embeddings(query) + + # Verify search was called twice with correct parameters + expected_calls = [ + (([0.1, 0.2, 0.3],), {"limit": 6}), + (([0.4, 0.5, 0.6],), {"limit": 6}), + ] + assert processor.vecstore.search.call_count == 2 + for i, (expected_args, expected_kwargs) in enumerate(expected_calls): + actual_call = processor.vecstore.search.call_args_list[i] + assert actual_call[0] == expected_args + assert actual_call[1] == expected_kwargs + + # Verify results are deduplicated and limited + assert len(result) == 3 + entity_values = [r.value for r in result] + assert "http://example.com/entity1" in entity_values + assert "http://example.com/entity2" in entity_values + assert "http://example.com/entity3" in entity_values + + @pytest.mark.asyncio + async def test_query_graph_embeddings_with_limit(self, processor): + """Test querying graph embeddings respects limit parameter""" + query = GraphEmbeddingsRequest( + user='test_user', + collection='test_collection', + vectors=[[0.1, 0.2, 0.3]], + limit=2 + ) + + # Mock search results - more results than limit + mock_results = [ + {"entity": {"entity": "http://example.com/entity1"}}, + {"entity": {"entity": "http://example.com/entity2"}}, + {"entity": {"entity": "http://example.com/entity3"}}, + {"entity": {"entity": "http://example.com/entity4"}}, + ] + processor.vecstore.search.return_value = mock_results + + result = await processor.query_graph_embeddings(query) + + # Verify search was called with 2*limit for better deduplication + processor.vecstore.search.assert_called_once_with([0.1, 0.2, 0.3], limit=4) + + # Verify results are limited to the requested limit + assert len(result) == 2 + + @pytest.mark.asyncio + async def test_query_graph_embeddings_deduplication(self, processor): + """Test that duplicate entities are properly deduplicated""" + query = GraphEmbeddingsRequest( + user='test_user', + collection='test_collection', + vectors=[[0.1, 0.2, 0.3], [0.4, 0.5, 0.6]], + limit=5 + ) + + # Mock search results with duplicates + mock_results_1 = [ + {"entity": {"entity": "http://example.com/entity1"}}, + {"entity": {"entity": "http://example.com/entity2"}}, + ] + mock_results_2 = [ + {"entity": {"entity": "http://example.com/entity2"}}, # Duplicate + {"entity": {"entity": "http://example.com/entity1"}}, # Duplicate + {"entity": {"entity": "http://example.com/entity3"}}, # New + ] + processor.vecstore.search.side_effect = [mock_results_1, mock_results_2] + + result = await processor.query_graph_embeddings(query) + + # Verify duplicates are removed + assert len(result) == 3 + entity_values = [r.value for r in result] + assert len(set(entity_values)) == 3 # All unique + assert "http://example.com/entity1" in entity_values + assert "http://example.com/entity2" in entity_values + assert "http://example.com/entity3" in entity_values + + @pytest.mark.asyncio + async def test_query_graph_embeddings_early_termination_on_limit(self, processor): + """Test that querying stops early when limit is reached""" + query = GraphEmbeddingsRequest( + user='test_user', + collection='test_collection', + vectors=[[0.1, 0.2, 0.3], [0.4, 0.5, 0.6]], + limit=2 + ) + + # Mock search results - first vector returns enough results + mock_results_1 = [ + {"entity": {"entity": "http://example.com/entity1"}}, + {"entity": {"entity": "http://example.com/entity2"}}, + {"entity": {"entity": "http://example.com/entity3"}}, + ] + processor.vecstore.search.return_value = mock_results_1 + + result = await processor.query_graph_embeddings(query) + + # Verify only first vector was searched (limit reached) + processor.vecstore.search.assert_called_once_with([0.1, 0.2, 0.3], limit=4) + + # Verify results are limited + assert len(result) == 2 + + @pytest.mark.asyncio + async def test_query_graph_embeddings_empty_vectors(self, processor): + """Test querying graph embeddings with empty vectors list""" + query = GraphEmbeddingsRequest( + user='test_user', + collection='test_collection', + vectors=[], + limit=5 + ) + + result = await processor.query_graph_embeddings(query) + + # Verify no search was called + processor.vecstore.search.assert_not_called() + + # Verify empty results + assert len(result) == 0 + + @pytest.mark.asyncio + async def test_query_graph_embeddings_empty_search_results(self, processor): + """Test querying graph embeddings with empty search results""" + query = GraphEmbeddingsRequest( + user='test_user', + collection='test_collection', + vectors=[[0.1, 0.2, 0.3]], + limit=5 + ) + + # Mock empty search results + processor.vecstore.search.return_value = [] + + result = await processor.query_graph_embeddings(query) + + # Verify search was called + processor.vecstore.search.assert_called_once_with([0.1, 0.2, 0.3], limit=10) + + # Verify empty results + assert len(result) == 0 + + @pytest.mark.asyncio + async def test_query_graph_embeddings_mixed_uri_literal_results(self, processor): + """Test querying graph embeddings with mixed URI and literal results""" + query = GraphEmbeddingsRequest( + user='test_user', + collection='test_collection', + vectors=[[0.1, 0.2, 0.3]], + limit=5 + ) + + # Mock search results with mixed types + mock_results = [ + {"entity": {"entity": "http://example.com/uri_entity"}}, + {"entity": {"entity": "literal entity text"}}, + {"entity": {"entity": "https://example.com/another_uri"}}, + {"entity": {"entity": "another literal"}}, + ] + processor.vecstore.search.return_value = mock_results + + result = await processor.query_graph_embeddings(query) + + # Verify all results are properly typed + assert len(result) == 4 + + # Check URI entities + uri_results = [r for r in result if r.is_uri] + assert len(uri_results) == 2 + uri_values = [r.value for r in uri_results] + assert "http://example.com/uri_entity" in uri_values + assert "https://example.com/another_uri" in uri_values + + # Check literal entities + literal_results = [r for r in result if not r.is_uri] + assert len(literal_results) == 2 + literal_values = [r.value for r in literal_results] + assert "literal entity text" in literal_values + assert "another literal" in literal_values + + @pytest.mark.asyncio + async def test_query_graph_embeddings_exception_handling(self, processor): + """Test exception handling during query processing""" + query = GraphEmbeddingsRequest( + user='test_user', + collection='test_collection', + vectors=[[0.1, 0.2, 0.3]], + limit=5 + ) + + # Mock search to raise exception + processor.vecstore.search.side_effect = Exception("Milvus connection failed") + + # Should raise the exception + with pytest.raises(Exception, match="Milvus connection failed"): + await processor.query_graph_embeddings(query) + + 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.milvus.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 + # Parse empty args to check defaults + args = parser.parse_args([]) + + assert hasattr(args, 'store_uri') + assert args.store_uri == 'http://localhost:19530' + + 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.milvus.service.GraphEmbeddingsQueryService.add_args'): + Processor.add_args(parser) + + # Test parsing with custom values + args = parser.parse_args([ + '--store-uri', 'http://custom-milvus:19530' + ]) + + assert args.store_uri == 'http://custom-milvus:19530' + + 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.milvus.service.GraphEmbeddingsQueryService.add_args'): + Processor.add_args(parser) + + # Test parsing with short form + args = parser.parse_args(['-t', 'http://short-milvus:19530']) + + assert args.store_uri == 'http://short-milvus:19530' + + @patch('trustgraph.query.graph_embeddings.milvus.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.milvus.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\n" + ) + + @pytest.mark.asyncio + async def test_query_graph_embeddings_zero_limit(self, processor): + """Test querying graph embeddings with zero limit""" + query = GraphEmbeddingsRequest( + user='test_user', + collection='test_collection', + vectors=[[0.1, 0.2, 0.3]], + limit=0 + ) + + # Mock search results + mock_results = [ + {"entity": {"entity": "http://example.com/entity1"}}, + ] + processor.vecstore.search.return_value = mock_results + + result = await processor.query_graph_embeddings(query) + + # Verify search was called with 0 limit + processor.vecstore.search.assert_called_once_with([0.1, 0.2, 0.3], limit=0) + + # Verify empty results due to zero limit + assert len(result) == 0 + + @pytest.mark.asyncio + async def test_query_graph_embeddings_different_vector_dimensions(self, processor): + """Test querying graph embeddings with different vector dimensions""" + query = GraphEmbeddingsRequest( + user='test_user', + collection='test_collection', + 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 + ], + limit=5 + ) + + # Mock search results for each vector + mock_results_1 = [{"entity": {"entity": "entity_2d"}}] + mock_results_2 = [{"entity": {"entity": "entity_4d"}}] + mock_results_3 = [{"entity": {"entity": "entity_3d"}}] + processor.vecstore.search.side_effect = [mock_results_1, mock_results_2, mock_results_3] + + result = await processor.query_graph_embeddings(query) + + # Verify all vectors were searched + assert processor.vecstore.search.call_count == 3 + + # Verify results from all dimensions + assert len(result) == 3 + entity_values = [r.value for r in result] + assert "entity_2d" in entity_values + assert "entity_4d" in entity_values + assert "entity_3d" in entity_values \ No newline at end of file diff --git a/tests/unit/test_storage/test_graph_embeddings_milvus_storage.py b/tests/unit/test_storage/test_graph_embeddings_milvus_storage.py new file mode 100644 index 00000000..ae300574 --- /dev/null +++ b/tests/unit/test_storage/test_graph_embeddings_milvus_storage.py @@ -0,0 +1,354 @@ +""" +Tests for Milvus graph embeddings storage service +""" + +import pytest +from unittest.mock import MagicMock, patch + +from trustgraph.storage.graph_embeddings.milvus.write import Processor +from trustgraph.schema import Value, EntityEmbeddings + + +class TestMilvusGraphEmbeddingsStorageProcessor: + """Test cases for Milvus 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 entities with embeddings + entity1 = EntityEmbeddings( + entity=Value(value='http://example.com/entity1', is_uri=True), + vectors=[[0.1, 0.2, 0.3], [0.4, 0.5, 0.6]] + ) + entity2 = EntityEmbeddings( + entity=Value(value='literal entity', 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.milvus.write.EntityVectors') as mock_entity_vectors: + mock_vecstore = MagicMock() + mock_entity_vectors.return_value = mock_vecstore + + processor = Processor( + taskgroup=MagicMock(), + id='test-milvus-ge-storage', + store_uri='http://localhost:19530' + ) + + return processor + + @patch('trustgraph.storage.graph_embeddings.milvus.write.EntityVectors') + def test_processor_initialization_with_defaults(self, mock_entity_vectors): + """Test processor initialization with default parameters""" + taskgroup_mock = MagicMock() + mock_vecstore = MagicMock() + mock_entity_vectors.return_value = mock_vecstore + + processor = Processor(taskgroup=taskgroup_mock) + + mock_entity_vectors.assert_called_once_with('http://localhost:19530') + assert processor.vecstore == mock_vecstore + + @patch('trustgraph.storage.graph_embeddings.milvus.write.EntityVectors') + def test_processor_initialization_with_custom_params(self, mock_entity_vectors): + """Test processor initialization with custom parameters""" + taskgroup_mock = MagicMock() + mock_vecstore = MagicMock() + mock_entity_vectors.return_value = mock_vecstore + + processor = Processor( + taskgroup=taskgroup_mock, + store_uri='http://custom-milvus:19530' + ) + + mock_entity_vectors.assert_called_once_with('http://custom-milvus:19530') + assert processor.vecstore == mock_vecstore + + @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.com/entity', is_uri=True), + vectors=[[0.1, 0.2, 0.3], [0.4, 0.5, 0.6]] + ) + message.entities = [entity] + + await processor.store_graph_embeddings(message) + + # Verify insert was called for each vector + expected_calls = [ + ([0.1, 0.2, 0.3], 'http://example.com/entity'), + ([0.4, 0.5, 0.6], 'http://example.com/entity'), + ] + + assert processor.vecstore.insert.call_count == 2 + for i, (expected_vec, expected_entity) in enumerate(expected_calls): + actual_call = processor.vecstore.insert.call_args_list[i] + assert actual_call[0][0] == expected_vec + assert actual_call[0][1] == expected_entity + + @pytest.mark.asyncio + async def test_store_graph_embeddings_multiple_entities(self, processor, mock_message): + """Test storing graph embeddings for multiple entities""" + await processor.store_graph_embeddings(mock_message) + + # Verify insert was called for each vector of each entity + expected_calls = [ + # Entity 1 vectors + ([0.1, 0.2, 0.3], 'http://example.com/entity1'), + ([0.4, 0.5, 0.6], 'http://example.com/entity1'), + # Entity 2 vectors + ([0.7, 0.8, 0.9], 'literal entity'), + ] + + assert processor.vecstore.insert.call_count == 3 + for i, (expected_vec, expected_entity) in enumerate(expected_calls): + actual_call = processor.vecstore.insert.call_args_list[i] + assert actual_call[0][0] == expected_vec + assert actual_call[0][1] == expected_entity + + @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] + + await processor.store_graph_embeddings(message) + + # Verify no insert was called for empty entity + processor.vecstore.insert.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] + + await processor.store_graph_embeddings(message) + + # Verify no insert was called for None entity + processor.vecstore.insert.assert_not_called() + + @pytest.mark.asyncio + async def test_store_graph_embeddings_mixed_valid_invalid_entities(self, processor): + """Test storing graph embeddings with mix of valid and invalid entities""" + message = MagicMock() + message.metadata = MagicMock() + message.metadata.user = 'test_user' + message.metadata.collection = 'test_collection' + + valid_entity = EntityEmbeddings( + entity=Value(value='http://example.com/valid', is_uri=True), + vectors=[[0.1, 0.2, 0.3]] + ) + empty_entity = EntityEmbeddings( + entity=Value(value='', is_uri=False), + vectors=[[0.4, 0.5, 0.6]] + ) + none_entity = EntityEmbeddings( + entity=Value(value=None, is_uri=False), + vectors=[[0.7, 0.8, 0.9]] + ) + message.entities = [valid_entity, empty_entity, none_entity] + + await processor.store_graph_embeddings(message) + + # Verify only valid entity was inserted + processor.vecstore.insert.assert_called_once_with( + [0.1, 0.2, 0.3], 'http://example.com/valid' + ) + + @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 = [] + + await processor.store_graph_embeddings(message) + + # Verify no insert was called + processor.vecstore.insert.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='http://example.com/entity', is_uri=True), + vectors=[] + ) + message.entities = [entity] + + await processor.store_graph_embeddings(message) + + # Verify no insert was called (no vectors to insert) + processor.vecstore.insert.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='http://example.com/entity', is_uri=True), + 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] + + await processor.store_graph_embeddings(message) + + # Verify all vectors were inserted regardless of dimension + expected_calls = [ + ([0.1, 0.2], 'http://example.com/entity'), + ([0.3, 0.4, 0.5, 0.6], 'http://example.com/entity'), + ([0.7, 0.8, 0.9], 'http://example.com/entity'), + ] + + assert processor.vecstore.insert.call_count == 3 + for i, (expected_vec, expected_entity) in enumerate(expected_calls): + actual_call = processor.vecstore.insert.call_args_list[i] + assert actual_call[0][0] == expected_vec + assert actual_call[0][1] == expected_entity + + @pytest.mark.asyncio + async def test_store_graph_embeddings_uri_and_literal_entities(self, processor): + """Test storing graph embeddings for both URI and literal entities""" + message = MagicMock() + message.metadata = MagicMock() + message.metadata.user = 'test_user' + message.metadata.collection = 'test_collection' + + uri_entity = EntityEmbeddings( + entity=Value(value='http://example.com/uri_entity', is_uri=True), + vectors=[[0.1, 0.2, 0.3]] + ) + literal_entity = EntityEmbeddings( + entity=Value(value='literal entity text', is_uri=False), + vectors=[[0.4, 0.5, 0.6]] + ) + message.entities = [uri_entity, literal_entity] + + await processor.store_graph_embeddings(message) + + # Verify both entities were inserted + expected_calls = [ + ([0.1, 0.2, 0.3], 'http://example.com/uri_entity'), + ([0.4, 0.5, 0.6], 'literal entity text'), + ] + + assert processor.vecstore.insert.call_count == 2 + for i, (expected_vec, expected_entity) in enumerate(expected_calls): + actual_call = processor.vecstore.insert.call_args_list[i] + assert actual_call[0][0] == expected_vec + assert actual_call[0][1] == expected_entity + + 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.milvus.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 + # Parse empty args to check defaults + args = parser.parse_args([]) + + assert hasattr(args, 'store_uri') + assert args.store_uri == 'http://localhost:19530' + + 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.milvus.write.GraphEmbeddingsStoreService.add_args'): + Processor.add_args(parser) + + # Test parsing with custom values + args = parser.parse_args([ + '--store-uri', 'http://custom-milvus:19530' + ]) + + assert args.store_uri == 'http://custom-milvus:19530' + + 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.milvus.write.GraphEmbeddingsStoreService.add_args'): + Processor.add_args(parser) + + # Test parsing with short form + args = parser.parse_args(['-t', 'http://short-milvus:19530']) + + assert args.store_uri == 'http://short-milvus:19530' + + @patch('trustgraph.storage.graph_embeddings.milvus.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.milvus.write import run, default_ident + + run() + + mock_launch.assert_called_once_with( + default_ident, + "\nAccepts entity/vector pairs and writes them to a Milvus store.\n" + ) \ No newline at end of file diff --git a/trustgraph-flow/trustgraph/query/graph_embeddings/milvus/service.py b/trustgraph-flow/trustgraph/query/graph_embeddings/milvus/service.py index d2cec084..498ab483 100755 --- a/trustgraph-flow/trustgraph/query/graph_embeddings/milvus/service.py +++ b/trustgraph-flow/trustgraph/query/graph_embeddings/milvus/service.py @@ -5,35 +5,21 @@ entities """ from .... direct.milvus_graph_embeddings import EntityVectors -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_store_uri = 'http://localhost:19530' -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) store_uri = params.get("store_uri", default_store_uri) super(Processor, self).__init__( **params | { - "input_queue": input_queue, - "output_queue": output_queue, - "subscriber": subscriber, - "input_schema": GraphEmbeddingsRequest, - "output_schema": GraphEmbeddingsResponse, "store_uri": store_uri, } ) @@ -46,29 +32,30 @@ 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() + entity_set = set() + entities = [] - # Sender-produced ID - id = msg.properties()["id"] + for vec in msg.vectors: - print(f"Handling input {id}...", flush=True) - - entities = set() - - for vec in v.vectors: - - resp = self.vecstore.search(vec, limit=v.limit) + resp = self.vecstore.search(vec, limit=msg.limit * 2) for r in resp: ent = r["entity"]["entity"] - entities.add(ent) + + # De-dupe entities + if ent not in entity_set: + entity_set.add(ent) + entities.append(ent) - # Convert set to list - entities = list(entities) + # Keep adding entities until limit + if len(entity_set) >= msg.limit: break + + # Keep adding entities until limit + if len(entity_set) >= msg.limit: break ents2 = [] @@ -78,36 +65,19 @@ class Processor(ConsumerProducer): entities = ents2 print("Send response...", flush=True) - r = GraphEmbeddingsResponse(entities=entities, error=None) - await self.send(r, properties={"id": id}) + return entities print("Done.", flush=True) 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( '-t', '--store-uri', @@ -117,5 +87,5 @@ class Processor(ConsumerProducer): def run(): - Processor.launch(module, __doc__) + Processor.launch(default_ident, __doc__) diff --git a/trustgraph-flow/trustgraph/storage/graph_embeddings/milvus/write.py b/trustgraph-flow/trustgraph/storage/graph_embeddings/milvus/write.py index 8d8b68b0..f140ab76 100755 --- a/trustgraph-flow/trustgraph/storage/graph_embeddings/milvus/write.py +++ b/trustgraph-flow/trustgraph/storage/graph_embeddings/milvus/write.py @@ -3,42 +3,29 @@ Accepts entity/vector pairs and writes them to a Milvus store. """ -from .... schema import GraphEmbeddings -from .... schema import graph_embeddings_store_queue -from .... log_level import LogLevel from .... direct.milvus_graph_embeddings import EntityVectors -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_store_uri = 'http://localhost:19530' -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) store_uri = params.get("store_uri", default_store_uri) super(Processor, self).__init__( **params | { - "input_queue": input_queue, - "subscriber": subscriber, - "input_schema": GraphEmbeddings, "store_uri": store_uri, } ) self.vecstore = EntityVectors(store_uri) - async def handle(self, msg): + async def store_graph_embeddings(self, message): - v = msg.value() - - for entity in v.entities: + for entity in message.entities: if entity.entity.value != "" and entity.entity.value is not None: for vec in entity.vectors: @@ -47,9 +34,7 @@ class Processor(Consumer): @staticmethod def add_args(parser): - Consumer.add_args( - parser, default_input_queue, default_subscriber, - ) + GraphEmbeddingsStoreService.add_args(parser) parser.add_argument( '-t', '--store-uri', @@ -59,5 +44,5 @@ class Processor(Consumer): def run(): - Processor.launch(module, __doc__) + Processor.launch(default_ident, __doc__)