mirror of
https://github.com/trustgraph-ai/trustgraph.git
synced 2026-07-22 03:31:02 +02:00
Fixing storage and adding tests
This commit is contained in:
parent
0768ac2c9b
commit
fcd8246d8e
4 changed files with 874 additions and 75 deletions
490
tests/unit/test_query/test_graph_embeddings_milvus_query.py
Normal file
490
tests/unit/test_query/test_graph_embeddings_milvus_query.py
Normal file
|
|
@ -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
|
||||||
354
tests/unit/test_storage/test_graph_embeddings_milvus_storage.py
Normal file
354
tests/unit/test_storage/test_graph_embeddings_milvus_storage.py
Normal file
|
|
@ -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"
|
||||||
|
)
|
||||||
|
|
@ -5,35 +5,21 @@ entities
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from .... direct.milvus_graph_embeddings import EntityVectors
|
from .... direct.milvus_graph_embeddings import EntityVectors
|
||||||
from .... schema import GraphEmbeddingsRequest, GraphEmbeddingsResponse
|
from .... schema import GraphEmbeddingsResponse
|
||||||
from .... schema import Error, Value
|
from .... schema import Error, Value
|
||||||
from .... schema import graph_embeddings_request_queue
|
from .... base import GraphEmbeddingsQueryService
|
||||||
from .... schema import graph_embeddings_response_queue
|
|
||||||
from .... base import ConsumerProducer
|
|
||||||
|
|
||||||
module = "ge-query"
|
default_ident = "ge-query"
|
||||||
|
|
||||||
default_input_queue = graph_embeddings_request_queue
|
|
||||||
default_output_queue = graph_embeddings_response_queue
|
|
||||||
default_subscriber = module
|
|
||||||
default_store_uri = 'http://localhost:19530'
|
default_store_uri = 'http://localhost:19530'
|
||||||
|
|
||||||
class Processor(ConsumerProducer):
|
class Processor(GraphEmbeddingsQueryService):
|
||||||
|
|
||||||
def __init__(self, **params):
|
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)
|
store_uri = params.get("store_uri", default_store_uri)
|
||||||
|
|
||||||
super(Processor, self).__init__(
|
super(Processor, self).__init__(
|
||||||
**params | {
|
**params | {
|
||||||
"input_queue": input_queue,
|
|
||||||
"output_queue": output_queue,
|
|
||||||
"subscriber": subscriber,
|
|
||||||
"input_schema": GraphEmbeddingsRequest,
|
|
||||||
"output_schema": GraphEmbeddingsResponse,
|
|
||||||
"store_uri": store_uri,
|
"store_uri": store_uri,
|
||||||
}
|
}
|
||||||
)
|
)
|
||||||
|
|
@ -46,29 +32,30 @@ class Processor(ConsumerProducer):
|
||||||
else:
|
else:
|
||||||
return Value(value=ent, is_uri=False)
|
return Value(value=ent, is_uri=False)
|
||||||
|
|
||||||
async def handle(self, msg):
|
async def query_graph_embeddings(self, msg):
|
||||||
|
|
||||||
try:
|
try:
|
||||||
|
|
||||||
v = msg.value()
|
entity_set = set()
|
||||||
|
entities = []
|
||||||
|
|
||||||
# Sender-produced ID
|
for vec in msg.vectors:
|
||||||
id = msg.properties()["id"]
|
|
||||||
|
|
||||||
print(f"Handling input {id}...", flush=True)
|
resp = self.vecstore.search(vec, limit=msg.limit * 2)
|
||||||
|
|
||||||
entities = set()
|
|
||||||
|
|
||||||
for vec in v.vectors:
|
|
||||||
|
|
||||||
resp = self.vecstore.search(vec, limit=v.limit)
|
|
||||||
|
|
||||||
for r in resp:
|
for r in resp:
|
||||||
ent = r["entity"]["entity"]
|
ent = r["entity"]["entity"]
|
||||||
entities.add(ent)
|
|
||||||
|
|
||||||
# Convert set to list
|
# De-dupe entities
|
||||||
entities = list(entities)
|
if ent not in entity_set:
|
||||||
|
entity_set.add(ent)
|
||||||
|
entities.append(ent)
|
||||||
|
|
||||||
|
# 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 = []
|
ents2 = []
|
||||||
|
|
||||||
|
|
@ -78,36 +65,19 @@ class Processor(ConsumerProducer):
|
||||||
entities = ents2
|
entities = ents2
|
||||||
|
|
||||||
print("Send response...", flush=True)
|
print("Send response...", flush=True)
|
||||||
r = GraphEmbeddingsResponse(entities=entities, error=None)
|
return entities
|
||||||
await self.send(r, properties={"id": id})
|
|
||||||
|
|
||||||
print("Done.", flush=True)
|
print("Done.", flush=True)
|
||||||
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
|
|
||||||
print(f"Exception: {e}")
|
print(f"Exception: {e}")
|
||||||
|
raise 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)
|
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def add_args(parser):
|
def add_args(parser):
|
||||||
|
|
||||||
ConsumerProducer.add_args(
|
GraphEmbeddingsQueryService.add_args(parser)
|
||||||
parser, default_input_queue, default_subscriber,
|
|
||||||
default_output_queue,
|
|
||||||
)
|
|
||||||
|
|
||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
'-t', '--store-uri',
|
'-t', '--store-uri',
|
||||||
|
|
@ -117,5 +87,5 @@ class Processor(ConsumerProducer):
|
||||||
|
|
||||||
def run():
|
def run():
|
||||||
|
|
||||||
Processor.launch(module, __doc__)
|
Processor.launch(default_ident, __doc__)
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -3,42 +3,29 @@
|
||||||
Accepts entity/vector pairs and writes them to a Milvus store.
|
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 .... direct.milvus_graph_embeddings import EntityVectors
|
||||||
from .... base import Consumer
|
from .... base import GraphEmbeddingsStoreService
|
||||||
|
|
||||||
module = "ge-write"
|
default_ident = "ge-write"
|
||||||
|
|
||||||
default_input_queue = graph_embeddings_store_queue
|
|
||||||
default_subscriber = module
|
|
||||||
default_store_uri = 'http://localhost:19530'
|
default_store_uri = 'http://localhost:19530'
|
||||||
|
|
||||||
class Processor(Consumer):
|
class Processor(GraphEmbeddingsStoreService):
|
||||||
|
|
||||||
def __init__(self, **params):
|
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)
|
store_uri = params.get("store_uri", default_store_uri)
|
||||||
|
|
||||||
super(Processor, self).__init__(
|
super(Processor, self).__init__(
|
||||||
**params | {
|
**params | {
|
||||||
"input_queue": input_queue,
|
|
||||||
"subscriber": subscriber,
|
|
||||||
"input_schema": GraphEmbeddings,
|
|
||||||
"store_uri": store_uri,
|
"store_uri": store_uri,
|
||||||
}
|
}
|
||||||
)
|
)
|
||||||
|
|
||||||
self.vecstore = EntityVectors(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 message.entities:
|
||||||
|
|
||||||
for entity in v.entities:
|
|
||||||
|
|
||||||
if entity.entity.value != "" and entity.entity.value is not None:
|
if entity.entity.value != "" and entity.entity.value is not None:
|
||||||
for vec in entity.vectors:
|
for vec in entity.vectors:
|
||||||
|
|
@ -47,9 +34,7 @@ class Processor(Consumer):
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def add_args(parser):
|
def add_args(parser):
|
||||||
|
|
||||||
Consumer.add_args(
|
GraphEmbeddingsStoreService.add_args(parser)
|
||||||
parser, default_input_queue, default_subscriber,
|
|
||||||
)
|
|
||||||
|
|
||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
'-t', '--store-uri',
|
'-t', '--store-uri',
|
||||||
|
|
@ -59,5 +44,5 @@ class Processor(Consumer):
|
||||||
|
|
||||||
def run():
|
def run():
|
||||||
|
|
||||||
Processor.launch(module, __doc__)
|
Processor.launch(default_ident, __doc__)
|
||||||
|
|
||||||
|
|
|
||||||
Loading…
Add table
Add a link
Reference in a new issue