diff --git a/trustgraph-base/trustgraph/api/api.py b/trustgraph-base/trustgraph/api/api.py index 557e7583..68c4b147 100644 --- a/trustgraph-base/trustgraph/api/api.py +++ b/trustgraph-base/trustgraph/api/api.py @@ -104,7 +104,7 @@ class Api: def graph_rag( 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 @@ -114,7 +114,7 @@ class Api: "collection": collection, "entity-limit": entity_limit, "triple-limit": triple_limit, - "max-subgraph-limit": subgraph_limit, + "max-subgraph-size": max_subgraph_size, } url = f"{self.url}graph-rag" diff --git a/trustgraph-cli/scripts/tg-invoke-graph-rag b/trustgraph-cli/scripts/tg-invoke-graph-rag index a573bdfd..dedec5e6 100755 --- a/trustgraph-cli/scripts/tg-invoke-graph-rag +++ b/trustgraph-cli/scripts/tg-invoke-graph-rag @@ -11,10 +11,13 @@ from trustgraph.api import Api default_url = os.getenv("TRUSTGRAPH_URL", 'http://localhost:8088/') default_user = 'trustgraph' default_collection = 'default' +default_entity_limit = 50 +default_triple_limit = 30 +default_max_subgraph_size = 1000 def question( url, question, user, collection, entity_limit, triple_limit, - subgraph_limit + max_subgraph_size ): rag = Api(url) @@ -22,7 +25,7 @@ def question( resp = rag.graph_rag( question=question, user=user, collection=collection, entity_limit=entity_limit, triple_limit=triple_limit, - subgraph_limit=subgraph_limit + max_subgraph_size=max_subgraph_size ) print(resp) @@ -71,9 +74,9 @@ def main(): ) parser.add_argument( - '-s', '--subgraph-limit', - default=default_subgraph_limit, - help=f'Subgraph limit (default: {subgraph_limit})' + '-s', '--max-subgraph-size', + default=default_max_subgraph_size, + help=f'Max subgraph size (default: {default_max_subgraph_size})' ) args = parser.parse_args() @@ -87,7 +90,7 @@ def main(): collection=args.collection, entity_limit=args.entity_limit, triple_limit=args.triple_limit, - subgraph_limit=args.subgraph_limit, + max_subgraph_size=args.max_subgraph_size, ) except Exception as e: diff --git a/trustgraph-flow/trustgraph/gateway/graph_rag.py b/trustgraph-flow/trustgraph/gateway/graph_rag.py index 6979428a..fc02a17f 100644 --- a/trustgraph-flow/trustgraph/gateway/graph_rag.py +++ b/trustgraph-flow/trustgraph/gateway/graph_rag.py @@ -23,9 +23,9 @@ class GraphRagRequestor(ServiceRequestor): query=body["query"], user=body.get("user", "trustgraph"), collection=body.get("collection", "default"), - entity_limit=body.get("entity-limit", 50), - triple_limit=body.get("triple-limit", 30), - doc_limit=body.get("doc-limit", 1000), + entity_limit=int(body.get("entity-limit", 50)), + triple_limit=int(body.get("triple-limit", 30)), + max_subgraph_size=int(body.get("max-subgraph-size", 1000)), ) def from_response(self, message): diff --git a/trustgraph-flow/trustgraph/graph_rag.py b/trustgraph-flow/trustgraph/graph_rag.py index 721933fd..aa020a69 100644 --- a/trustgraph-flow/trustgraph/graph_rag.py +++ b/trustgraph-flow/trustgraph/graph_rag.py @@ -22,14 +22,14 @@ class Query: def __init__( 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.user = user self.collection = collection self.verbose = verbose self.entity_limit = entity_limit - self.query_limit = query_limit + self.triple_limit = triple_limit self.max_subgraph_size = max_subgraph_size def get_vector(self, query): @@ -99,7 +99,7 @@ class Query: res = self.rag.triples_client.request( user=self.user, collection=self.collection, s=e, p=None, o=None, - limit=self.query_limit + limit=self.triple_limit ) for triple in res: @@ -110,7 +110,7 @@ class Query: res = self.rag.triples_client.request( user=self.user, collection=self.collection, s=None, p=e, o=None, - limit=self.query_limit + limit=self.triple_limit ) for triple in res: @@ -121,7 +121,7 @@ class Query: res = self.rag.triples_client.request( user=self.user, collection=self.collection, s=None, p=None, o=e, - limit=self.query_limit, + limit=self.triple_limit, ) for triple in res: @@ -255,8 +255,8 @@ class GraphRag: print("Construct prompt...", flush=True) q = Query( - rag=self, user=user, collection=collection, verbose=self.verbose - entity_limit=entity_limit, query_limit=query_limit, + rag=self, user=user, collection=collection, verbose=self.verbose, + entity_limit=entity_limit, triple_limit=triple_limit, max_subgraph_size=max_subgraph_size, ) diff --git a/trustgraph-flow/trustgraph/retrieval/graph_rag/rag.py b/trustgraph-flow/trustgraph/retrieval/graph_rag/rag.py index 67ecfca0..d94d9a8f 100755 --- a/trustgraph-flow/trustgraph/retrieval/graph_rag/rag.py +++ b/trustgraph-flow/trustgraph/retrieval/graph_rag/rag.py @@ -115,17 +115,17 @@ class Processor(ConsumerProducer): if v.entity_limit: entity_limit = v.entity_limit else: - entity_limit = self.entity_limit + entity_limit = self.default_entity_limit if v.triple_limit: triple_limit = v.triple_limit else: - triple_limit = self.triple_limit + triple_limit = self.default_triple_limit if v.max_subgraph_size: max_subgraph_size = v.max_subgraph_size else: - max_subgraph_size = self.max_subgraph_size + max_subgraph_size = self.default_max_subgraph_size response = self.rag.query( query=v.query, user=v.user, collection=v.collection,