Fix bugs, graph-rag working

This commit is contained in:
Cyber MacGeddon 2025-03-12 22:32:55 +00:00
parent 8830e99e57
commit 5ca910e089
5 changed files with 24 additions and 21 deletions

View file

@ -104,7 +104,7 @@ class Api:
def graph_rag( def graph_rag(
self, question, user="trustgraph", collection="default", self, question, user="trustgraph", collection="default",
entity_limit=50, triple_limit=30, subgraph_limit=1000, entity_limit=50, triple_limit=30, max_subgraph_size=1000,
): ):
# The input consists of a question # The input consists of a question
@ -114,7 +114,7 @@ class Api:
"collection": collection, "collection": collection,
"entity-limit": entity_limit, "entity-limit": entity_limit,
"triple-limit": triple_limit, "triple-limit": triple_limit,
"max-subgraph-limit": subgraph_limit, "max-subgraph-size": max_subgraph_size,
} }
url = f"{self.url}graph-rag" url = f"{self.url}graph-rag"

View file

@ -11,10 +11,13 @@ from trustgraph.api import Api
default_url = os.getenv("TRUSTGRAPH_URL", 'http://localhost:8088/') default_url = os.getenv("TRUSTGRAPH_URL", 'http://localhost:8088/')
default_user = 'trustgraph' default_user = 'trustgraph'
default_collection = 'default' default_collection = 'default'
default_entity_limit = 50
default_triple_limit = 30
default_max_subgraph_size = 1000
def question( def question(
url, question, user, collection, entity_limit, triple_limit, url, question, user, collection, entity_limit, triple_limit,
subgraph_limit max_subgraph_size
): ):
rag = Api(url) rag = Api(url)
@ -22,7 +25,7 @@ def question(
resp = rag.graph_rag( resp = rag.graph_rag(
question=question, user=user, collection=collection, question=question, user=user, collection=collection,
entity_limit=entity_limit, triple_limit=triple_limit, entity_limit=entity_limit, triple_limit=triple_limit,
subgraph_limit=subgraph_limit max_subgraph_size=max_subgraph_size
) )
print(resp) print(resp)
@ -71,9 +74,9 @@ def main():
) )
parser.add_argument( parser.add_argument(
'-s', '--subgraph-limit', '-s', '--max-subgraph-size',
default=default_subgraph_limit, default=default_max_subgraph_size,
help=f'Subgraph limit (default: {subgraph_limit})' help=f'Max subgraph size (default: {default_max_subgraph_size})'
) )
args = parser.parse_args() args = parser.parse_args()
@ -87,7 +90,7 @@ def main():
collection=args.collection, collection=args.collection,
entity_limit=args.entity_limit, entity_limit=args.entity_limit,
triple_limit=args.triple_limit, triple_limit=args.triple_limit,
subgraph_limit=args.subgraph_limit, max_subgraph_size=args.max_subgraph_size,
) )
except Exception as e: except Exception as e:

View file

@ -23,9 +23,9 @@ class GraphRagRequestor(ServiceRequestor):
query=body["query"], query=body["query"],
user=body.get("user", "trustgraph"), user=body.get("user", "trustgraph"),
collection=body.get("collection", "default"), collection=body.get("collection", "default"),
entity_limit=body.get("entity-limit", 50), entity_limit=int(body.get("entity-limit", 50)),
triple_limit=body.get("triple-limit", 30), triple_limit=int(body.get("triple-limit", 30)),
doc_limit=body.get("doc-limit", 1000), max_subgraph_size=int(body.get("max-subgraph-size", 1000)),
) )
def from_response(self, message): def from_response(self, message):

View file

@ -22,14 +22,14 @@ class Query:
def __init__( def __init__(
self, rag, user, collection, verbose, self, rag, user, collection, verbose,
entity_limit=50, query_limit=30, max_subgraph_size=1000, entity_limit=50, triple_limit=30, max_subgraph_size=1000,
): ):
self.rag = rag self.rag = rag
self.user = user self.user = user
self.collection = collection self.collection = collection
self.verbose = verbose self.verbose = verbose
self.entity_limit = entity_limit self.entity_limit = entity_limit
self.query_limit = query_limit self.triple_limit = triple_limit
self.max_subgraph_size = max_subgraph_size self.max_subgraph_size = max_subgraph_size
def get_vector(self, query): def get_vector(self, query):
@ -99,7 +99,7 @@ class Query:
res = self.rag.triples_client.request( res = self.rag.triples_client.request(
user=self.user, collection=self.collection, user=self.user, collection=self.collection,
s=e, p=None, o=None, s=e, p=None, o=None,
limit=self.query_limit limit=self.triple_limit
) )
for triple in res: for triple in res:
@ -110,7 +110,7 @@ class Query:
res = self.rag.triples_client.request( res = self.rag.triples_client.request(
user=self.user, collection=self.collection, user=self.user, collection=self.collection,
s=None, p=e, o=None, s=None, p=e, o=None,
limit=self.query_limit limit=self.triple_limit
) )
for triple in res: for triple in res:
@ -121,7 +121,7 @@ class Query:
res = self.rag.triples_client.request( res = self.rag.triples_client.request(
user=self.user, collection=self.collection, user=self.user, collection=self.collection,
s=None, p=None, o=e, s=None, p=None, o=e,
limit=self.query_limit, limit=self.triple_limit,
) )
for triple in res: for triple in res:
@ -255,8 +255,8 @@ class GraphRag:
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, verbose=self.verbose,
entity_limit=entity_limit, query_limit=query_limit, entity_limit=entity_limit, triple_limit=triple_limit,
max_subgraph_size=max_subgraph_size, max_subgraph_size=max_subgraph_size,
) )

View file

@ -115,17 +115,17 @@ class Processor(ConsumerProducer):
if v.entity_limit: if v.entity_limit:
entity_limit = v.entity_limit entity_limit = v.entity_limit
else: else:
entity_limit = self.entity_limit entity_limit = self.default_entity_limit
if v.triple_limit: if v.triple_limit:
triple_limit = v.triple_limit triple_limit = v.triple_limit
else: else:
triple_limit = self.triple_limit triple_limit = self.default_triple_limit
if v.max_subgraph_size: if v.max_subgraph_size:
max_subgraph_size = v.max_subgraph_size max_subgraph_size = v.max_subgraph_size
else: else:
max_subgraph_size = self.max_subgraph_size max_subgraph_size = self.default_max_subgraph_size
response = self.rag.query( response = self.rag.query(
query=v.query, user=v.user, collection=v.collection, query=v.query, user=v.user, collection=v.collection,