@@ -110,37 +110,21 @@ def return_result(
110110 except asyncio .CancelledError :
111111 return KnowledgeGraphDict (nodes = {}, edges = {}), []
112112
113- # TODO: split batch by job and run splits concurrently
114- async def batch_subquery (self , batch : list [SubqContext ]) -> None :
115- """Produce query payloads and make them as a single batch query to the backend(s)."""
116- loggers = dict [int , TRAPILogger ]()
117-
118- query_mapping = list [tuple [SubqContext , QueryGraphDict , Transpiler ]]()
119- payloads = list [Any ]()
120- for subq in batch :
121- subq_id = hash ((subq .job , subq .branch .superposition_id ))
122- if subq_id not in loggers :
123- loggers [subq_id ] = TRAPILogger (subq .job )
124-
125- new_qgraphs , new_transpilers , new_payloads = self .make_payloads (
126- subq .branch , loggers [subq_id ]
127- )
128- payloads .extend (new_payloads )
129- query_mapping .extend (
130- [
131- (subq , qg , trans )
132- for qg , trans in zip (new_qgraphs , new_transpilers , strict = True )
133- ]
134- )
135-
113+ async def handle_subquery_batch (
114+ self ,
115+ payload_batch : list [Any ],
116+ query_mapping : list [tuple [SubqContext , QueryGraphDict , Transpiler ]],
117+ job_log : TRAPILogger ,
118+ ) -> None :
119+ """Run a given subquery payload batch and transform the results, sending them to the callback."""
136120 start = time .time ()
137121 logger .info (
138- f"Subquerying Tier 1 with batch of { len (payloads )} subqueries (originating from batch of { len ( batch ) } ) ..."
122+ f"Subquerying Tier 1 with batch of { len (payload_batch )} subqueries..."
139123 )
140124 query_driver = tier_manager .get_driver (1 )
141125 try :
142126 response_records = cast (
143- list [list [ESEdge ]], await query_driver .run_query (payloads )
127+ list [list [ESEdge ]], await query_driver .run_query (payload_batch )
144128 )
145129 split = time .time ()
146130 logger .success (f"Got results in { math .ceil ((split - start ) * 1000 )} ms" )
@@ -158,7 +142,7 @@ async def batch_subquery(self, batch: list[SubqContext]) -> None:
158142 try :
159143 append_aggregator_source (edge , Infores ("infores:retriever" ))
160144 except ValueError :
161- loggers [ subq_id ] .warning (
145+ job_log .warning (
162146 f"Edge f{ edge_id } has an invalid provenance chain."
163147 )
164148
@@ -187,22 +171,61 @@ async def batch_subquery(self, batch: list[SubqContext]) -> None:
187171
188172 for subq_id , result in results .items ():
189173 for callback in self .subscriptions .get (subq_id , []):
190- callback ((result ["knowledge_graph" ], loggers [ subq_id ] .get_logs ()))
174+ callback ((result ["knowledge_graph" ], job_log .get_logs ()))
191175 end = time .time ()
192176 logger .success (
193177 f"Transformed results and sent to original callers in { math .ceil ((end - split ) * 1000 )} ms"
194178 )
195179
196180 except Exception as e :
197- for subq_id , job_log in loggers .items ():
198- job_log .with_exception (
199- "An unhandled error occurred in the query driver." , exception = e
200- )
181+ job_log .with_exception (
182+ "An unhandled error occurred in the query driver." , exception = e
183+ )
184+ for subq , _ , _ in query_mapping :
185+ subq_id = hash ((subq .job , subq .branch .superposition_id ))
201186 for callback in self .subscriptions .get (subq_id , []):
202187 callback (
203188 (KnowledgeGraphDict (nodes = {}, edges = {}), job_log .get_logs ())
204189 )
205190
191+ async def batch_subquery (self , batch : list [SubqContext ]) -> None :
192+ """Produce query payloads and make them as a single batch query to the backend(s)."""
193+ loggers = dict [str , TRAPILogger ]()
194+
195+ query_mapping = dict [
196+ str , list [tuple [SubqContext , QueryGraphDict , Transpiler ]]
197+ ]()
198+ payloads = dict [str , list [Any ]]()
199+ for subq in batch :
200+ if subq .job not in loggers :
201+ loggers [subq .job ] = TRAPILogger (subq .job )
202+
203+ new_qgraphs , new_transpilers , new_payloads = self .make_payloads (
204+ subq .branch , loggers [subq .job ]
205+ )
206+
207+ # Separate payloads by job so queries can't time-contaminate each other
208+ if subq .job not in payloads :
209+ payloads [subq .job ] = []
210+ query_mapping [subq .job ] = []
211+ payloads [subq .job ].extend (new_payloads )
212+
213+ query_mapping [subq .job ].extend (
214+ [
215+ (subq , qg , trans )
216+ for qg , trans in zip (new_qgraphs , new_transpilers , strict = True )
217+ ]
218+ )
219+
220+ for job , job_payloads in payloads .items ():
221+ self .tasks .append (
222+ asyncio .create_task (
223+ self .handle_subquery_batch (
224+ job_payloads , query_mapping [job ], loggers [job ]
225+ )
226+ )
227+ )
228+
206229 def make_payloads (
207230 self , branch : Branch , job_log : TRAPILogger
208231 ) -> tuple [list [QueryGraphDict ], list [Transpiler ], list [Any ]]:
0 commit comments