From 6eb16473c555569a4abc6609c10ea7ebb1a613d1 Mon Sep 17 00:00:00 2001 From: Cyber MacGeddon Date: Mon, 21 Apr 2025 12:36:34 +0100 Subject: [PATCH] Linkage is complete --- trustgraph-base/trustgraph/base/__init__.py | 2 + .../base/graph_embeddings_client.py | 43 ++++ .../trustgraph/base/triples_client.py | 53 +++++ .../retrieval/graph_rag/graph_rag.py | 158 ++++-------- .../trustgraph/retrieval/graph_rag/rag.py | 224 +++++++----------- 5 files changed, 221 insertions(+), 259 deletions(-) create mode 100644 trustgraph-base/trustgraph/base/graph_embeddings_client.py create mode 100644 trustgraph-base/trustgraph/base/triples_client.py diff --git a/trustgraph-base/trustgraph/base/__init__.py b/trustgraph-base/trustgraph/base/__init__.py index 6be488bf..23a099a2 100644 --- a/trustgraph-base/trustgraph/base/__init__.py +++ b/trustgraph-base/trustgraph/base/__init__.py @@ -23,4 +23,6 @@ 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 +from . graph_embeddings_client import GraphEmbeddingsClientSpec +from . triples_client import TriplesClientSpec diff --git a/trustgraph-base/trustgraph/base/graph_embeddings_client.py b/trustgraph-base/trustgraph/base/graph_embeddings_client.py new file mode 100644 index 00000000..67b290fb --- /dev/null +++ b/trustgraph-base/trustgraph/base/graph_embeddings_client.py @@ -0,0 +1,43 @@ + +from . request_response_spec import RequestResponse, RequestResponseSpec +from .. schema import GraphEmbeddingsRequest, GraphEmbeddingsResponse +from .. knowledge import Uri, Literal + +def to_value(x): + if x.e: return Uri(x.v) + return Literal(x.v) + +class GraphEmbeddingsClient(RequestResponse): + async def query(self, vectors, limit=20, user="trustgraph", + collection="default", timeout=30): + + resp = await self.request( + GraphEmbeddingsRequest( + vectors = vectors, + limit = limit, + user = user, + collection = collection + ), + timeout=timeout + ) + + if resp.error: + raise RuntimeError(resp.error.message) + + return [ + to_value(v) + for v in resp.entities + ] + +class GraphEmbeddingsClientSpec(RequestResponseSpec): + def __init__( + self, request_name, response_name, + ): + super(GraphEmbeddingsClientSpec, self).__init__( + request_name = request_name, + request_schema = GraphEmbeddingsRequest, + response_name = response_name, + response_schema = GraphEmbeddingsResponse, + impl = GraphEmbeddingsClient, + ) + diff --git a/trustgraph-base/trustgraph/base/triples_client.py b/trustgraph-base/trustgraph/base/triples_client.py new file mode 100644 index 00000000..17e53ca1 --- /dev/null +++ b/trustgraph-base/trustgraph/base/triples_client.py @@ -0,0 +1,53 @@ + +from . request_response_spec import RequestResponse, RequestResponseSpec +from .. schema import TriplesQueryRequest, TriplesQueryResponse, Value +from .. knowledge import Uri, Literal + +def to_value(x): + if x.e: return Uri(x.v) + return Literal(x.v) + +def from_value(x): + if x is None: return None + if isinstance(x, Uri): + return Value(value=str(x), is_uri=True) + else: + return Value(value=str(x), is_uri=False) + +class TriplesClient(RequestResponse): + async def query(self, s=None, p=None, o=None, limit=20, + user="trustgraph", collection="default", + timeout=30): + + resp = await self.request( + TriplesQueryRequest( + s = from_value(s), + p = from_value(p), + o = from_value(o), + limit = limit, + user = user, + collection = collection, + ), + timeout=timeout + ) + + if resp.error: + raise RuntimeError(resp.error.message) + + return [ + to_value(v) + for v in resp.triples + ] + +class TriplesClientSpec(RequestResponseSpec): + def __init__( + self, request_name, response_name, + ): + super(TriplesClientSpec, self).__init__( + request_name = request_name, + request_schema = TriplesQueryRequest, + response_name = response_name, + response_schema = TriplesQueryResponse, + impl = TriplesClient, + ) + diff --git a/trustgraph-flow/trustgraph/retrieval/graph_rag/graph_rag.py b/trustgraph-flow/trustgraph/retrieval/graph_rag/graph_rag.py index 6a4e11c5..d115a51e 100644 --- a/trustgraph-flow/trustgraph/retrieval/graph_rag/graph_rag.py +++ b/trustgraph-flow/trustgraph/retrieval/graph_rag/graph_rag.py @@ -1,19 +1,5 @@ -from . clients.graph_embeddings_client import GraphEmbeddingsClient -from . clients.triples_query_client import TriplesQueryClient -from . clients.embeddings_client import EmbeddingsClient -from . clients.prompt_client import PromptClient - -from . schema import GraphEmbeddingsRequest, GraphEmbeddingsResponse -from . schema import TriplesQueryRequest, TriplesQueryResponse -from . schema import prompt_request_queue -from . schema import prompt_response_queue -from . schema import embeddings_request_queue -from . schema import embeddings_response_queue -from . schema import graph_embeddings_request_queue -from . schema import graph_embeddings_response_queue -from . schema import triples_request_queue -from . schema import triples_response_queue +import asyncio LABEL="http://www.w3.org/2000/01/rdf-schema#label" DEFINITION="http://www.w3.org/2004/02/skos/core#definition" @@ -34,26 +20,26 @@ class Query: self.max_subgraph_size = max_subgraph_size self.max_path_length = max_path_length - def get_vector(self, query): + async def get_vector(self, query): if self.verbose: print("Compute embeddings...", flush=True) - qembeds = self.rag.embeddings.request(query) + qembeds = await self.rag.embeddings.request(query) if self.verbose: print("Done.", flush=True) return qembeds - def get_entities(self, query): + async def get_entities(self, query): - vectors = self.get_vector(query) + vectors = await self.get_vector(query) if self.verbose: print("Get entities...", flush=True) - entities = self.rag.ge_client.request( + entities = await self.rag.ge_client.request( user=self.user, collection=self.collection, vectors=vectors, limit=self.entity_limit, ) @@ -70,12 +56,12 @@ class Query: return entities - def maybe_label(self, e): + async def maybe_label(self, e): if e in self.rag.label_cache: return self.rag.label_cache[e] - res = self.rag.triples_client.request( + res = await self.rag.triples_client.request( user=self.user, collection=self.collection, s=e, p=LABEL, o=None, limit=1, ) @@ -87,7 +73,7 @@ class Query: self.rag.label_cache[e] = res[0].o.value return self.rag.label_cache[e] - def follow_edges(self, ent, subgraph, path_length): + async def follow_edges(self, ent, subgraph, path_length): # Not needed? if path_length <= 0: @@ -97,7 +83,7 @@ class Query: if len(subgraph) >= self.max_subgraph_size: return - res = self.rag.triples_client.request( + res = await self.rag.triples_client.request( user=self.user, collection=self.collection, s=ent, p=None, o=None, limit=self.triple_limit @@ -108,9 +94,9 @@ class Query: (triple.s.value, triple.p.value, triple.o.value) ) if path_length > 1: - self.follow_edges(triple.o.value, subgraph, path_length-1) + await self.follow_edges(triple.o.value, subgraph, path_length-1) - res = self.rag.triples_client.request( + res = await self.rag.triples_client.request( user=self.user, collection=self.collection, s=None, p=ent, o=None, limit=self.triple_limit @@ -121,7 +107,7 @@ class Query: (triple.s.value, triple.p.value, triple.o.value) ) - res = self.rag.triples_client.request( + res = await self.rag.triples_client.request( user=self.user, collection=self.collection, s=None, p=None, o=ent, limit=self.triple_limit, @@ -132,11 +118,11 @@ class Query: (triple.s.value, triple.p.value, triple.o.value) ) if path_length > 1: - self.follow_edges(triple.s.value, subgraph, path_length-1) + await self.follow_edges(triple.s.value, subgraph, path_length-1) - def get_subgraph(self, query): + async def get_subgraph(self, query): - entities = self.get_entities(query) + entities = await self.get_entities(query) if self.verbose: print("Get subgraph...", flush=True) @@ -144,15 +130,15 @@ class Query: subgraph = set() for ent in entities: - self.follow_edges(ent, subgraph, self.max_path_length) + await self.follow_edges(ent, subgraph, self.max_path_length) subgraph = list(subgraph) return subgraph - def get_labelgraph(self, query): + async def get_labelgraph(self, query): - subgraph = self.get_subgraph(query) + subgraph = await self.get_subgraph(query) sg2 = [] @@ -161,9 +147,9 @@ class Query: if edge[1] == LABEL: continue - s = self.maybe_label(edge[0]) - p = self.maybe_label(edge[1]) - o = self.maybe_label(edge[2]) + s = await self.maybe_label(edge[0]) + p = await self.maybe_label(edge[1]) + o = await self.maybe_label(edge[2]) sg2.append((s, p, o)) @@ -182,111 +168,47 @@ class Query: class GraphRag: def __init__( - self, - pulsar_host="pulsar://pulsar:6650", - pulsar_api_key=None, - pr_request_queue=None, - pr_response_queue=None, - emb_request_queue=None, - emb_response_queue=None, - ge_request_queue=None, - ge_response_queue=None, - tpl_request_queue=None, - tpl_response_queue=None, - verbose=False, - module="test", + self, prompt_client, embeddings_client, graph_embeddings_client, + triples_client, verbose=False, ): - self.verbose=verbose + self.verbose = verbose - if pr_request_queue is None: - pr_request_queue = prompt_request_queue - - if pr_response_queue is None: - pr_response_queue = prompt_response_queue - - if emb_request_queue is None: - emb_request_queue = embeddings_request_queue - - if emb_response_queue is None: - emb_response_queue = embeddings_response_queue - - if ge_request_queue is None: - ge_request_queue = graph_embeddings_request_queue - - if ge_response_queue is None: - ge_response_queue = graph_embeddings_response_queue - - if tpl_request_queue is None: - tpl_request_queue = triples_request_queue - - if tpl_response_queue is None: - tpl_response_queue = triples_response_queue - - if self.verbose: - print("Initialising...", flush=True) - - self.ge_client = GraphEmbeddingsClient( - pulsar_host=pulsar_host, - pulsar_api_key=pulsar_api_key, - subscriber=module + "-ge", - input_queue=ge_request_queue, - output_queue=ge_response_queue, - ) - - self.triples_client = TriplesQueryClient( - pulsar_host=pulsar_host, - pulsar_api_key=pulsar_api_key, - subscriber=module + "-tpl", - input_queue=tpl_request_queue, - output_queue=tpl_response_queue - ) - - self.embeddings = EmbeddingsClient( - pulsar_host=pulsar_host, - pulsar_api_key=pulsar_api_key, - input_queue=emb_request_queue, - output_queue=emb_response_queue, - subscriber=module + "-emb", - ) + self.prompt_client = prompt_client + self.embeddings_client = embeddings_client + self.graph_embeddings_client = graph_embeddings_client + self.triples_client = triples_client self.label_cache = {} - self.prompt = PromptClient( - pulsar_host=pulsar_host, - pulsar_api_key=pulsar_api_key, - input_queue=pr_request_queue, - output_queue=pr_response_queue, - subscriber=module + "-prompt", - ) - if self.verbose: print("Initialised", flush=True) - def query( - self, query, user="trustgraph", collection="default", - entity_limit=50, triple_limit=30, max_subgraph_size=1000, - max_path_length=2, + async def query( + self, query, user = "trustgraph", collection = "default", + entity_limit = 50, triple_limit = 30, max_subgraph_size = 1000, + max_path_length = 2, ): if self.verbose: print("Construct prompt...", flush=True) q = Query( - rag=self, user=user, collection=collection, verbose=self.verbose, - entity_limit=entity_limit, triple_limit=triple_limit, - max_subgraph_size=max_subgraph_size, - max_path_length=max_path_length, + rag = self, user = user, collection = collection, + verbose = self.verbose, entity_limit = entity_limit, + triple_limit = triple_limit, + max_subgraph_size = max_subgraph_size, + max_path_length = max_path_length, ) - kg = q.get_labelgraph(query) + kg = await q.get_labelgraph(query) if self.verbose: print("Invoke LLM...", flush=True) print(kg) print(query) - resp = self.prompt.request_kg_prompt(query, kg) + resp = await self.prompt.request_kg_prompt(query, kg) if self.verbose: print("Done", flush=True) diff --git a/trustgraph-flow/trustgraph/retrieval/graph_rag/rag.py b/trustgraph-flow/trustgraph/retrieval/graph_rag/rag.py index 6c965695..99143118 100755 --- a/trustgraph-flow/trustgraph/retrieval/graph_rag/rag.py +++ b/trustgraph-flow/trustgraph/retrieval/graph_rag/rag.py @@ -5,57 +5,18 @@ Input is query, output is response. """ from ... schema import GraphRagQuery, GraphRagResponse, Error -from ... schema import graph_rag_request_queue, graph_rag_response_queue -from ... schema import prompt_request_queue -from ... schema import prompt_response_queue -from ... schema import embeddings_request_queue -from ... schema import embeddings_response_queue -from ... schema import graph_embeddings_request_queue -from ... schema import graph_embeddings_response_queue -from ... schema import triples_request_queue -from ... schema import triples_response_queue -from ... log_level import LogLevel -from ... graph_rag import GraphRag -from ... base import ConsumerProducer +from . graph_rag import GraphRag +from ... base import FlowProcessor, ConsumerSpec, ProducerSpec +from ... base import PromptClientSpec, EmbeddingsClientSpec +from ... base import GraphEmbeddingsClientSpec, TriplesClientSpec -module = "graph-rag" +default_ident = "graph-rag" -default_input_queue = graph_rag_request_queue -default_output_queue = graph_rag_response_queue -default_subscriber = module - -class Processor(ConsumerProducer): +class Processor(FlowProcessor): 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) - - pr_request_queue = params.get( - "prompt_request_queue", prompt_request_queue - ) - pr_response_queue = params.get( - "prompt_response_queue", prompt_response_queue - ) - emb_request_queue = params.get( - "embeddings_request_queue", embeddings_request_queue - ) - emb_response_queue = params.get( - "embeddings_response_queue", embeddings_response_queue - ) - ge_request_queue = params.get( - "graph_embeddings_request_queue", graph_embeddings_request_queue - ) - ge_response_queue = params.get( - "graph_embeddings_response_queue", graph_embeddings_response_queue - ) - tpl_request_queue = params.get( - "triples_request_queue", triples_request_queue - ) - tpl_response_queue = params.get( - "triples_response_queue", triples_response_queue - ) + id = params.get("id", default_ident) entity_limit = params.get("entity_limit", 50) triple_limit = params.get("triple_limit", 30) @@ -64,49 +25,74 @@ class Processor(ConsumerProducer): super(Processor, self).__init__( **params | { - "input_queue": input_queue, - "output_queue": output_queue, - "subscriber": subscriber, - "input_schema": GraphRagQuery, - "output_schema": GraphRagResponse, + "id": id, "entity_limit": entity_limit, "triple_limit": triple_limit, "max_subgraph_size": max_subgraph_size, - "prompt_request_queue": pr_request_queue, - "prompt_response_queue": pr_response_queue, - "embeddings_request_queue": emb_request_queue, - "embeddings_response_queue": emb_response_queue, - "graph_embeddings_request_queue": ge_request_queue, - "graph_embeddings_response_queue": ge_response_queue, - "triples_request_queue": triples_request_queue, - "triples_response_queue": triples_response_queue, + "max_path_length": max_path_length, } ) - self.rag = GraphRag( - pulsar_host=self.pulsar_host, - pulsar_api_key=self.pulsar_api_key, - pr_request_queue=pr_request_queue, - pr_response_queue=pr_response_queue, - emb_request_queue=emb_request_queue, - emb_response_queue=emb_response_queue, - ge_request_queue=ge_request_queue, - ge_response_queue=ge_response_queue, - tpl_request_queue=triples_request_queue, - tpl_response_queue=triples_response_queue, - verbose=True, - module=module, - ) - self.default_entity_limit = entity_limit self.default_triple_limit = triple_limit self.default_max_subgraph_size = max_subgraph_size self.default_max_path_length = max_path_length - async def handle(self, msg): + self.register_specification( + ConsumerSpec( + name = "request", + schema = GraphRagQuery, + handler = self.on_request, + ) + ) + + self.register_specification( + EmbeddingsClientSpec( + request_name = "embeddings-request", + response_name = "embeddings-response", + ) + ) + + self.register_specification( + GraphEmbeddingsClientSpec( + request_name = "graph-embeddings-request", + response_name = "graph-embeddings-response", + ) + ) + + self.register_specification( + TriplesClientSpec( + request_name = "triples-request", + response_name = "triples-response", + ) + ) + + self.register_specification( + PromptClientSpec( + request_name = "prompt-request", + response_name = "prompt-response", + ) + ) + + self.register_specification( + ProducerSpec( + name = "response", + schema = GraphRagResponse, + ) + ) + + async def on_request(self, msg, consumer, flow): try: + self.rag = GraphRag( + embeddings_client = flow("embeddings-request"), + graph_embeddings_client = flow("graph-embeddings-request"), + triples_client = flow("triples-request"), + prompt_client = flow("prompt-request"), + verbose=True, + ) + v = msg.value() # Sender-produced ID @@ -134,16 +120,20 @@ class Processor(ConsumerProducer): else: max_path_length = self.default_max_path_length - response = self.rag.query( - query=v.query, user=v.user, collection=v.collection, - entity_limit=entity_limit, triple_limit=triple_limit, - max_subgraph_size=max_subgraph_size, - max_path_length=max_path_length, + response = await self.rag.query( + query = v.query, user = v.user, collection = v.collection, + entity_limit = entity_limit, triple_limit = triple_limit, + max_subgraph_size = max_subgraph_size, + max_path_length = max_path_length, ) - print("Send response...", flush=True) - r = GraphRagResponse(response=response, error=None) - await self.send(r, properties={"id": id}) + await flow("response").send( + GraphRagResponse( + response = response, + error = None + ), + properties = {"id": id} + ) print("Done.", flush=True) @@ -153,12 +143,15 @@ class Processor(ConsumerProducer): print("Send error response...", flush=True) - r = GraphRagResponse( - error=Error( - type = "llm-error", - message = str(e), + await flow("response").send( + GraphRagResponse( + response = None, + error = Error( + type = "graph-rag-error", + message = str(e), + ), ), - response=None, + properties = {"id": id} ) await self.send(r, properties={"id": id}) @@ -168,10 +161,7 @@ class Processor(ConsumerProducer): @staticmethod def add_args(parser): - ConsumerProducer.add_args( - parser, default_input_queue, default_subscriber, - default_output_queue, - ) + FlowProcessor.add_args(parser) parser.add_argument( '-e', '--entity-limit', @@ -201,55 +191,7 @@ class Processor(ConsumerProducer): help=f'Default max path length (default: 2)' ) - parser.add_argument( - '--prompt-request-queue', - default=prompt_request_queue, - help=f'Prompt request queue (default: {prompt_request_queue})', - ) - - parser.add_argument( - '--prompt-response-queue', - default=prompt_response_queue, - help=f'Prompt response queue (default: {prompt_response_queue})', - ) - - parser.add_argument( - '--embeddings-request-queue', - default=embeddings_request_queue, - help=f'Embeddings request queue (default: {embeddings_request_queue})', - ) - - parser.add_argument( - '--embeddings-response-queue', - default=embeddings_response_queue, - help=f'Embeddings response queue (default: {embeddings_response_queue})', - ) - - parser.add_argument( - '--graph-embeddings-request-queue', - default=graph_embeddings_request_queue, - help=f'Graph embeddings request queue (default: {graph_embeddings_request_queue})', - ) - - parser.add_argument( - '--graph-embeddings-response-queue', - default=graph_embeddings_response_queue, - help=f'Graph embeddings response queue (default: {graph_embeddings_response_queue})', - ) - - parser.add_argument( - '--triples-request-queue', - default=triples_request_queue, - help=f'Triples request queue (default: {triples_request_queue})', - ) - - parser.add_argument( - '--triples-response-queue', - default=triples_response_queue, - help=f'Triples response queue (default: {triples_response_queue})', - ) - def run(): - Processor.launch(module, __doc__) + Processor.launch(default_ident, __doc__)