| 67 | } |
| 68 | |
| 69 | Result<AggregateNodeArgs<HashAggregateKernel>> GroupByNode::MakeAggregateNodeArgs( |
| 70 | const std::shared_ptr<Schema>& input_schema, const std::vector<FieldRef>& keys, |
| 71 | const std::vector<FieldRef>& segment_keys, const std::vector<Aggregate>& aggs, |
| 72 | ExecContext* ctx, const bool is_cpu_parallel) { |
| 73 | // Find input field indices for key fields |
| 74 | std::vector<int> key_field_ids(keys.size()); |
| 75 | for (size_t i = 0; i < keys.size(); ++i) { |
| 76 | ARROW_ASSIGN_OR_RAISE(auto match, keys[i].FindOne(*input_schema)); |
| 77 | key_field_ids[i] = match[0]; |
| 78 | } |
| 79 | |
| 80 | // Find input field indices for segment key fields |
| 81 | std::vector<int> segment_key_field_ids(segment_keys.size()); |
| 82 | for (size_t i = 0; i < segment_keys.size(); ++i) { |
| 83 | ARROW_ASSIGN_OR_RAISE(auto match, segment_keys[i].FindOne(*input_schema)); |
| 84 | segment_key_field_ids[i] = match[0]; |
| 85 | } |
| 86 | |
| 87 | // Check key fields and segment key fields are disjoint |
| 88 | std::unordered_set<int> key_field_id_set(key_field_ids.begin(), key_field_ids.end()); |
| 89 | for (const auto& segment_key_field_id : segment_key_field_ids) { |
| 90 | if (key_field_id_set.find(segment_key_field_id) != key_field_id_set.end()) { |
| 91 | return Status::Invalid("Group-by aggregation with field '", |
| 92 | input_schema->field(segment_key_field_id)->name(), |
| 93 | "' as both key and segment key"); |
| 94 | } |
| 95 | } |
| 96 | |
| 97 | // Find input field indices for aggregates |
| 98 | std::vector<std::vector<int>> agg_src_fieldsets(aggs.size()); |
| 99 | for (size_t i = 0; i < aggs.size(); ++i) { |
| 100 | const auto& target_fieldset = aggs[i].target; |
| 101 | for (const auto& target : target_fieldset) { |
| 102 | ARROW_ASSIGN_OR_RAISE(auto match, target.FindOne(*input_schema)); |
| 103 | agg_src_fieldsets[i].push_back(match[0]); |
| 104 | } |
| 105 | } |
| 106 | |
| 107 | // Build vector of aggregate source field data types |
| 108 | std::vector<std::vector<TypeHolder>> agg_src_types(aggs.size()); |
| 109 | for (size_t i = 0; i < aggs.size(); ++i) { |
| 110 | for (const auto& agg_src_field_id : agg_src_fieldsets[i]) { |
| 111 | agg_src_types[i].push_back(input_schema->field(agg_src_field_id)->type().get()); |
| 112 | } |
| 113 | } |
| 114 | |
| 115 | // Build vector of segment key field data types |
| 116 | std::vector<TypeHolder> segment_key_types(segment_keys.size()); |
| 117 | for (size_t i = 0; i < segment_keys.size(); ++i) { |
| 118 | auto segment_key_field_id = segment_key_field_ids[i]; |
| 119 | segment_key_types[i] = input_schema->field(segment_key_field_id)->type().get(); |
| 120 | } |
| 121 | |
| 122 | ARROW_ASSIGN_OR_RAISE(auto segmenter, RowSegmenter::Make(std::move(segment_key_types), |
| 123 | /*nullable_keys=*/false, ctx)); |
| 124 | |
| 125 | // Construct aggregates |
| 126 | ARROW_ASSIGN_OR_RAISE(auto agg_kernels, GetKernels(ctx, aggs, agg_src_types)); |
nothing calls this directly
no test coverage detected