Skip to content

Commit e20e8af

Browse files
committed
SubqueryDispatcher: batch by job
Prevents time contamination between jobs
1 parent 17524ce commit e20e8af

1 file changed

Lines changed: 54 additions & 31 deletions

File tree

src/retriever/lookup/subquery.py

Lines changed: 54 additions & 31 deletions
Original file line numberDiff line numberDiff line change
@@ -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

Comments
 (0)