mirror of
https://github.com/trustgraph-ai/trustgraph.git
synced 2026-07-22 03:31:02 +02:00
Linkage is complete
This commit is contained in:
parent
aa47ae9970
commit
6eb16473c5
5 changed files with 221 additions and 259 deletions
|
|
@ -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
|
||||||
|
|
||||||
|
|
|
||||||
43
trustgraph-base/trustgraph/base/graph_embeddings_client.py
Normal file
43
trustgraph-base/trustgraph/base/graph_embeddings_client.py
Normal 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,
|
||||||
|
)
|
||||||
|
|
||||||
53
trustgraph-base/trustgraph/base/triples_client.py
Normal file
53
trustgraph-base/trustgraph/base/triples_client.py
Normal 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,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
@ -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)
|
||||||
|
|
|
||||||
|
|
@ -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__)
|
||||||
|
|
||||||
|
|
|
||||||
Loading…
Add table
Add a link
Reference in a new issue