88#include < variant>
99#include < vector>
1010
11+ #include < arrow/acero/options.h>
12+ #include < arrow/compute/exec.h>
1113#include < fmt/format.h>
1214#include < fmt/ranges.h>
1315#include < oneapi/tbb/blocked_range.h>
2123#include " silo/query_engine/actions/action.h"
2224#include " silo/query_engine/bad_request.h"
2325#include " silo/query_engine/copy_on_write_bitmap.h"
26+ #include " silo/query_engine/exec_node/arrow_util.h"
27+ #include " silo/query_engine/exec_node/json_value_type_array_builder.h"
2428#include " silo/query_engine/query_result.h"
29+ #include " silo/storage/column/column_type_visitor.h"
2530#include " silo/storage/column/sequence_column.h"
2631#include " silo/storage/table_partition.h"
2732
@@ -196,12 +201,13 @@ void Mutations<SymbolType>::validateOrderByFields(const schema::TableSchema& /*t
196201}
197202
198203template <typename SymbolType>
199- void Mutations<SymbolType>::addMutationsToOutput(
204+ arrow::Status Mutations<SymbolType>::addMutationsToOutput(
200205 const std::string& sequence_name,
201206 const storage::column::SequenceColumnMetadata<SymbolType>& sequence_column_metadata,
207+ double min_proportion,
202208 const PrefilteredBitmaps& bitmap_filter,
203- std::vector<QueryResultEntry >& output
204- ) const {
209+ std::unordered_map<std::string_view, exec_node::JsonValueTypeArrayBuilder >& output_builder
210+ ) {
205211 const size_t sequence_length = sequence_column_metadata.reference_sequence .size ();
206212
207213 const SymbolMap<SymbolType, std::vector<uint32_t >> count_of_mutations_per_position =
@@ -228,56 +234,62 @@ void Mutations<SymbolType>::addMutationsToOutput(
228234 const uint32_t count = count_of_mutations_per_position.at (symbol)[pos];
229235 if (count > threshold_count) {
230236 const double proportion = static_cast <double >(count) / static_cast <double >(total);
231- std::map<std::string, common::JsonValueType> fields_for_row;
232- if (std::ranges::find (fields, MUTATION_FIELD_NAME ) != fields.end ()) {
233- fields_for_row.emplace (
234- MUTATION_FIELD_NAME ,
235- fmt::format (
236- " {}{}{}" ,
237- SymbolType::symbolToChar (symbol_in_reference_genome),
238- pos + 1 ,
239- SymbolType::symbolToChar (symbol)
240- )
241- );
237+ if (auto builder = output_builder.find (MUTATION_FIELD_NAME );
238+ builder != output_builder.end ()) {
239+ ARROW_RETURN_NOT_OK (builder->second .insert ({fmt::format (
240+ " {}{}{}" ,
241+ SymbolType::symbolToChar (symbol_in_reference_genome),
242+ pos + 1 ,
243+ SymbolType::symbolToChar (symbol)
244+ )}));
242245 }
243- if (std::ranges::find (fields, MUTATION_FROM_FIELD_NAME ) != fields. end ()) {
244- fields_for_row. emplace (
245- MUTATION_FROM_FIELD_NAME ,
246- std::string (1 , SymbolType::symbolToChar (symbol_in_reference_genome))
247- );
246+ if (auto builder = output_builder. find ( MUTATION_FROM_FIELD_NAME );
247+ builder != output_builder. end ()) {
248+ ARROW_RETURN_NOT_OK (builder-> second . insert (
249+ { std::string (1 , SymbolType::symbolToChar (symbol_in_reference_genome))}
250+ )) ;
248251 }
249- if (std::ranges::find (fields, MUTATION_TO_FIELD_NAME ) != fields.end ()) {
250- fields_for_row.emplace (
251- MUTATION_TO_FIELD_NAME , std::string (1 , SymbolType::symbolToChar (symbol))
252+ if (auto builder = output_builder.find (MUTATION_TO_FIELD_NAME );
253+ builder != output_builder.end ()) {
254+ ARROW_RETURN_NOT_OK (
255+ builder->second .insert ({std::string (1 , SymbolType::symbolToChar (symbol))})
252256 );
253257 }
254- if (std::ranges::find (fields, POSITION_FIELD_NAME ) != fields.end ()) {
255- fields_for_row.emplace (POSITION_FIELD_NAME , static_cast <int32_t >(pos + 1 ));
258+ if (auto builder = output_builder.find (POSITION_FIELD_NAME );
259+ builder != output_builder.end ()) {
260+ ARROW_RETURN_NOT_OK (builder->second .insert ({static_cast <int32_t >(pos + 1 )}));
256261 }
257- if (std::ranges::find (fields, SEQUENCE_FIELD_NAME ) != fields.end ()) {
258- fields_for_row.emplace (SEQUENCE_FIELD_NAME , sequence_name);
262+ if (auto builder = output_builder.find (SEQUENCE_FIELD_NAME );
263+ builder != output_builder.end ()) {
264+ ARROW_RETURN_NOT_OK (builder->second .insert ({sequence_name}));
259265 }
260- if (std::ranges::find (fields, PROPORTION_FIELD_NAME ) != fields.end ()) {
261- fields_for_row.emplace (PROPORTION_FIELD_NAME , proportion);
266+ if (auto builder = output_builder.find (PROPORTION_FIELD_NAME );
267+ builder != output_builder.end ()) {
268+ ARROW_RETURN_NOT_OK (builder->second .insert ({proportion}));
262269 }
263- if (std::ranges::find (fields, COUNT_FIELD_NAME ) != fields.end ()) {
264- fields_for_row.emplace (COUNT_FIELD_NAME , static_cast <int32_t >(count));
270+ if (auto builder = output_builder.find (COUNT_FIELD_NAME );
271+ builder != output_builder.end ()) {
272+ ARROW_RETURN_NOT_OK (builder->second .insert ({static_cast <int32_t >(count)}));
265273 }
266- if (std::ranges::find (fields, COVERAGE_FIELD_NAME ) != fields.end ()) {
267- fields_for_row.emplace (COVERAGE_FIELD_NAME , static_cast <int32_t >(total));
274+ if (auto builder = output_builder.find (COVERAGE_FIELD_NAME );
275+ builder != output_builder.end ()) {
276+ ARROW_RETURN_NOT_OK (builder->second .insert ({static_cast <int32_t >(total)}));
268277 }
269- output.push_back ({fields_for_row});
270278 }
271279 }
272280 }
273281 }
282+ return arrow::Status::OK ();
274283}
275284
285+ using silo::query_engine::filter::operators::Operator;
286+
276287template <typename SymbolType>
277- QueryResult Mutations<SymbolType>::execute (
288+ arrow::Result<QueryPlan> Mutations<SymbolType>::toQueryPlanImpl (
278289 std::shared_ptr<const storage::Table> table,
279- std::vector<CopyOnWriteBitmap> bitmap_filter
280- ) const {
290+ std::shared_ptr<std::vector<std::unique_ptr<Operator>>> partition_filter_operators,
291+ const config::QueryOptions& query_options
292+ ) {
281293 std::vector<std::string> sequence_names_to_evaluate;
282294 for (const auto & sequence_name : sequence_names) {
283295 auto column_identifier = table->schema .getColumn (sequence_name);
@@ -295,54 +307,115 @@ QueryResult Mutations<SymbolType>::execute(
295307 }
296308 }
297309
298- std::unordered_map<std::string, Mutations<SymbolType>::PrefilteredBitmaps> bitmaps_to_evaluate =
299- preFilterBitmaps (*table, bitmap_filter);
310+ auto output_fields = getOutputSchema (table->schema );
311+
312+ double given_min_proportion = min_proportion;
313+
314+ std::function<arrow::Future<std::optional<arrow::ExecBatch>>()> producer =
315+ [table,
316+ given_min_proportion,
317+ output_fields,
318+ partition_filter_operators,
319+ sequence_names_to_evaluate,
320+ produced = false ]() mutable -> arrow::Future<std::optional<arrow::ExecBatch>> {
321+ if (produced == true ) {
322+ std::optional<arrow::ExecBatch> result = std::nullopt ;
323+ return arrow::Future{result};
324+ }
325+ produced = true ;
326+ std::vector<CopyOnWriteBitmap> partition_filters;
327+ partition_filters.reserve (partition_filter_operators->size ());
328+ for (const auto & partition_filter_operator : *partition_filter_operators) {
329+ partition_filters.emplace_back (partition_filter_operator->evaluate ());
330+ }
300331
301- std::vector<QueryResultEntry> mutation_proportions;
302- for (const auto & sequence_name : sequence_names_to_evaluate) {
303- const storage::column::SequenceColumnMetadata<SymbolType>* sequence_column_metadata =
304- table->schema .getColumnMetadata <typename SymbolType::Column>(sequence_name).value ();
332+ std::unordered_map<std::string, Mutations<SymbolType>::PrefilteredBitmaps>
333+ bitmaps_to_evaluate = preFilterBitmaps (*table, partition_filters);
305334
306- if (bitmaps_to_evaluate.contains (sequence_name)) {
307- addMutationsToOutput (
308- sequence_name,
309- *sequence_column_metadata,
310- bitmaps_to_evaluate.at (sequence_name),
311- mutation_proportions
335+ std::unordered_map<std::string_view, exec_node::JsonValueTypeArrayBuilder> output_builder;
336+ for (const auto & output_field : output_fields) {
337+ output_builder.emplace (
338+ output_field.name , exec_node::columnTypeToArrowType (output_field.type )
312339 );
313340 }
314- }
315- return QueryResult::fromVector (std::move (mutation_proportions));
341+
342+ for (const auto & sequence_name : sequence_names_to_evaluate) {
343+ const storage::column::SequenceColumnMetadata<SymbolType>* sequence_column_metadata =
344+ table->schema .getColumnMetadata <typename SymbolType::Column>(sequence_name).value ();
345+
346+ if (bitmaps_to_evaluate.contains (sequence_name)) {
347+ ARROW_RETURN_NOT_OK (addMutationsToOutput (
348+ sequence_name,
349+ *sequence_column_metadata,
350+ given_min_proportion,
351+ bitmaps_to_evaluate.at (sequence_name),
352+ output_builder
353+ ));
354+ }
355+ }
356+ // Order of result_columns is relevant as it needs to be consistent with vector in schema
357+ std::vector<arrow::Datum> result_columns;
358+ for (const auto & output_field : output_fields) {
359+ if (auto array_builder = output_builder.find (output_field.name );
360+ array_builder != output_builder.end ()) {
361+ arrow::Datum datum;
362+ ARROW_ASSIGN_OR_RAISE (datum, array_builder->second .toDatum ());
363+ result_columns.push_back (std::move (datum));
364+ }
365+ }
366+ ARROW_ASSIGN_OR_RAISE (
367+ std::optional<arrow::ExecBatch> result, arrow::ExecBatch::Make (result_columns)
368+ );
369+ return arrow::Future{result};
370+ };
371+
372+ ARROW_ASSIGN_OR_RAISE (auto arrow_plan, arrow::acero::ExecPlan::Make ());
373+
374+ arrow::acero::SourceNodeOptions options{
375+ exec_node::columnsToArrowSchema (getOutputSchema (table->schema )),
376+ std::move (producer),
377+ arrow::Ordering::Implicit ()
378+ };
379+ ARROW_ASSIGN_OR_RAISE (
380+ auto node, arrow::acero::MakeExecNode (" source" , arrow_plan.get (), {}, options)
381+ );
382+
383+ ARROW_ASSIGN_OR_RAISE (node, addSortNode (arrow_plan.get (), node, table->schema ));
384+
385+ ARROW_ASSIGN_OR_RAISE (node, addLimitAndOffsetNode (arrow_plan.get (), node));
386+
387+ return QueryPlan::makeQueryPlan (arrow_plan, node);
316388}
317389
318390template <typename SymbolType>
319391std::vector<schema::ColumnIdentifier> Mutations<SymbolType>::getOutputSchema(
320392 const silo::schema::TableSchema& table_schema
321393) const {
394+ using silo::schema::ColumnType;
322395 std::vector<schema::ColumnIdentifier> output_fields;
323396 if (std::ranges::find (fields, MUTATION_FIELD_NAME ) != fields.end ()) {
324- output_fields.emplace_back (std::string (MUTATION_FIELD_NAME ), schema:: ColumnType::STRING );
397+ output_fields.emplace_back (std::string (MUTATION_FIELD_NAME ), ColumnType::STRING );
325398 }
326399 if (std::ranges::find (fields, MUTATION_FROM_FIELD_NAME ) != fields.end ()) {
327- output_fields.emplace_back (std::string (MUTATION_FROM_FIELD_NAME ), schema:: ColumnType::STRING );
400+ output_fields.emplace_back (std::string (MUTATION_FROM_FIELD_NAME ), ColumnType::STRING );
328401 }
329402 if (std::ranges::find (fields, MUTATION_TO_FIELD_NAME ) != fields.end ()) {
330- output_fields.emplace_back (std::string (MUTATION_TO_FIELD_NAME ), schema:: ColumnType::STRING );
403+ output_fields.emplace_back (std::string (MUTATION_TO_FIELD_NAME ), ColumnType::STRING );
331404 }
332405 if (std::ranges::find (fields, SEQUENCE_FIELD_NAME ) != fields.end ()) {
333- output_fields.emplace_back (std::string (SEQUENCE_FIELD_NAME ), schema:: ColumnType::STRING );
406+ output_fields.emplace_back (std::string (SEQUENCE_FIELD_NAME ), ColumnType::STRING );
334407 }
335408 if (std::ranges::find (fields, POSITION_FIELD_NAME ) != fields.end ()) {
336- output_fields.emplace_back (std::string (POSITION_FIELD_NAME ), schema:: ColumnType::INT );
409+ output_fields.emplace_back (std::string (POSITION_FIELD_NAME ), ColumnType::INT );
337410 }
338411 if (std::ranges::find (fields, PROPORTION_FIELD_NAME ) != fields.end ()) {
339- output_fields.emplace_back (std::string (PROPORTION_FIELD_NAME ), schema:: ColumnType::FLOAT );
412+ output_fields.emplace_back (std::string (PROPORTION_FIELD_NAME ), ColumnType::FLOAT );
340413 }
341414 if (std::ranges::find (fields, COVERAGE_FIELD_NAME ) != fields.end ()) {
342- output_fields.emplace_back (std::string (COVERAGE_FIELD_NAME ), schema:: ColumnType::INT );
415+ output_fields.emplace_back (std::string (COVERAGE_FIELD_NAME ), ColumnType::INT );
343416 }
344417 if (std::ranges::find (fields, COUNT_FIELD_NAME ) != fields.end ()) {
345- output_fields.emplace_back (std::string (COUNT_FIELD_NAME ), schema:: ColumnType::INT );
418+ output_fields.emplace_back (std::string (COUNT_FIELD_NAME ), ColumnType::INT );
346419 }
347420 return output_fields;
348421}
0 commit comments