ProvSQL C/C++ API
Adding support for provenance and uncertainty management to PostgreSQL databases
Loading...
Searching...
No Matches
CaseCmpExpander.cpp
Go to the documentation of this file.
1/**
2 * @file CaseCmpExpander.cpp
3 * @brief Implementation of @c runCaseCmpExpander. See @c CaseCmpExpander.h.
4 */
5#include "CaseCmpExpander.h"
6
7#include <string>
8#include <vector>
9
10#include "having_semantics.hpp" // extract_constant_string, map_cmp_op, parse_decimal_scaled
11
12extern "C" {
13#include "provsql_utils.h" // gate_type enum
14}
15
16namespace provsql {
17
18namespace {
19
20/// Whether @p op holds of a pair whose comparison has sign @p sign.
21bool op_holds(ComparisonOperator op, int sign)
22{
23 switch (op) {
24 case ComparisonOperator::EQ: return sign == 0;
25 case ComparisonOperator::NE: return sign != 0;
26 case ComparisonOperator::LT: return sign < 0;
27 case ComparisonOperator::LE: return sign <= 0;
28 case ComparisonOperator::GT: return sign > 0;
29 case ComparisonOperator::GE: return sign >= 0;
30 }
31 return false;
32}
33
34/**
35 * @brief The constant an operand carries, where it carries one.
36 *
37 * An arm of a lowered @c CASE is a bare @c gate_value; the other side of a
38 * comparison the rewriter builds is the @c semimod(𝟙, value) of a value of
39 * the row. @p is_null says the constant is SQL's NULL, which no comparison
40 * holds of.
41 */
42bool constant_of(GenericCircuit &gc, gate_t g, std::string &out, bool &is_null)
43{
44 const gate_type t = gc.getGateType(g);
45
46 if (t == gate_value)
47 out = gc.getExtra(g);
48 else if (t == gate_semimod) {
50 return false;
51 } else
52 return false;
53
54 is_null = out.empty() || out == "NULL" ||
55 gc.getUUID(g) == provsql_having_detail::GATE_NULL_UUID;
56 return true;
57}
58
59/// Decide a comparison between two constants, if both are decimal numbers.
60bool decide_constant_cmp(GenericCircuit &gc, gate_t l, gate_t r,
61 ComparisonOperator op, bool &holds)
62{
63 std::string ls, rs;
64 bool lnull = false, rnull = false;
65 long lm = 0, rm = 0;
66 int lscale = 0, rscale = 0, scale;
67
68 if (!constant_of(gc, l, ls, lnull) || !constant_of(gc, r, rs, rnull))
69 return false;
70 if (lnull || rnull) {
71 holds = false; /* unknown, and an unknown guard does not hold */
72 return true;
73 }
76 return false;
77 scale = lscale > rscale ? lscale : rscale;
78 if (!provsql_having_detail::rescale_to(lm, lscale, scale, lm) ||
79 !provsql_having_detail::rescale_to(rm, rscale, scale, rm))
80 return false;
81 holds = op_holds(op, lm < rm ? -1 : (lm > rm ? 1 : 0));
82 return true;
83}
84
85/**
86 * @brief Expand one comparison whose operand at @p side is a @c gate_case.
87 *
88 * The cmp gate itself becomes the @c gate_plus of the arm terms, keeping its
89 * id (and so every parent that reads it).
90 */
91void expand(GenericCircuit &gc, gate_t cmp, unsigned side, gate_t one)
92{
93 /* Copy before creating gates: the wire vectors may move. */
94 const std::vector<gate_t> cw = gc.getWires(cmp);
95 const gate_t selection = cw[side], other = cw[1 - side];
96 const std::vector<gate_t> aw = gc.getWires(selection);
97 const std::pair<unsigned, unsigned> infos = gc.getInfos(cmp);
98 const std::size_t n = aw.size(), arms = (n - 1) / 2;
99 bool okop = false;
101 std::vector<gate_t> terms;
102
103 for (std::size_t i = 0; i <= arms; ++i) {
104 const bool is_default = (i == arms);
105 const gate_t value = is_default ? aw[n - 1] : aw[2 * i + 1];
106 const gate_t l = (side == 0) ? value : other;
107 const gate_t r = (side == 0) ? other : value;
108 std::vector<gate_t> factors;
109 gate_t arm_cmp;
110 bool holds = false;
111
112 /* The comparison of this arm, in the operand order of the original. */
113 if (okop && decide_constant_cmp(gc, l, r, op, holds)) {
114 if (!holds)
115 continue; /* the term is 𝟘: the arm never satisfies the comparison */
116 arm_cmp = gc.addAnonymousGate(gate_one, {});
117 } else {
118 arm_cmp = gc.addAnonymousGate(gate_cmp, {l, r});
119 gc.setInfos(arm_cmp, infos.first, infos.second);
120 }
121
122 /* The arm is selected where its guard holds and none before it does. */
123 for (std::size_t j = 0; j < i; ++j)
124 factors.push_back(gc.addAnonymousGate(gate_monus, {one, aw[2 * j]}));
125 if (!is_default)
126 factors.push_back(aw[2 * i]);
127 factors.push_back(arm_cmp);
128 terms.push_back(factors.size() == 1
129 ? factors[0]
130 : gc.addAnonymousGate(gate_times, std::move(factors)));
131 }
132
133 if (terms.empty())
134 gc.resolveGateToZero(cmp); /* no arm ever satisfies it */
135 else
136 gc.resolveToPlus(cmp, std::move(terms));
137}
138
139} // namespace
140
142{
143 unsigned expanded = 0;
144 gate_t one{};
145 bool have_one = false;
146
147 /* A comparison between two guarded selections takes one round per side, a
148 * nested selection one per level; the bound is a guard against a circuit
149 * that would build them without end. */
150 for (unsigned round = 0; round < 32; ++round) {
151 std::vector<std::pair<gate_t, unsigned>> todo;
152 const auto nb = gc.getNbGates();
153
154 for (std::size_t i = 0; i < nb; ++i) {
155 const auto g = static_cast<gate_t>(i);
156 if (gc.getGateType(g) != gate_cmp)
157 continue;
158 const auto &w = gc.getWires(g);
159 if (w.size() != 2)
160 continue;
161 /* The wires of a gate_case are (guard, value)* followed by the
162 * default, so an odd number of them, at least one. */
163 for (unsigned side = 0; side < 2; ++side)
164 if (gc.getGateType(w[side]) == gate_case) {
165 const std::size_t n = gc.getWires(w[side]).size();
166 if (n >= 1 && n % 2 == 1)
167 todo.emplace_back(g, side);
168 break;
169 }
170 }
171 if (todo.empty())
172 break;
173
174 if (!have_one) {
175 one = gc.addAnonymousGate(gate_one, {});
176 have_one = true;
177 }
178 for (const auto &t : todo) {
179 if (gc.getGateType(t.first) != gate_cmp)
180 continue; /* defensive: already rewritten */
181 expand(gc, t.first, t.second, one);
182 ++expanded;
183 }
184 }
185 return expanded;
186}
187
188} // namespace provsql
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
Expansion of a comparison one of whose operands is a guarded selection (gate_case) into the compariso...
gate_t
Strongly-typed gate identifier.
Definition Circuit.h:49
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
uuid getUUID(gate_t g) const
Return the UUID string associated with gate g.
Definition Circuit.hpp:46
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.
void resolveToPlus(gate_t g, std::vector< gate_t > w)
Rewrite an arbitrary gate as a gate_plus over w.
void resolveGateToZero(gate_t g)
Replace an arbitrary gate (typically gate_times) by gate_zero.
void setInfos(gate_t g, unsigned info1, unsigned info2)
Set the integer annotation pair for gate g.
std::string getExtra(gate_t g) const
Return the string extra for gate g.
gate_t addAnonymousGate(gate_type type, std::vector< gate_t > wires_)
Allocate a fresh gate of type type over wires_.
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 parse_decimal_scaled(const std::string &s, long &mantissa, int &scale)
ComparisonOperator map_cmp_op(GenericCircuit &c, gate_t cmp_gate, bool &ok)
bool extract_constant_string(GenericCircuit &c, gate_t x, std::string &C_out)
unsigned runCaseCmpExpander(GenericCircuit &gc)
Expand every gate_cmp with a gate_case operand in gc.
Core types, constants, and utilities shared across ProvSQL.
@ gate_case
N-ary guarded selection over scalar (RV) children: wires are [guard_1, value_1, .....