diff --git a/trustgraph-base/trustgraph/base/__init__.py b/trustgraph-base/trustgraph/base/__init__.py index d5e539c6..9086fd8c 100644 --- a/trustgraph-base/trustgraph/base/__init__.py +++ b/trustgraph-base/trustgraph/base/__init__.py @@ -17,4 +17,5 @@ from . embeddings_service import EmbeddingsService from . embeddings_client import EmbeddingsClientSpec from . text_completion_client import TextCompletionClientSpec from . prompt_client import PromptClientSpec +from . triples_store_service import TriplesStoreService diff --git a/trustgraph-base/trustgraph/base/document_embeddings_store_service.py b/trustgraph-base/trustgraph/base/document_embeddings_store_service.py new file mode 100755 index 00000000..c426988d --- /dev/null +++ b/trustgraph-base/trustgraph/base/document_embeddings_store_service.py @@ -0,0 +1,49 @@ + +""" +Document embeddings store base class +""" + +from .. schema import DocumentEmbeddings +from .. base import FlowProcessor, ConsumerSpec + +default_ident = "document-embeddings-write" + +class DocumentEmbeddingsStoreService(FlowProcessor): + + def __init__(self, **params): + + id = params.get("id") + + super(DocumentEmbeddingsStoreService, self).__init__( + **params | { "id": id } + ) + + self.register_specification( + ConsumerSpec( + name = "input", + schema = DocumentEmbeddings, + handler = self.on_message + ) + ) + + async def on_message(self, msg, consumer, flow): + + try: + + request = msg.value() + + await self.store_document_embeddings(request) + + except TooManyRequests as e: + raise e + + except Exception as e: + + print(f"Exception: {e}") + raise e + + @staticmethod + def add_args(parser): + + FlowProcessor.add_args(parser) + diff --git a/trustgraph-base/trustgraph/base/graph_embeddings_store_service.py b/trustgraph-base/trustgraph/base/graph_embeddings_store_service.py new file mode 100755 index 00000000..1f31a629 --- /dev/null +++ b/trustgraph-base/trustgraph/base/graph_embeddings_store_service.py @@ -0,0 +1,49 @@ + +""" +Graph embeddings store base class +""" + +from .. schema import GraphEmbeddings +from .. base import FlowProcessor, ConsumerSpec + +default_ident = "graph-embeddings-write" + +class GraphEmbeddingsStoreService(FlowProcessor): + + def __init__(self, **params): + + id = params.get("id") + + super(GraphEmbeddingsStoreService, self).__init__( + **params | { "id": id } + ) + + self.register_specification( + ConsumerSpec( + name = "input", + schema = GraphEmbeddings, + handler = self.on_message + ) + ) + + async def on_message(self, msg, consumer, flow): + + try: + + request = msg.value() + + await self.store_graph_embeddings(request) + + except TooManyRequests as e: + raise e + + except Exception as e: + + print(f"Exception: {e}") + raise e + + @staticmethod + def add_args(parser): + + FlowProcessor.add_args(parser) + diff --git a/trustgraph-base/trustgraph/base/triples_store_service.py b/trustgraph-base/trustgraph/base/triples_store_service.py new file mode 100755 index 00000000..74f95f57 --- /dev/null +++ b/trustgraph-base/trustgraph/base/triples_store_service.py @@ -0,0 +1,47 @@ + +""" +Triples store base class +""" + +from .. schema import Triples +from .. base import FlowProcessor, ConsumerSpec + +default_ident = "triples-write" + +class TriplesStoreService(FlowProcessor): + + def __init__(self, **params): + + id = params.get("id") + + super(TriplesStoreService, self).__init__(**params | { "id": id }) + + self.register_specification( + ConsumerSpec( + name = "input", + schema = Triples, + handler = self.on_message + ) + ) + + async def on_message(self, msg, consumer, flow): + + try: + + request = msg.value() + + await self.store_triples(request) + + except TooManyRequests as e: + raise e + + except Exception as e: + + print(f"Exception: {e}") + raise e + + @staticmethod + def add_args(parser): + + FlowProcessor.add_args(parser) + diff --git a/trustgraph-flow/trustgraph/storage/triples/cassandra/write.py b/trustgraph-flow/trustgraph/storage/triples/cassandra/write.py index f452804b..10a78409 100755 --- a/trustgraph-flow/trustgraph/storage/triples/cassandra/write.py +++ b/trustgraph-flow/trustgraph/storage/triples/cassandra/write.py @@ -10,35 +10,26 @@ import argparse import time from .... direct.cassandra import TrustGraph -from .... schema import Triples -from .... schema import triples_store_queue -from .... log_level import LogLevel -from .... base import Consumer +from .... base import TriplesStoreService -module = "triples-write" +default_ident = "triples-write" -default_input_queue = triples_store_queue -default_subscriber = module default_graph_host='localhost' -class Processor(Consumer): +class Processor(TriplesStoreService): def __init__(self, **params): - input_queue = params.get("input_queue", default_input_queue) - subscriber = params.get("subscriber", default_subscriber) + id = params.get("id", default_ident) + graph_host = params.get("graph_host", default_graph_host) graph_username = params.get("graph_username", None) graph_password = params.get("graph_password", None) super(Processor, self).__init__( **params | { - "input_queue": input_queue, - "subscriber": subscriber, - "input_schema": Triples, "graph_host": graph_host, - "graph_username": graph_username, - "graph_password": graph_password, + "graph_username": graph_username } ) @@ -47,11 +38,9 @@ class Processor(Consumer): self.password = graph_password self.table = None - async def handle(self, msg): + async def store_triples(self, message): - v = msg.value() - - table = (v.metadata.user, v.metadata.collection) + table = (message.metadata.user, message.metadata.collection) if self.table is None or self.table != table: @@ -86,9 +75,7 @@ class Processor(Consumer): @staticmethod def add_args(parser): - Consumer.add_args( - parser, default_input_queue, default_subscriber, - ) + TriplesStoreService.add_args(parser) parser.add_argument( '-g', '--graph-host', @@ -110,5 +97,5 @@ class Processor(Consumer): def run(): - Processor.launch(module, __doc__) + Processor.launch(default_ident, __doc__)