ProvSQL C/C++ API
Adding support for provenance and uncertainty management to PostgreSQL databases
Loading...
Searching...
No Matches
Aggregation.cpp
Go to the documentation of this file.
1/**
2 * @file Aggregation.cpp
3 * @brief Aggregation operator and accumulator implementations.
4 *
5 * Implements the two factory functions declared in @c Aggregation.h:
6 * - @c getAggregationOperator(): maps a PostgreSQL aggregate function OID
7 * (looked up by name via @c get_func_name()) to an @c AggregationOperator
8 * enum value.
9 * - @c makeAggregator(): constructs a concrete @c Aggregator subclass
10 * for the given operator/type combination.
11 *
12 * Each built aggregation function × value-type combination has its own
13 * @c Aggregator subclass defined locally in this file (e.g. @c SumAgg<long>,
14 * @c MinAgg<double>, @c ChooseAgg<long>). Only the aggregates the Monte-Carlo
15 * sampler and the subset enumerator evaluate directly are built: the numeric
16 * ones (SUM / COUNT / MIN / MAX / AVG) and CHOOSE (the categorical analog,
17 * decided by the exhaustive subset enumerator). The boolean (bool_or /
18 * bool_and) and array_agg aggregates are resolved to a Boolean subcircuit by
19 * the m-semiring HAVING rewrite (@c having_semantics) and so never reach
20 * @c makeAggregator.
21 */
22#include "Aggregation.h"
23
24#include <string>
25#include <stdexcept>
26
27extern "C" {
28#include "utils/lsyscache.h"
29#include "utils/elog.h"
30#include "provsql_utils.h"
31}
32
33#include "provsql_error.h"
34
36{
37 char *fname = get_func_name(oid);
38
39 if(fname == nullptr)
40 provsql_error("Invalid OID for aggregation function: %d", oid);
41
42 std::string func_name {fname};
43 pfree(fname);
44
46
47 if(func_name == "count") {
49 } else if(func_name == "sum") {
51 } else if(func_name == "min") {
53 } else if(func_name == "max") {
55 } else if(func_name == "choose") {
57 } else if(func_name == "avg") {
59 } else if(func_name == "array_agg" || func_name == "array_collect") {
61 } else if(func_name == "bool_and" || func_name == "every") {
63 } else if(func_name == "bool_or") {
65 } else {
66 provsql_unsupported(PROVSQL_GAP, "aggregate-no-evaluator", "Aggregation operator %s not supported", func_name.c_str());
67 }
68
69 return op;
70}
71
72ComparisonOperator cmpOpFromOid(Oid op_oid, bool &ok)
73{
74 ok = false;
75 char *opname = get_opname(op_oid);
76 if(opname == nullptr)
78
79 std::string s {opname};
80 pfree(opname);
81
82 ok = true;
83 if(s == "=") return ComparisonOperator::EQ;
84 if(s == "<>") return ComparisonOperator::NE;
85 if(s == "<") return ComparisonOperator::LT;
86 if(s == "<=") return ComparisonOperator::LE;
87 if(s == ">") return ComparisonOperator::GT;
88 if(s == ">=") return ComparisonOperator::GE;
89
90 ok = false;
92}
93
121
122template <class ...>
123struct False : std::bool_constant<false> { };
124
125/**
126 * @brief Base aggregator template for scalar types (int, float, bool, string).
127 *
128 * @tparam T The C++ type of the accumulated value.
129 */
130template <class T>
132protected:
133 T value{}; ///< Current accumulated value
134 bool has = false; ///< @c true once the first non-NULL input has been seen
135
136public:
137 /** @brief Return the accumulated value, or NULL if no inputs were seen. */
138 AggValue finalize() const override {
139 if (has) return AggValue {value}; else return AggValue{};
140 }
141 /** @brief Return the value type corresponding to @c T. */
142 ValueType inputType() const override {
143 if constexpr (std::is_same_v<T,long>)
144 return ValueType::INT;
145 else if constexpr (std::is_same_v<T,double>)
146 return ValueType::FLOAT;
147 else if constexpr (std::is_same_v<T,bool>)
148 return ValueType::BOOLEAN;
149 else if constexpr (std::is_same_v<T,std::string>)
150 return ValueType::STRING;
151 else
152 static_assert(False<T>{});
153 }
154};
155
156/** @brief Aggregator implementing SUM for integer or float types. */
157template <class T>
159 using StandardAgg<T>::value;
160 using StandardAgg<T>::has;
161
162 void add(const AggValue& x) override {
163 if (x.getType() == ValueType::NONE) return;
164 const T& v = std::get<T>(x.v);
165 value += v;
166 has = true;
167 }
168};
169
170/** @brief Aggregator implementing MIN for integer or float types. */
171template <class T>
172struct MinAgg : StandardAgg<T> {
173 using StandardAgg<T>::value;
174 using StandardAgg<T>::has;
175
176 void add(const AggValue& x) override {
177 if (x.getType() == ValueType::NONE) return;
178 const T& v = std::get<T>(x.v);
179 if(has) {
180 if(v < value) value = v;
181 } else {
182 value = v;
183 has = true;
184 }
185 }
186};
187
188/** @brief Aggregator implementing MAX for integer or float types. */
189template <class T>
190struct MaxAgg : StandardAgg<T> {
191 using StandardAgg<T>::value;
192 using StandardAgg<T>::has;
193
194 void add(const AggValue& x) override {
195 if (x.getType() == ValueType::NONE) return;
196 const T& v = std::get<T>(x.v);
197 if(has) {
198 if(v > value) value = v;
199 } else {
200 value = v;
201 has = true;
202 }
203 }
204};
205
206/** @brief Aggregator implementing CHOOSE (returns the first non-NULL input). */
207template <class T>
209 using StandardAgg<T>::value;
210 using StandardAgg<T>::has;
211
212 void add(const AggValue& x) override {
213 if (x.getType() == ValueType::NONE) return;
214 if(!has)
215 value = std::get<T>(x.v);
216 has = true;
217 }
218};
219
220/** @brief Aggregator implementing AVG; always returns a float result. */
221template <class T>
223protected:
224 double sum = 0; ///< Running sum of all non-NULL input values
225 unsigned count = 0; ///< Number of non-NULL inputs seen so far
226 bool has = false; ///< @c true once the first non-NULL input has been seen
227
228public:
229 void add(const AggValue& x) override {
230 if (x.getType() == ValueType::NONE) return;
231 const T& v = std::get<T>(x.v);
232 sum += v;
233 ++count;
234 has = true;
235 }
236 AggValue finalize() const override {
237 if (has) return AggValue {sum/count}; else return AggValue{};
238 }
239 ValueType inputType() const override {
240 if constexpr (std::is_same_v<T,long>)
241 return ValueType::INT;
242 else if constexpr (std::is_same_v<T,double>)
243 return ValueType::FLOAT;
244 else
245 static_assert(False<T>{});
246 }
247 ValueType resultType() const override {
248 return ValueType::FLOAT;
249 }
250};
251
252// Constructs the deterministic accumulator the Monte-Carlo sampler and the
253// exhaustive subset enumerator push per-world values into. The numeric
254// aggregates (SUM / COUNT / MIN / MAX / AVG) and CHOOSE are built; the boolean
255// (bool_or / bool_and) and array_agg aggregates never reach this factory: the
256// m-semiring HAVING rewrite in having_semantics resolves them to a Boolean
257// subcircuit before probability evaluation, so no such gate_agg survives to the
258// sampler. They are rejected explicitly rather than handled.
259std::unique_ptr<Aggregator> makeAggregator(AggregationOperator op, ValueType t) {
260 switch (op) {
262 // Each row contributes 1 (count(*)) or 0/1 (count(expr)), so the count is
263 // the sum of the contributions; the operator stays COUNT so the empty set
264 // reads as 0 rather than a sum's NULL.
265 if (t == ValueType::INT) return std::make_unique<SumAgg<long> >();
266 throw std::runtime_error("COUNT expects an integer-valued contribution");
268 switch (t) {
269 case ValueType::INT: return std::make_unique<SumAgg<long> >();
270 case ValueType::FLOAT: return std::make_unique<SumAgg<double> >();
271 default: throw std::runtime_error("SUM not supported for this type");
272 }
274 switch (t) {
275 case ValueType::INT: return std::make_unique<MinAgg<long> >();
276 case ValueType::FLOAT: return std::make_unique<MinAgg<double> >();
277 default: throw std::runtime_error("MIN not supported for this type");
278 }
280 switch (t) {
281 case ValueType::INT: return std::make_unique<MaxAgg<long> >();
282 case ValueType::FLOAT: return std::make_unique<MaxAgg<double> >();
283 default: throw std::runtime_error("MAX not supported for this type");
284 }
286 switch (t) {
287 case ValueType::INT: return std::make_unique<AvgAgg<long> >();
288 case ValueType::FLOAT: return std::make_unique<AvgAgg<double> >();
289 default: throw std::runtime_error("AVG not supported for this type");
290 }
292 switch(t) {
293 case ValueType::BOOLEAN: return std::make_unique<ChooseAgg<bool> >();
294 case ValueType::INT: return std::make_unique<ChooseAgg<long> >();
295 case ValueType::FLOAT: return std::make_unique<ChooseAgg<double> >();
296 case ValueType::STRING: return std::make_unique<ChooseAgg<std::string> >();
297 default: throw std::runtime_error("CHOOSE not supported for this type");
298 }
303 // Resolved to a Boolean subcircuit by the HAVING rewrite; never sampled.
304 throw std::runtime_error(
305 "makeAggregator: boolean/array_agg aggregates are handled by the "
306 "m-semiring HAVING rewrite, not the deterministic sampler");
307 }
308
309 throw std::logic_error("Unhandled AggregationOperator");
310}
ArithmeticOperator arithOpFromTag(unsigned tag, bool &ok)
Map a gate_arith operator tag to an ArithmeticOperator.
ComparisonOperator cmpOpFromOid(Oid op_oid, bool &ok)
Map a PostgreSQL comparison-operator OID to a ComparisonOperator.
AggregationOperator getAggregationOperator(Oid oid)
Map a PostgreSQL aggregate function OID to an AggregationOperator.
std::unique_ptr< Aggregator > makeAggregator(AggregationOperator op, ValueType t)
Create a concrete Aggregator for the given operator and value type.
Typed aggregation value, operator, and aggregator abstractions.
AggregationOperator
SQL aggregation functions tracked by ProvSQL.
Definition Aggregation.h:51
@ OR
Boolean OR aggregate.
Definition Aggregation.h:58
@ MAX
MAX → input type.
Definition Aggregation.h:55
@ COUNT
COUNT(*) or COUNT(expr) → integer.
Definition Aggregation.h:52
@ AND
Boolean AND aggregate.
Definition Aggregation.h:57
@ SUM
SUM → integer or float.
Definition Aggregation.h:53
@ ARRAY_AGG
Array aggregation.
Definition Aggregation.h:60
@ NONE
No aggregation (returns NULL).
Definition Aggregation.h:61
@ MIN
MIN → input type.
Definition Aggregation.h:54
@ CHOOSE
Arbitrary selection (pick one element).
Definition Aggregation.h:59
@ AVG
AVG → float.
Definition Aggregation.h:56
ComparisonOperator
SQL comparison operators used in gate_cmp circuit gates.
Definition Aggregation.h:39
@ LT
Less than (<).
Definition Aggregation.h:43
@ GT
Greater than (>).
Definition Aggregation.h:45
@ LE
Less than or equal (<=).
Definition Aggregation.h:42
@ NE
Not equal (<>).
Definition Aggregation.h:41
@ GE
Greater than or equal (>=).
Definition Aggregation.h:44
ValueType
Runtime type tag for aggregate values.
Definition Aggregation.h:98
@ INT
Signed 64-bit integer.
Definition Aggregation.h:99
@ STRING
Text string.
@ NONE
No value (NULL).
@ BOOLEAN
Boolean.
@ FLOAT
Double-precision float.
ArithmeticOperator
Arithmetic operations carried by gate_arith circuit gates.
Definition Aggregation.h:74
@ POW
binary power
Definition Aggregation.h:83
@ MAX
n-ary maximum (order statistic)
Definition Aggregation.h:81
@ DIV
binary quotient
Definition Aggregation.h:78
@ ROUND
rounding, to a number of digits given by a second wire
Definition Aggregation.h:87
@ PERCENTILE
continuous percentile over interleaved [indicator, value] wires
Definition Aggregation.h:86
@ FLOOR
unary floor
Definition Aggregation.h:88
@ CEIL
unary ceiling
Definition Aggregation.h:89
@ NEG
unary negation
Definition Aggregation.h:80
@ INTDIV
binary quotient truncated toward zero (integer division)
Definition Aggregation.h:79
@ ABS
unary absolute value
Definition Aggregation.h:90
@ EXP
unary exponential
Definition Aggregation.h:85
@ TIMES
n-ary product
Definition Aggregation.h:76
@ AS_FLOAT8
unary: the value as double precision reads it
Definition Aggregation.h:91
@ AS_FLOAT4
unary: the value as real reads it
Definition Aggregation.h:92
@ MIN
n-ary minimum (order statistic)
Definition Aggregation.h:82
@ LN
unary natural logarithm
Definition Aggregation.h:84
@ MINUS
binary difference
Definition Aggregation.h:77
Uniform error-reporting macros for ProvSQL.
#define PROVSQL_GAP
#define provsql_error(fmt,...)
Report a fatal ProvSQL error and abort the current transaction.
#define provsql_unsupported(scope, tag, fmt,...)
Refuse a query ProvSQL cannot track, and abort the transaction.
Core types, constants, and utilities shared across ProvSQL.
provsql_arith_op
Arithmetic operator tags used by gate_arith.
@ PROVSQL_ARITH_PERCENTILE
continuous percentile (order-statistic aggregate): wires are interleaved [ind_1, x_1,...
@ PROVSQL_ARITH_DIV
binary, child0 / child1
@ PROVSQL_ARITH_ASFLOAT8
unary, child0 as double precision
@ PROVSQL_ARITH_LN
unary, natural logarithm of child0 (a negative draw raises at evaluation)
@ PROVSQL_ARITH_ROUND
child0 rounded half away from zero, to child1 decimal digits where a second child is given (SQL round...
@ PROVSQL_ARITH_PLUS
n-ary, sum of children
@ PROVSQL_ARITH_POW
binary, child0 ^ child1 (real branch only: a negative base drawn with a non-integer exponent raises a...
@ PROVSQL_ARITH_ABS
unary, |child0|
@ PROVSQL_ARITH_FLOOR
unary, greatest integer <= child0
@ PROVSQL_ARITH_NEG
unary, -child0
@ PROVSQL_ARITH_INTDIV
binary, child0 / child1 truncated toward zero: SQL's division of two integers
@ PROVSQL_ARITH_MINUS
binary, child0 - child1
@ PROVSQL_ARITH_CEIL
unary, least integer >= child0
@ PROVSQL_ARITH_EXP
unary, e^child0
@ PROVSQL_ARITH_TIMES
n-ary, product of children
@ PROVSQL_ARITH_MIN
n-ary, min of children (order statistic; least / min aggregate)
@ PROVSQL_ARITH_MAX
n-ary, max of children (order statistic; greatest / max aggregate)
@ PROVSQL_ARITH_ASFLOAT4
unary, child0 as real reads it
A dynamically-typed aggregate value.
ValueType getType() const
Return the runtime type tag of this value.
std::variant< long, double, bool, std::string, std::vector< long >, std::vector< double >, std::vector< bool >, std::vector< std::string > > v
The variant holding the actual value.
Abstract interface for an incremental aggregate accumulator.
Aggregator implementing AVG; always returns a float result.
bool has
true once the first non-NULL input has been seen
ValueType resultType() const override
Return the type of the value returned by finalize().
ValueType inputType() const override
Return the type of the input values accepted by add().
unsigned count
Number of non-NULL inputs seen so far.
void add(const AggValue &x) override
Incorporate one input value into the running aggregate.
AggValue finalize() const override
Return the final aggregate result.
double sum
Running sum of all non-NULL input values.
Aggregator implementing CHOOSE (returns the first non-NULL input).
void add(const AggValue &x) override
Incorporate one input value into the running aggregate.
Aggregator implementing MAX for integer or float types.
void add(const AggValue &x) override
Incorporate one input value into the running aggregate.
Aggregator implementing MIN for integer or float types.
void add(const AggValue &x) override
Incorporate one input value into the running aggregate.
Base aggregator template for scalar types (int, float, bool, string).
AggValue finalize() const override
Return the accumulated value, or NULL if no inputs were seen.
ValueType inputType() const override
Return the value type corresponding to T.
bool has
true once the first non-NULL input has been seen
T value
Current accumulated value.
Aggregator implementing SUM for integer or float types.
void add(const AggValue &x) override
Incorporate one input value into the running aggregate.