Linkage is complete

This commit is contained in:
Cyber MacGeddon 2025-04-21 12:36:34 +01:00
parent aa47ae9970
commit 6eb16473c5
5 changed files with 221 additions and 259 deletions

View file

@ -23,4 +23,6 @@ from . document_embeddings_store_service import DocumentEmbeddingsStoreService
from . triples_query_service import TriplesQueryService from . triples_query_service import TriplesQueryService
from . graph_embeddings_query_service import GraphEmbeddingsQueryService from . graph_embeddings_query_service import GraphEmbeddingsQueryService
from . document_embeddings_query_service import DocumentEmbeddingsQueryService from . document_embeddings_query_service import DocumentEmbeddingsQueryService
from . graph_embeddings_client import GraphEmbeddingsClientSpec
from . triples_client import TriplesClientSpec

View file

@ -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,
)

View file

@ -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,
)

View file

@ -1,19 +1,5 @@
from . clients.graph_embeddings_client import GraphEmbeddingsClient import asyncio
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
LABEL="http://www.w3.org/2000/01/rdf-schema#label" LABEL="http://www.w3.org/2000/01/rdf-schema#label"
DEFINITION="http://www.w3.org/2004/02/skos/core#definition" 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_subgraph_size = max_subgraph_size
self.max_path_length = max_path_length self.max_path_length = max_path_length
def get_vector(self, query): async def get_vector(self, query):
if self.verbose: if self.verbose:
print("Compute embeddings...", flush=True) print("Compute embeddings...", flush=True)
qembeds = self.rag.embeddings.request(query) qembeds = await self.rag.embeddings.request(query)
if self.verbose: if self.verbose:
print("Done.", flush=True) print("Done.", flush=True)
return qembeds 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: if self.verbose:
print("Get entities...", flush=True) print("Get entities...", flush=True)
entities = self.rag.ge_client.request( entities = await self.rag.ge_client.request(
user=self.user, collection=self.collection, user=self.user, collection=self.collection,
vectors=vectors, limit=self.entity_limit, vectors=vectors, limit=self.entity_limit,
) )
@ -70,12 +56,12 @@ class Query:
return entities return entities
def maybe_label(self, e): async def maybe_label(self, e):
if e in self.rag.label_cache: if e in self.rag.label_cache:
return self.rag.label_cache[e] 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, user=self.user, collection=self.collection,
s=e, p=LABEL, o=None, limit=1, s=e, p=LABEL, o=None, limit=1,
) )
@ -87,7 +73,7 @@ class Query:
self.rag.label_cache[e] = res[0].o.value self.rag.label_cache[e] = res[0].o.value
return self.rag.label_cache[e] 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? # Not needed?
if path_length <= 0: if path_length <= 0:
@ -97,7 +83,7 @@ class Query:
if len(subgraph) >= self.max_subgraph_size: if len(subgraph) >= self.max_subgraph_size:
return return
res = self.rag.triples_client.request( res = await self.rag.triples_client.request(
user=self.user, collection=self.collection, user=self.user, collection=self.collection,
s=ent, p=None, o=None, s=ent, p=None, o=None,
limit=self.triple_limit limit=self.triple_limit
@ -108,9 +94,9 @@ class Query:
(triple.s.value, triple.p.value, triple.o.value) (triple.s.value, triple.p.value, triple.o.value)
) )
if path_length > 1: 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, user=self.user, collection=self.collection,
s=None, p=ent, o=None, s=None, p=ent, o=None,
limit=self.triple_limit limit=self.triple_limit
@ -121,7 +107,7 @@ class Query:
(triple.s.value, triple.p.value, triple.o.value) (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, user=self.user, collection=self.collection,
s=None, p=None, o=ent, s=None, p=None, o=ent,
limit=self.triple_limit, limit=self.triple_limit,
@ -132,11 +118,11 @@ class Query:
(triple.s.value, triple.p.value, triple.o.value) (triple.s.value, triple.p.value, triple.o.value)
) )
if path_length > 1: 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: if self.verbose:
print("Get subgraph...", flush=True) print("Get subgraph...", flush=True)
@ -144,15 +130,15 @@ class Query:
subgraph = set() subgraph = set()
for ent in entities: 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) subgraph = list(subgraph)
return 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 = [] sg2 = []
@ -161,9 +147,9 @@ class Query:
if edge[1] == LABEL: if edge[1] == LABEL:
continue continue
s = self.maybe_label(edge[0]) s = await self.maybe_label(edge[0])
p = self.maybe_label(edge[1]) p = await self.maybe_label(edge[1])
o = self.maybe_label(edge[2]) o = await self.maybe_label(edge[2])
sg2.append((s, p, o)) sg2.append((s, p, o))
@ -182,111 +168,47 @@ class Query:
class GraphRag: class GraphRag:
def __init__( def __init__(
self, self, prompt_client, embeddings_client, graph_embeddings_client,
pulsar_host="pulsar://pulsar:6650", triples_client, verbose=False,
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.verbose=verbose self.verbose = verbose
if pr_request_queue is None: self.prompt_client = prompt_client
pr_request_queue = prompt_request_queue self.embeddings_client = embeddings_client
self.graph_embeddings_client = graph_embeddings_client
if pr_response_queue is None: self.triples_client = triples_client
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.label_cache = {} 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: if self.verbose:
print("Initialised", flush=True) print("Initialised", flush=True)
def query( async def query(
self, query, user="trustgraph", collection="default", self, query, user = "trustgraph", collection = "default",
entity_limit=50, triple_limit=30, max_subgraph_size=1000, entity_limit = 50, triple_limit = 30, max_subgraph_size = 1000,
max_path_length=2, max_path_length = 2,
): ):
if self.verbose: if self.verbose:
print("Construct prompt...", flush=True) print("Construct prompt...", flush=True)
q = Query( q = Query(
rag=self, user=user, collection=collection, verbose=self.verbose, rag = self, user = user, collection = collection,
entity_limit=entity_limit, triple_limit=triple_limit, verbose = self.verbose, entity_limit = entity_limit,
max_subgraph_size=max_subgraph_size, triple_limit = triple_limit,
max_path_length=max_path_length, 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: if self.verbose:
print("Invoke LLM...", flush=True) print("Invoke LLM...", flush=True)
print(kg) print(kg)
print(query) print(query)
resp = self.prompt.request_kg_prompt(query, kg) resp = await self.prompt.request_kg_prompt(query, kg)
if self.verbose: if self.verbose:
print("Done", flush=True) print("Done", flush=True)

View file

@ -5,57 +5,18 @@ Input is query, output is response.
""" """
from ... schema import GraphRagQuery, GraphRagResponse, Error from ... schema import GraphRagQuery, GraphRagResponse, Error
from ... schema import graph_rag_request_queue, graph_rag_response_queue from . graph_rag import GraphRag
from ... schema import prompt_request_queue from ... base import FlowProcessor, ConsumerSpec, ProducerSpec
from ... schema import prompt_response_queue from ... base import PromptClientSpec, EmbeddingsClientSpec
from ... schema import embeddings_request_queue from ... base import GraphEmbeddingsClientSpec, TriplesClientSpec
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
module = "graph-rag" default_ident = "graph-rag"
default_input_queue = graph_rag_request_queue class Processor(FlowProcessor):
default_output_queue = graph_rag_response_queue
default_subscriber = module
class Processor(ConsumerProducer):
def __init__(self, **params): def __init__(self, **params):
input_queue = params.get("input_queue", default_input_queue) id = params.get("id", default_ident)
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
)
entity_limit = params.get("entity_limit", 50) entity_limit = params.get("entity_limit", 50)
triple_limit = params.get("triple_limit", 30) triple_limit = params.get("triple_limit", 30)
@ -64,49 +25,74 @@ class Processor(ConsumerProducer):
super(Processor, self).__init__( super(Processor, self).__init__(
**params | { **params | {
"input_queue": input_queue, "id": id,
"output_queue": output_queue,
"subscriber": subscriber,
"input_schema": GraphRagQuery,
"output_schema": GraphRagResponse,
"entity_limit": entity_limit, "entity_limit": entity_limit,
"triple_limit": triple_limit, "triple_limit": triple_limit,
"max_subgraph_size": max_subgraph_size, "max_subgraph_size": max_subgraph_size,
"prompt_request_queue": pr_request_queue, "max_path_length": max_path_length,
"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,
} }
) )
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_entity_limit = entity_limit
self.default_triple_limit = triple_limit self.default_triple_limit = triple_limit
self.default_max_subgraph_size = max_subgraph_size self.default_max_subgraph_size = max_subgraph_size
self.default_max_path_length = max_path_length 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: 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() v = msg.value()
# Sender-produced ID # Sender-produced ID
@ -134,16 +120,20 @@ class Processor(ConsumerProducer):
else: else:
max_path_length = self.default_max_path_length max_path_length = self.default_max_path_length
response = self.rag.query( response = await self.rag.query(
query=v.query, user=v.user, collection=v.collection, query = v.query, user = v.user, collection = v.collection,
entity_limit=entity_limit, triple_limit=triple_limit, entity_limit = entity_limit, triple_limit = triple_limit,
max_subgraph_size=max_subgraph_size, max_subgraph_size = max_subgraph_size,
max_path_length=max_path_length, max_path_length = max_path_length,
) )
print("Send response...", flush=True) await flow("response").send(
r = GraphRagResponse(response=response, error=None) GraphRagResponse(
await self.send(r, properties={"id": id}) response = response,
error = None
),
properties = {"id": id}
)
print("Done.", flush=True) print("Done.", flush=True)
@ -153,12 +143,15 @@ class Processor(ConsumerProducer):
print("Send error response...", flush=True) print("Send error response...", flush=True)
r = GraphRagResponse( await flow("response").send(
error=Error( GraphRagResponse(
type = "llm-error", response = None,
message = str(e), error = Error(
type = "graph-rag-error",
message = str(e),
),
), ),
response=None, properties = {"id": id}
) )
await self.send(r, properties={"id": id}) await self.send(r, properties={"id": id})
@ -168,10 +161,7 @@ class Processor(ConsumerProducer):
@staticmethod @staticmethod
def add_args(parser): def add_args(parser):
ConsumerProducer.add_args( FlowProcessor.add_args(parser)
parser, default_input_queue, default_subscriber,
default_output_queue,
)
parser.add_argument( parser.add_argument(
'-e', '--entity-limit', '-e', '--entity-limit',
@ -201,55 +191,7 @@ class Processor(ConsumerProducer):
help=f'Default max path length (default: 2)' 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(): def run():
Processor.launch(module, __doc__) Processor.launch(default_ident, __doc__)