Skip to content

Commit 96c6d70

Browse files
committed
refactor: migrate Mutations action to arrow framework
1 parent baf0fcc commit 96c6d70

13 files changed

Lines changed: 296 additions & 177 deletions

src/silo/query_engine/actions/action.cpp

Lines changed: 4 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -308,7 +308,8 @@ std::vector<schema::ColumnIdentifier> columnNamesToFields(
308308

309309
arrow::Result<QueryPlan> Action::toQueryPlanImpl(
310310
std::shared_ptr<const storage::Table> table,
311-
const std::vector<std::unique_ptr<filter::operators::Operator>>& partition_filter_operators,
311+
std::shared_ptr<std::vector<std::unique_ptr<filter::operators::Operator>>>
312+
partition_filter_operators,
312313
const config::QueryOptions& query_options
313314
) {
314315
ARROW_ASSIGN_OR_RAISE(auto arrow_plan, arrow::acero::ExecPlan::Make());
@@ -325,7 +326,8 @@ arrow::Result<QueryPlan> Action::toQueryPlanImpl(
325326

326327
QueryPlan Action::toQueryPlan(
327328
std::shared_ptr<const storage::Table> table,
328-
const std::vector<std::unique_ptr<filter::operators::Operator>>& partition_filter_operators,
329+
std::shared_ptr<std::vector<std::unique_ptr<filter::operators::Operator>>>
330+
partition_filter_operators,
329331
const config::QueryOptions& query_options
330332
) {
331333
auto query_plan = toQueryPlanImpl(table, partition_filter_operators, query_options);

src/silo/query_engine/actions/action.h

Lines changed: 4 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -51,7 +51,8 @@ class Action {
5151

5252
QueryPlan toQueryPlan(
5353
std::shared_ptr<const storage::Table> table,
54-
const std::vector<std::unique_ptr<filter::operators::Operator>>& partition_filter_operators,
54+
std::shared_ptr<std::vector<std::unique_ptr<filter::operators::Operator>>>
55+
partition_filter_operators,
5556
const config::QueryOptions& query_options
5657
);
5758

@@ -77,7 +78,8 @@ class Action {
7778
// If this method is not overloaded, a LegacyResultProducer will be created instead
7879
virtual arrow::Result<QueryPlan> toQueryPlanImpl(
7980
std::shared_ptr<const storage::Table> table,
80-
const std::vector<std::unique_ptr<filter::operators::Operator>>& partition_filter_operators,
81+
std::shared_ptr<std::vector<std::unique_ptr<filter::operators::Operator>>>
82+
partition_filter_operators,
8183
const config::QueryOptions& query_options
8284
);
8385

src/silo/query_engine/actions/mutations.cpp

Lines changed: 131 additions & 58 deletions
Original file line numberDiff line numberDiff line change
@@ -8,6 +8,8 @@
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>
@@ -21,7 +23,10 @@
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

198203
template <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+
276287
template <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

318390
template <typename SymbolType>
319391
std::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
}

src/silo/query_engine/actions/mutations.h

Lines changed: 19 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -8,10 +8,14 @@
88
#include <utility>
99
#include <vector>
1010

11+
#include <arrow/array/builder_base.h>
12+
#include <arrow/array/builder_binary.h>
13+
#include <arrow/array/builder_primitive.h>
1114
#include <nlohmann/json_fwd.hpp>
1215

1316
#include "silo/common/symbol_map.h"
1417
#include "silo/query_engine/actions/action.h"
18+
#include "silo/query_engine/exec_node/json_value_type_array_builder.h"
1519
#include "silo/query_engine/query_result.h"
1620
#include "silo/storage/column/sequence_column.h"
1721
#include "silo/storage/table.h"
@@ -79,19 +83,29 @@ class Mutations : public Action {
7983
const PrefilteredBitmaps& bitmap_filter
8084
);
8185

82-
void addMutationsToOutput(
86+
static arrow::Status addMutationsToOutput(
8387
const std::string& sequence_name,
8488
const storage::column::SequenceColumnMetadata<SymbolType>& sequence_store,
89+
double min_proportion,
8590
const PrefilteredBitmaps& bitmap_filter,
86-
std::vector<QueryResultEntry>& output
87-
) const;
91+
std::unordered_map<std::string_view, exec_node::JsonValueTypeArrayBuilder>& output_builder
92+
);
8893

8994
void validateOrderByFields(const schema::TableSchema& schema) const override;
9095

91-
[[nodiscard]] QueryResult execute(
96+
QueryResult execute(
9297
std::shared_ptr<const storage::Table> table,
9398
std::vector<CopyOnWriteBitmap> bitmap_filter
94-
) const override;
99+
) const override {
100+
SILO_PANIC("Legacy execute called on already migrated action. Programming error.");
101+
}
102+
103+
arrow::Result<QueryPlan> toQueryPlanImpl(
104+
std::shared_ptr<const storage::Table> table,
105+
std::shared_ptr<std::vector<std::unique_ptr<filter::operators::Operator>>>
106+
partition_filter_operators,
107+
const config::QueryOptions& query_options
108+
) override;
95109

96110
public:
97111
explicit Mutations(

src/silo/query_engine/actions/simple_select_action.cpp

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -33,7 +33,8 @@ void SimpleSelectAction::validateOrderByFields(const schema::TableSchema& schema
3333

3434
arrow::Result<QueryPlan> SimpleSelectAction::toQueryPlanImpl(
3535
std::shared_ptr<const storage::Table> table,
36-
const std::vector<std::unique_ptr<filter::operators::Operator>>& partition_filter_operators,
36+
std::shared_ptr<std::vector<std::unique_ptr<filter::operators::Operator>>>
37+
partition_filter_operators,
3738
const config::QueryOptions& query_options
3839
) {
3940
validateOrderByFields(table->schema);

src/silo/query_engine/actions/simple_select_action.h

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -16,7 +16,8 @@ class SimpleSelectAction : public Action {
1616

1717
arrow::Result<QueryPlan> toQueryPlanImpl(
1818
std::shared_ptr<const storage::Table> table,
19-
const std::vector<std::unique_ptr<filter::operators::Operator>>& partition_filter_operators,
19+
std::shared_ptr<std::vector<std::unique_ptr<filter::operators::Operator>>>
20+
partition_filter_operators,
2021
const config::QueryOptions& query_options
2122
) override;
2223
};

0 commit comments

Comments
 (0)