From 4943b87f7a3a134980f567d81a92555a11115e4a Mon Sep 17 00:00:00 2001 From: Cyber MacGeddon Date: Sat, 19 Apr 2025 15:51:26 +0100 Subject: [PATCH] Updated query/store --- trustgraph-base/trustgraph/base/__init__.py | 3 + .../base/document_embeddings_query_service.py | 84 +++++++++++++ .../base/graph_embeddings_query_service.py | 11 +- .../trustgraph/base/triples_query_service.py | 9 +- .../query/doc_embeddings/qdrant/service.py | 55 ++------- .../query/graph_embeddings/qdrant/service.py | 4 + .../query/triples/cassandra/service.py | 114 ++++++------------ .../storage/doc_embeddings/qdrant/write.py | 30 ++--- 8 files changed, 162 insertions(+), 148 deletions(-) create mode 100755 trustgraph-base/trustgraph/base/document_embeddings_query_service.py diff --git a/trustgraph-base/trustgraph/base/__init__.py b/trustgraph-base/trustgraph/base/__init__.py index ddf47f97..6be488bf 100644 --- a/trustgraph-base/trustgraph/base/__init__.py +++ b/trustgraph-base/trustgraph/base/__init__.py @@ -20,4 +20,7 @@ from . prompt_client import PromptClientSpec from . triples_store_service import TriplesStoreService from . graph_embeddings_store_service import GraphEmbeddingsStoreService from . document_embeddings_store_service import DocumentEmbeddingsStoreService +from . triples_query_service import TriplesQueryService +from . graph_embeddings_query_service import GraphEmbeddingsQueryService +from . document_embeddings_query_service import DocumentEmbeddingsQueryService diff --git a/trustgraph-base/trustgraph/base/document_embeddings_query_service.py b/trustgraph-base/trustgraph/base/document_embeddings_query_service.py new file mode 100755 index 00000000..d84def88 --- /dev/null +++ b/trustgraph-base/trustgraph/base/document_embeddings_query_service.py @@ -0,0 +1,84 @@ + +""" +Document embeddings query service. Input is vectors. Output is list of +embeddings. +""" + +from .. schema import DocumentEmbeddingsRequest, DocumentEmbeddingsResponse +from .. schema import Error, Value + +from . flow_processor import FlowProcessor +from . consumer_spec import ConsumerSpec +from . producer_spec import ProducerSpec + +default_ident = "ge-query" + +class DocumentEmbeddingsQueryService(FlowProcessor): + + def __init__(self, **params): + + id = params.get("id") + + super(DocumentEmbeddingsQueryService, self).__init__( + **params | { "id": id } + ) + + self.register_specification( + ConsumerSpec( + name = "request", + schema = DocumentEmbeddingsRequest, + handler = self.on_message + ) + ) + + self.register_specification( + ProducerSpec( + name = "response", + schema = DocumentEmbeddingsResponse, + ) + ) + + async def on_message(self, msg, consumer, flow): + + try: + + request = msg.value() + + # Sender-produced ID + id = msg.properties()["id"] + + print(f"Handling input {id}...", flush=True) + + entities = self.query_document_embeddings(request) + + print("Send response...", flush=True) + r = DocumentEmbeddingsResponse(entities=entities, error=None) + await flow("response").send(r, properties={"id": id}) + + print("Done.", flush=True) + + except Exception as e: + + print(f"Exception: {e}") + + print("Send error response...", flush=True) + + r = DocumentEmbeddingsResponse( + error=Error( + type = "document-embeddings-query-error", + message = str(e), + ), + response=None, + ) + + await flow("response").send(r, properties={"id": id}) + + @staticmethod + def add_args(parser): + + FlowProcessor.add_args(parser) + +def run(): + + Processor.launch(default_ident, __doc__) + diff --git a/trustgraph-base/trustgraph/base/graph_embeddings_query_service.py b/trustgraph-base/trustgraph/base/graph_embeddings_query_service.py index 04efed0c..003ba664 100755 --- a/trustgraph-base/trustgraph/base/graph_embeddings_query_service.py +++ b/trustgraph-base/trustgraph/base/graph_embeddings_query_service.py @@ -4,11 +4,12 @@ Graph embeddings query service. Input is vectors. Output is list of embeddings. """ -from .... schema import GraphEmbeddingsRequest, GraphEmbeddingsResponse -from .... schema import Error, Value +from .. schema import GraphEmbeddingsRequest, GraphEmbeddingsResponse +from .. schema import Error, Value -from .... base import FlowProcessor, GraphEmbeddingsQueryService, ConsumerSpec -from .... base import ProducerSpec +from . flow_processor import FlowProcessor +from . consumer_spec import ConsumerSpec +from . producer_spec import ProducerSpec default_ident = "ge-query" @@ -75,7 +76,7 @@ class GraphEmbeddingsQueryService(FlowProcessor): @staticmethod def add_args(parser): - GraphEmbeddingsQueryService.add_args(parser) + FlowProcessor.add_args(parser) def run(): diff --git a/trustgraph-base/trustgraph/base/triples_query_service.py b/trustgraph-base/trustgraph/base/triples_query_service.py index 0538a7f7..ef90081c 100755 --- a/trustgraph-base/trustgraph/base/triples_query_service.py +++ b/trustgraph-base/trustgraph/base/triples_query_service.py @@ -4,11 +4,12 @@ Triples query service. Input is a (s, p, o) triple, some values may be null. Output is a list of triples. """ -from .... schema import TriplesQueryRequest, TriplesQueryResponse, Error -from .... schema import Value, Triple +from .. schema import TriplesQueryRequest, TriplesQueryResponse, Error +from .. schema import Value, Triple -from .... base import FlowProcessor, TriplesQueryService, ConsumerSpec -from .... base import ProducerSpec +from . flow_processor import FlowProcessor +from . consumer_spec import ConsumerSpec +from . producer_spec import ProducerSpec default_ident = "triples-query" diff --git a/trustgraph-flow/trustgraph/query/doc_embeddings/qdrant/service.py b/trustgraph-flow/trustgraph/query/doc_embeddings/qdrant/service.py index b2d4cdb3..83cfc766 100755 --- a/trustgraph-flow/trustgraph/query/doc_embeddings/qdrant/service.py +++ b/trustgraph-flow/trustgraph/query/doc_embeddings/qdrant/service.py @@ -7,45 +7,32 @@ of chunks from qdrant_client import QdrantClient from qdrant_client.models import PointStruct from qdrant_client.models import Distance, VectorParams -import uuid -from .... schema import DocumentEmbeddingsRequest, DocumentEmbeddingsResponse +from .... schema import DocumentEmbeddingsResponse from .... schema import Error, Value -from .... schema import document_embeddings_request_queue -from .... schema import document_embeddings_response_queue -from .... base import ConsumerProducer +from .... base import DocumentEmbeddingsQueryService -module = "de-query" +default_ident = "de-query" -default_input_queue = document_embeddings_request_queue -default_output_queue = document_embeddings_response_queue -default_subscriber = module default_store_uri = 'http://localhost:6333' -class Processor(ConsumerProducer): +class Processor(DocumentEmbeddingsQueryService): 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) + #optional api key api_key = params.get("api_key", None) super(Processor, self).__init__( **params | { - "input_queue": input_queue, - "output_queue": output_queue, - "subscriber": subscriber, - "input_schema": DocumentEmbeddingsRequest, - "output_schema": DocumentEmbeddingsResponse, "store_uri": store_uri, "api_key": api_key, } ) - self.client = QdrantClient(url=store_uri, api_key=api_key) + self.qdrant = QdrantClient(url=store_uri, api_key=api_key) async def handle(self, msg): @@ -68,7 +55,7 @@ class Processor(ConsumerProducer): str(dim) ) - search_result = self.client.query_points( + search_result = self.qdrant.query_points( collection_name=collection, query=vec, limit=v.limit, @@ -79,37 +66,17 @@ class Processor(ConsumerProducer): ent = r.payload["doc"] chunks.append(ent) - print("Send response...", flush=True) - r = DocumentEmbeddingsResponse(documents=chunks, error=None) - await self.send(r, properties={"id": id}) - - print("Done.", flush=True) + return chunks except Exception as e: print(f"Exception: {e}") - - print("Send error response...", flush=True) - - r = DocumentEmbeddingsResponse( - error=Error( - type = "llm-error", - message = str(e), - ), - documents=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, - ) + DocumentEmbeddingsQueryService.add_args(parser) parser.add_argument( '-t', '--store-uri', @@ -125,5 +92,5 @@ class Processor(ConsumerProducer): def run(): - Processor.launch(module, __doc__) + Processor.launch(default_ident, __doc__) diff --git a/trustgraph-flow/trustgraph/query/graph_embeddings/qdrant/service.py b/trustgraph-flow/trustgraph/query/graph_embeddings/qdrant/service.py index 74a5d182..d818f672 100755 --- a/trustgraph-flow/trustgraph/query/graph_embeddings/qdrant/service.py +++ b/trustgraph-flow/trustgraph/query/graph_embeddings/qdrant/service.py @@ -12,6 +12,8 @@ from .... schema import GraphEmbeddingsResponse from .... schema import Error, Value from .... base import GraphEmbeddingsQueryService +default_ident = "ge-query" + default_store_uri = 'http://localhost:6333' class Processor(GraphEmbeddingsQueryService): @@ -19,6 +21,8 @@ class Processor(GraphEmbeddingsQueryService): def __init__(self, **params): store_uri = params.get("store_uri", default_store_uri) + + #optional api key api_key = params.get("api_key", None) super(Processor, self).__init__( diff --git a/trustgraph-flow/trustgraph/query/triples/cassandra/service.py b/trustgraph-flow/trustgraph/query/triples/cassandra/service.py index b06765e2..48818dad 100755 --- a/trustgraph-flow/trustgraph/query/triples/cassandra/service.py +++ b/trustgraph-flow/trustgraph/query/triples/cassandra/service.py @@ -7,38 +7,24 @@ null. Output is a list of triples. from .... direct.cassandra import TrustGraph from .... schema import TriplesQueryRequest, TriplesQueryResponse, Error from .... schema import Value, Triple -from .... schema import triples_request_queue -from .... schema import triples_response_queue -from .... base import ConsumerProducer +from .... base import TriplesQueryService -module = "triples-query" +default_ident = "triples-query" -default_input_queue = triples_request_queue -default_output_queue = triples_response_queue -default_subscriber = module default_graph_host='localhost' -class Processor(ConsumerProducer): +class Processor(TriplesQueryService): 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) 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, - "output_queue": output_queue, - "subscriber": subscriber, - "input_schema": TriplesQueryRequest, - "output_schema": TriplesQueryResponse, "graph_host": graph_host, "graph_username": graph_username, - "graph_password": graph_password, } ) @@ -53,25 +39,23 @@ class Processor(ConsumerProducer): else: return Value(value=ent, is_uri=False) - async def handle(self, msg): + async def query_triples(self, query): try: - v = msg.value() - - table = (v.user, v.collection) + table = (query.user, query.collection) if table != self.table: if self.username and self.password: self.tg = TrustGraph( hosts=self.graph_host, - keyspace=v.user, table=v.collection, + keyspace=query.user, table=query.collection, username=self.username, password=self.password ) else: self.tg = TrustGraph( hosts=self.graph_host, - keyspace=v.user, table=v.collection, + keyspace=query.user, table=query.collection, ) self.table = table @@ -82,63 +66,63 @@ class Processor(ConsumerProducer): triples = [] - if v.s is not None: - if v.p is not None: - if v.o is not None: + if query.s is not None: + if query.p is not None: + if query.o is not None: resp = self.tg.get_spo( - v.s.value, v.p.value, v.o.value, - limit=v.limit + query.s.value, query.p.value, query.o.value, + limit=query.limit ) - triples.append((v.s.value, v.p.value, v.o.value)) + triples.append((query.s.value, query.p.value, query.o.value)) else: resp = self.tg.get_sp( - v.s.value, v.p.value, - limit=v.limit + query.s.value, query.p.value, + limit=query.limit ) for t in resp: - triples.append((v.s.value, v.p.value, t.o)) + triples.append((query.s.value, query.p.value, t.o)) else: - if v.o is not None: + if query.o is not None: resp = self.tg.get_os( - v.o.value, v.s.value, - limit=v.limit + query.o.value, query.s.value, + limit=query.limit ) for t in resp: - triples.append((v.s.value, t.p, v.o.value)) + triples.append((query.s.value, t.p, query.o.value)) else: resp = self.tg.get_s( - v.s.value, - limit=v.limit + query.s.value, + limit=query.limit ) for t in resp: - triples.append((v.s.value, t.p, t.o)) + triples.append((query.s.value, t.p, t.o)) else: - if v.p is not None: - if v.o is not None: + if query.p is not None: + if query.o is not None: resp = self.tg.get_po( - v.p.value, v.o.value, - limit=v.limit + query.p.value, query.o.value, + limit=query.limit ) for t in resp: - triples.append((t.s, v.p.value, v.o.value)) + triples.append((t.s, query.p.value, query.o.value)) else: resp = self.tg.get_p( - v.p.value, - limit=v.limit + query.p.value, + limit=query.limit ) for t in resp: - triples.append((t.s, v.p.value, t.o)) + triples.append((t.s, query.p.value, t.o)) else: - if v.o is not None: + if query.o is not None: resp = self.tg.get_o( - v.o.value, - limit=v.limit + query.o.value, + limit=query.limit ) for t in resp: - triples.append((t.s, t.p, v.o.value)) + triples.append((t.s, t.p, query.o.value)) else: resp = self.tg.get_all( - limit=v.limit + limit=query.limit ) for t in resp: triples.append((t.s, t.p, t.o)) @@ -152,37 +136,19 @@ class Processor(ConsumerProducer): for t in triples ] - print("Send response...", flush=True) - r = TriplesQueryResponse(triples=triples, error=None) - await self.send(r, properties={"id": id}) + return triples print("Done.", flush=True) except Exception as e: print(f"Exception: {e}") - - print("Send error response...", flush=True) - - r = TriplesQueryResponse( - error=Error( - type = "llm-error", - message = str(e), - ), - response=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, - ) + TriplesQueryService.add_args(parser) parser.add_argument( '-g', '--graph-host', @@ -205,5 +171,5 @@ class Processor(ConsumerProducer): def run(): - Processor.launch(module, __doc__) + Processor.launch(default_ident, __doc__) diff --git a/trustgraph-flow/trustgraph/storage/doc_embeddings/qdrant/write.py b/trustgraph-flow/trustgraph/storage/doc_embeddings/qdrant/write.py index 2213dbe3..e10d1eda 100644 --- a/trustgraph-flow/trustgraph/storage/doc_embeddings/qdrant/write.py +++ b/trustgraph-flow/trustgraph/storage/doc_embeddings/qdrant/write.py @@ -8,31 +8,21 @@ from qdrant_client.models import PointStruct from qdrant_client.models import Distance, VectorParams import uuid -from .... schema import DocumentEmbeddings -from .... schema import document_embeddings_store_queue -from .... log_level import LogLevel -from .... base import Consumer +from .... base import DocumentEmbeddingsStoreService -module = "de-write" +default_ident = "de-write" -default_input_queue = document_embeddings_store_queue -default_subscriber = module default_store_uri = 'http://localhost:6333' -class Processor(Consumer): +class Processor(DocumentEmbeddingsStoreService): 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) api_key = params.get("api_key", None) super(Processor, self).__init__( **params | { - "input_queue": input_queue, - "subscriber": subscriber, - "input_schema": DocumentEmbeddings, "store_uri": store_uri, "api_key": api_key, } @@ -40,13 +30,11 @@ class Processor(Consumer): self.last_collection = None - self.client = QdrantClient(url=store_uri) + self.qdrant = QdrantClient(url=store_uri, api_key=api_key) - async def handle(self, msg): + async def store_document_embeddings(self, message): - v = msg.value() - - for emb in v.chunks: + for emb in message.chunks: chunk = emb.chunk.decode("utf-8") if chunk == "": return @@ -55,7 +43,7 @@ class Processor(Consumer): dim = len(vec) collection = ( - "d_" + v.metadata.user + "_" + v.metadata.collection + "_" + + "d_" + message.metadata.user + "_" + message.metadata.collection + "_" + str(dim) ) @@ -92,7 +80,7 @@ class Processor(Consumer): @staticmethod def add_args(parser): - Consumer.add_args( + DocumentEmbeddingsStoreService.add_args( parser, default_input_queue, default_subscriber, ) @@ -110,5 +98,5 @@ class Processor(Consumer): def run(): - Processor.launch(module, __doc__) + Processor.launch(default_ident, __doc__)