diff --git a/trustgraph-flow/trustgraph/storage/graph_embeddings/milvus/write.py b/trustgraph-flow/trustgraph/storage/graph_embeddings/milvus/write.py index 2e192cd6..148e866a 100755 --- a/trustgraph-flow/trustgraph/storage/graph_embeddings/milvus/write.py +++ b/trustgraph-flow/trustgraph/storage/graph_embeddings/milvus/write.py @@ -9,10 +9,24 @@ from .... direct.milvus_graph_embeddings import EntityVectors from .... base import GraphEmbeddingsStoreService, CollectionConfigHandler from .... base import AsyncProcessor, Consumer, Producer from .... base import ConsumerMetrics, ProducerMetrics +from .... schema import IRI, LITERAL # Module logger logger = logging.getLogger(__name__) + +def get_term_value(term): + """Extract the string value from a Term""" + if term is None: + return None + if term.type == IRI: + return term.iri + elif term.type == LITERAL: + return term.value + else: + # For blank nodes or other types, use id or value + return term.id or term.value + default_ident = "ge-write" default_store_uri = 'http://localhost:19530' @@ -36,11 +50,12 @@ class Processor(CollectionConfigHandler, GraphEmbeddingsStoreService): async def store_graph_embeddings(self, message): for entity in message.entities: + entity_value = get_term_value(entity.entity) - if entity.entity.value != "" and entity.entity.value is not None: + if entity_value != "" and entity_value is not None: for vec in entity.vectors: self.vecstore.insert( - vec, entity.entity.value, + vec, entity_value, message.metadata.user, message.metadata.collection ) diff --git a/trustgraph-flow/trustgraph/storage/graph_embeddings/pinecone/write.py b/trustgraph-flow/trustgraph/storage/graph_embeddings/pinecone/write.py index 0bee6ceb..c92d7661 100755 --- a/trustgraph-flow/trustgraph/storage/graph_embeddings/pinecone/write.py +++ b/trustgraph-flow/trustgraph/storage/graph_embeddings/pinecone/write.py @@ -14,10 +14,24 @@ import logging from .... base import GraphEmbeddingsStoreService, CollectionConfigHandler from .... base import AsyncProcessor, Consumer, Producer from .... base import ConsumerMetrics, ProducerMetrics +from .... schema import IRI, LITERAL # Module logger logger = logging.getLogger(__name__) + +def get_term_value(term): + """Extract the string value from a Term""" + if term is None: + return None + if term.type == IRI: + return term.iri + elif term.type == LITERAL: + return term.value + else: + # For blank nodes or other types, use id or value + return term.id or term.value + default_ident = "ge-write" default_api_key = os.getenv("PINECONE_API_KEY", "not-specified") default_cloud = "aws" @@ -100,8 +114,9 @@ class Processor(CollectionConfigHandler, GraphEmbeddingsStoreService): return for entity in message.entities: + entity_value = get_term_value(entity.entity) - if entity.entity.value == "" or entity.entity.value is None: + if entity_value == "" or entity_value is None: continue for vec in entity.vectors: @@ -126,7 +141,7 @@ class Processor(CollectionConfigHandler, GraphEmbeddingsStoreService): { "id": vector_id, "values": vec, - "metadata": { "entity": entity.entity.value }, + "metadata": { "entity": entity_value }, } ]