MCPcopy Create free account
hub / github.com/apache/arrow / MakeAggregateNodeArgs

Method MakeAggregateNodeArgs

cpp/src/arrow/acero/groupby_aggregate_node.cc:69–178  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

67}
68
69Result<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));

Callers

nothing calls this directly

Calls 13

push_backMethod · 0.80
WithNameMethod · 0.80
InvalidFunction · 0.50
NotImplementedFunction · 0.50
schemaFunction · 0.50
sizeMethod · 0.45
beginMethod · 0.45
endMethod · 0.45
findMethod · 0.45
nameMethod · 0.45
fieldMethod · 0.45
getMethod · 0.45

Tested by

no test coverage detected