ProvSQL C/C++ API
Adding support for provenance and uncertainty management to PostgreSQL databases
Loading...
Searching...
No Matches
CmpEvaluatorCommon.cpp
Go to the documentation of this file.
1/**
2 * @file CmpEvaluatorCommon.cpp
3 * @brief Implementation of the shared HAVING @c gate_cmp evaluator
4 * machinery. See @c CmpEvaluatorCommon.h.
5 */
7
8#include <algorithm>
9
10#include "having_semantics.hpp" // extract_constant_string, semimod_extract_string_and_K, map_cmp_op, flip_op
11extern "C" {
12#include "provsql_utils.h" // gate_type enum
13}
14
15namespace provsql {
16
18{
19 const auto &cw = gc.getWires(cmp);
20 if (cw.size() != 2) return false;
21
22 bool okop = false;
24 if (!okop) return false;
25
26 /* Identify the aggregate side; the other is the threshold constant.
27 * The reversed order (const compared to agg) calls for op flipping.
28 * The aggregate side is s * agg + d, s = ±1 and d a sum of constants,
29 * once its constant arithmetic is peeled. */
30 gate_t agg_side, const_side;
31 std::vector<gate_t> via;
32 std::vector<std::pair<int, std::string>> offsets; // d as signed constants
33 int sign = 1;
34 auto peel = [&](gate_t g) -> gate_t {
35 while (gc.getGateType(g) == gate_arith) {
36 const auto &w = gc.getWires(g);
37 const unsigned aop = static_cast<unsigned>(gc.getInfos(g).first);
38 gate_t inner{};
39 if (aop == PROVSQL_ARITH_PLUS) {
40 int non_const = 0;
41 for (gate_t ch : w)
42 if (gc.getGateType(ch) != gate_value) { inner = ch; ++non_const; }
43 if (non_const != 1) return g;
44 for (gate_t ch : w)
45 if (ch != inner) offsets.emplace_back(sign, gc.getExtra(ch));
46 } else if (aop == PROVSQL_ARITH_MINUS && w.size() == 2) {
47 if (gc.getGateType(w[1]) == gate_value) { // X - c
48 inner = w[0];
49 offsets.emplace_back(-sign, gc.getExtra(w[1]));
50 } else if (gc.getGateType(w[0]) == gate_value) { // c - X
51 inner = w[1];
52 offsets.emplace_back(sign, gc.getExtra(w[0]));
53 sign = -sign;
54 } else
55 return g;
56 } else if (aop == PROVSQL_ARITH_NEG && w.size() == 1) {
57 inner = w[0];
58 sign = -sign;
59 } else if ((aop == PROVSQL_ARITH_ASFLOAT8 ||
60 aop == PROVSQL_ARITH_ASFLOAT4) && w.size() == 1) {
61 /* The value read in its own type: the comparison is the same one on
62 * the value below, so a cast the query writes over an aggregate keeps
63 * the closed forms (sum(x)::float8 > 5 is sum(x) > 5). */
64 inner = w[0];
65 } else
66 return g;
67 via.push_back(g);
68 g = inner;
69 }
70 return g;
71 };
72 agg_side = peel(cw[0]);
73 const_side = cw[1];
74 if (gc.getGateType(agg_side) != gate_agg) {
75 via.clear(); offsets.clear(); sign = 1;
76 agg_side = peel(cw[1]);
77 const_side = cw[0];
78 if (gc.getGateType(agg_side) != gate_agg) return false;
80 }
81
82 /* The comparison domain is the aggregate result type (info2 of the
83 * gate_agg). The closed-form evaluators handle the ordered numeric
84 * domains (int / numeric / float) by scaling every value and the
85 * threshold to a common integer grid from their decimal text -- so a
86 * numeric(p,d) or finite-decimal float column is exact and fractional
87 * thresholds work. Text aggregates, exponential / non-decimal values,
88 * and grids too wide for a @c long are declined here and left to the
89 * enumeration path. */
90 const unsigned aggtype =
91 gc.getInfos(agg_side).second & PROVSQL_AGG_TYPE_MASK; // strip scalar flag
92 if (provsql_having_detail::aggtype_is_text(aggtype)) return false;
93
94 std::string c_str;
95 if (!provsql_having_detail::extract_constant_string(gc, const_side, c_str))
96 return false;
97 long c_mant = 0; int c_scale = 0;
98 if (!provsql_having_detail::parse_decimal_scaled(c_str, c_mant, c_scale))
99 return false;
100
101 /* No child: only a scalar aggregation over no row, whose value is that of
102 * the empty input (a window frame that may be empty while its row exists);
103 * a group without rows is no gate_agg. */
104 const auto &agg_children = gc.getWires(agg_side);
105 if (agg_children.empty() &&
106 !(gc.getInfos(agg_side).second & PROVSQL_AGG_SCALAR_FLAG))
107 return false;
108
109 std::vector<gate_t> semimods, ks;
110 std::vector<long> m_mant;
111 std::vector<int> m_scale;
112 semimods.reserve(agg_children.size());
113 ks.reserve(agg_children.size());
114 m_mant.reserve(agg_children.size());
115 m_scale.reserve(agg_children.size());
116
117 for (gate_t ch : agg_children) {
118 if (gc.getGateType(ch) != gate_semimod) return false;
119 std::string m_str;
120 gate_t k_gate{};
121 if (!provsql_having_detail::semimod_extract_string_and_K(gc, ch, m_str, k_gate))
122 return false;
123 long mm = 0; int sc = 0;
125 return false;
126 semimods.push_back(ch);
127 ks.push_back(k_gate);
128 m_mant.push_back(mm);
129 m_scale.push_back(sc);
130 }
131
132 /* The gate records the aggregate the query wrote: count(*) and count(expr)
133 * both keep COUNT even though their contributions are 1 and 0/1, so there is
134 * nothing to infer from the values here. */
135 AggregationOperator agg_kind =
136 getAggregationOperator(gc.getInfos(agg_side).first);
137
138 std::vector<long> d_mant(offsets.size());
139 std::vector<int> d_scale(offsets.size());
140 for (std::size_t i = 0; i < offsets.size(); ++i)
141 if (!provsql_having_detail::parse_decimal_scaled(offsets[i].second,
142 d_mant[i], d_scale[i]))
143 return false;
144
145 /* Rescale every value, the threshold and the offsets to a common integer
146 * grid. */
147 int target = c_scale;
148 for (int s : m_scale) target = std::max(target, s);
149 for (int s : d_scale) target = std::max(target, s);
150 long C = 0;
151 if (!provsql_having_detail::rescale_to(c_mant, c_scale, target, C)) return false;
152 std::vector<long> ms(m_mant.size());
153 for (std::size_t i = 0; i < m_mant.size(); ++i)
154 if (!provsql_having_detail::rescale_to(m_mant[i], m_scale[i], target, ms[i]))
155 return false;
156
157 /* s * agg + d op C is agg op C - d (s = 1), or agg flip(op) d - C. */
158 long d = 0;
159 for (std::size_t i = 0; i < offsets.size(); ++i) {
160 long v = 0;
161 if (!provsql_having_detail::rescale_to(d_mant[i], d_scale[i], target, v))
162 return false;
163 d += offsets[i].first * v;
164 }
165 if (sign > 0)
166 C -= d;
167 else {
168 C = d - C;
170 }
171
172 out.via = std::move(via);
173 out.agg = agg_side;
174 out.semimods = std::move(semimods);
175 out.ks = std::move(ks);
176 out.ms = std::move(ms);
177 out.agg_kind = agg_kind;
178 out.op = op;
179 out.C = C;
180 return true;
181}
182
183bool aggPrivateToCmp(const AggCmpMatch &match, const std::vector<unsigned> &ref)
184{
185 if (ref[static_cast<std::size_t>(match.agg)] != 1)
186 return false;
187 for (gate_t g : match.via)
188 if (ref[static_cast<std::size_t>(g)] != 1)
189 return false;
190 return true;
191}
192
193std::vector<unsigned> computeRefCounts(const GenericCircuit &gc)
194{
195 const auto nb = gc.getNbGates();
196 std::vector<unsigned> ref(nb, 0);
197 for (std::size_t i = 0; i < nb; ++i) {
198 auto g = static_cast<gate_t>(i);
199 for (gate_t w : gc.getWires(g)) {
200 const auto idx = static_cast<std::size_t>(w);
201 if (idx < ref.size()) ++ref[idx];
202 }
203 }
204 return ref;
205}
206
208 const std::vector<unsigned> &ref, bool &ok)
209{
210 switch (gc.getGateType(g)) {
211 case gate_one: return 1.0;
212 case gate_zero: return 0.0;
213 case gate_input:
214 if (ref[static_cast<std::size_t>(g)] != 1) { ok = false; return 0.0; }
215 return gc.getProb(g);
216 case gate_times: {
217 if (ref[static_cast<std::size_t>(g)] != 1) { ok = false; return 0.0; }
218 double pr = 1.0;
219 for (gate_t c : gc.getWires(g)) {
220 pr *= contributorProb(gc, c, ref, ok);
221 if (!ok) return 0.0;
222 }
223 return pr;
224 }
225 case gate_plus: {
226 if (ref[static_cast<std::size_t>(g)] != 1) { ok = false; return 0.0; }
227 double q = 1.0;
228 for (gate_t c : gc.getWires(g)) {
229 q *= (1.0 - contributorProb(gc, c, ref, ok));
230 if (!ok) return 0.0;
231 }
232 return 1.0 - q;
233 }
234 case gate_monus: {
235 /* a (-) b = a AND NOT b ; with disjoint private leaves a and b are
236 * independent, so Pr = Pr(a) * (1 - Pr(b)). Children are
237 * [minuend, subtrahend] (see GenericCircuit evaluate<S>). */
238 if (ref[static_cast<std::size_t>(g)] != 1) { ok = false; return 0.0; }
239 const auto &w = gc.getWires(g);
240 if (w.size() != 2) { ok = false; return 0.0; }
241 double pa = contributorProb(gc, w[0], ref, ok);
242 if (!ok) return 0.0;
243 double pb = contributorProb(gc, w[1], ref, ok);
244 if (!ok) return 0.0;
245 return pa * (1.0 - pb);
246 }
247 default:
248 ok = false;
249 return 0.0;
250 }
251}
252
253} // namespace provsql
AggregationOperator getAggregationOperator(Oid oid)
Map a PostgreSQL aggregate function OID to an AggregationOperator.
AggregationOperator
SQL aggregation functions tracked by ProvSQL.
Definition Aggregation.h:51
ComparisonOperator
SQL comparison operators used in gate_cmp circuit gates.
Definition Aggregation.h:39
gate_t
Strongly-typed gate identifier.
Definition Circuit.h:49
Shared machinery for the closed-form HAVING gate_cmp probability evaluators (Poisson-binomial COUNT,...
std::vector< gate_t > & getWires(gate_t g)
Return a mutable reference to the child-wire list of gate g.
Definition Circuit.h:140
gateType getGateType(gate_t g) const
Return the type of gate g.
Definition Circuit.h:130
std::vector< gate_t >::size_type getNbGates() const
Return the total number of gates in the circuit.
Definition Circuit.h:103
In-memory provenance circuit with semiring-generic evaluation.
std::string getExtra(gate_t g) const
Return the string extra for gate g.
double getProb(gate_t g) const
Return the probability for gate g.
std::pair< unsigned, unsigned > getInfos(gate_t g) const
Return the integer annotation pair for gate g.
Provenance evaluation helper for HAVING-clause circuits.
bool rescale_to(long mantissa, int scale, int target_scale, long &out)
bool aggtype_is_text(unsigned oid)
ComparisonOperator flip_op(ComparisonOperator op)
bool parse_decimal_scaled(const std::string &s, long &mantissa, int &scale)
ComparisonOperator map_cmp_op(GenericCircuit &c, gate_t cmp_gate, bool &ok)
bool semimod_extract_string_and_K(GenericCircuit &c, gate_t semimod_gate, std::string &m_out, gate_t &k_gate_out)
bool extract_constant_string(GenericCircuit &c, gate_t x, std::string &C_out)
bool aggPrivateToCmp(const AggCmpMatch &match, const std::vector< unsigned > &ref)
Whether the aggregate of match is consumed by its comparison alone: agg and every gate of via has ref...
std::vector< unsigned > computeRefCounts(const GenericCircuit &gc)
Reference count of every gate as a wire-target across the whole circuit.
bool matchAggCmp(GenericCircuit &gc, gate_t cmp, AggCmpMatch &out)
Try to match cmp against gate_cmp(gate_agg(α, semimod_i(K_i, m_i)*), gate_value(C)).
double contributorProb(const GenericCircuit &gc, gate_t g, const std::vector< unsigned > &ref, bool &ok)
Read-once marginal probability of a count/aggregate contributor (the K side of a semimod).
Core types, constants, and utilities shared across ProvSQL.
@ PROVSQL_ARITH_ASFLOAT8
unary, child0 as double precision
@ PROVSQL_ARITH_PLUS
n-ary, sum of children
@ PROVSQL_ARITH_NEG
unary, -child0
@ PROVSQL_ARITH_MINUS
binary, child0 - child1
@ PROVSQL_ARITH_ASFLOAT4
unary, child0 as real reads it
#define PROVSQL_AGG_TYPE_MASK
@ gate_arith
n-ary arithmetic gate over scalar-valued children (info1 holds operator tag)
#define PROVSQL_AGG_SCALAR_FLAG
Scalar-aggregation flag, stored in the upper bit of a gate_agg's info2 (whose low 31 bits hold the ag...
Result of matching a gate_cmp against the canonical HAVING aggregate-comparison shape.
gate_t agg
the gate_agg operand of the cmp
long C
the constant threshold, on the same integer grid as ms
std::vector< gate_t > ks
the K side of each semimod (contributor root)
std::vector< gate_t > semimods
the per-child gate_semimod parents
std::vector< gate_t > via
the gate_arith gates of constant arithmetic between the cmp and agg, folded into op and C
std::vector< long > ms
the M side of each semimod (per-row value), scaled to a common integer grid (numeric / decimal-float ...
AggregationOperator agg_kind
effective aggregate (SUM-of-1s remapped to COUNT)
ComparisonOperator op
comparator, flipped if the agg sits on the right