ProvSQL C/C++ API
Adding support for provenance and uncertainty management to PostgreSQL databases
Loading...
Searching...
No Matches
RangeCheck.cpp
Go to the documentation of this file.
1/**
2 * @file RangeCheck.cpp
3 * @brief Implementation of the support-based bound check pass.
4 * See @c RangeCheck.h for the full docstring.
5 */
6#include "RangeCheck.h"
7
8#include <algorithm>
9#include <cmath>
10#include <limits>
11#include <stack>
12#include <unordered_map>
13#include <unordered_set>
14#include <vector>
15
16#include "Aggregation.h" // ComparisonOperator + cmpOpFromOid
17#include "AnalyticEvaluator.h" // cdfAt for shape_mass under truncation
18#include "CircuitFromMMap.h" // getGenericCircuit
19#include "ConjugatePosterior.h" // conjugatePosterior (observe-evidence shapes)
20#include "Expectation.h" // lift_conditioning
21#include "RandomVariable.h" // parse_distribution_spec
22#include "distributions/Distribution.h" // makeDistribution -> per-family support()
23#include "provsql_utils_cpp.h" // uuid2string
24
25#include <type_traits> // std::is_same_v in truncateShape
26#include <variant>
27extern "C" {
28#include "postgres.h"
29#include "fmgr.h"
30#include "funcapi.h" // get_call_result_type, BlessTupleDesc
31#include "access/htup_details.h" // heap_form_tuple
32#include "utils/uuid.h"
33#include "provsql_utils.h" // gate_type, provsql_arith_op
34#include "provsql_error.h"
35
36PG_FUNCTION_INFO_V1(rv_support);
37}
38
39namespace provsql {
40
41namespace {
42
43/**
44 * @brief Closed interval @c [lo, hi] on the extended real line.
45 *
46 * @c -INFINITY / @c +INFINITY are used for unbounded ends (e.g. the
47 * support of a normal RV is @c {-INF, +INF}). Empty intervals are
48 * not generated by any constructor below; comparators against an
49 * empty interval would be vacuous and we consider them undecidable.
50 */
51struct Interval {
52 double lo;
53 double hi;
54
55 static Interval point(double v) { return {v, v}; }
56 static Interval all() { return {-std::numeric_limits<double>::infinity(),
57 +std::numeric_limits<double>::infinity()}; }
58 bool isAll() const {
59 return std::isinf(lo) && lo < 0 && std::isinf(hi) && hi > 0;
60 }
61};
62
63Interval add(Interval a, Interval b) { return {a.lo + b.lo, a.hi + b.hi}; }
64Interval sub(Interval a, Interval b) { return {a.lo - b.hi, a.hi - b.lo}; }
65Interval neg(Interval a) { return {-a.hi, -a.lo}; }
66
67/* Interval product: take the min/max of the four corner products.
68 * Handles signed bounds correctly (no special case for negative). */
69Interval mul(Interval a, Interval b)
70{
71 double p1 = a.lo * b.lo, p2 = a.lo * b.hi;
72 double p3 = a.hi * b.lo, p4 = a.hi * b.hi;
73 return {std::min({p1, p2, p3, p4}), std::max({p1, p2, p3, p4})};
74}
75
76/* Interval division: if the divisor straddles zero, the result is
77 * unbounded in both directions; otherwise compute via @c mul(a, 1/b).
78 * The conservative all-real fallback is correct (any real value is
79 * possible) but throws away precision &ndash; division by an interval
80 * crossing zero is rare in our tests. */
81/* exp is monotone increasing and total (exp(-inf) = 0, exp(inf) = inf). */
82Interval expInt(Interval a) { return {std::exp(a.lo), std::exp(a.hi)}; }
83
84/* ln is monotone increasing on [0, inf); the part of an interval below
85 * 0 is a domain violation the sampler raises on, so the bound covers
86 * the draws that do evaluate (lo <= 0 maps to -inf via ln 0). */
87Interval lnInt(Interval a)
88{
89 const double lo = a.lo > 0.0 ? std::log(a.lo)
90 : -std::numeric_limits<double>::infinity();
91 const double hi = a.hi > 0.0 ? std::log(a.hi)
92 : -std::numeric_limits<double>::infinity();
93 return {lo, hi};
94}
95
96/* x^y over the interval box. For a base interval entirely >= 0, x^y is
97 * monotone in each variable separately (in x for fixed y, in y for
98 * fixed x), so the extrema sit at the corners; 0^negative diverges to
99 * +inf, which std::pow reports directly. A base interval extending
100 * below 0 keeps the conservative all-real bound: integer-exponent draws
101 * are legitimate there, and non-integer ones raise in the sampler. */
102Interval powInt(Interval b, Interval e)
103{
104 if (!(b.lo >= 0.0))
105 return Interval::all();
106 double lo = std::numeric_limits<double>::infinity();
107 double hi = -std::numeric_limits<double>::infinity();
108 for (double x : {b.lo, b.hi})
109 for (double y : {e.lo, e.hi}) {
110 const double v = std::pow(x, y);
111 if (std::isnan(v))
112 return Interval::all();
113 lo = std::min(lo, v);
114 hi = std::max(hi, v);
115 }
116 return {lo, hi};
117}
118
119Interval divInt(Interval a, Interval b)
120{
121 if (b.lo <= 0.0 && b.hi >= 0.0)
122 return Interval::all();
123 Interval inv = {1.0 / b.hi, 1.0 / b.lo};
124 return mul(a, inv);
125}
126
127/**
128 * @brief Recursively compute the interval of @p g's value across
129 * worlds. Memoised in @p cache.
130 *
131 * Recognised gate types:
132 * - @c gate_value: point interval on the parsed scalar.
133 * - @c gate_rv: distribution support (uniform exact, exponential
134 * on @c [0, +∞), normal on @c (-∞, +∞)).
135 * - @c gate_arith: propagated via the interval-arith helpers above.
136 *
137 * Anything else (e.g. an aggregate gate reached via a HAVING cmp)
138 * yields the all-real interval, which downstream conservatively
139 * treats as undecidable.
140 */
141Interval intervalOf(const GenericCircuit &gc, gate_t g,
142 std::unordered_map<gate_t, Interval> &cache)
143{
144 auto it = cache.find(g);
145 if (it != cache.end()) return it->second;
146
147 Interval result = Interval::all();
148 auto type = gc.getGateType(g);
149
150 switch (type) {
151 case gate_value:
152 /* A value RangeCheck cannot read as a double (e.g. a text constant
153 * from an agg_token = text comparison) carries no numeric interval
154 * constraint: leave it unconstrained rather than aborting the whole
155 * load-time pass. Mirrors the "undecidable -> all()" default. */
156 try {
157 result = Interval::point(parseDoubleStrict(gc.getExtra(g)));
158 } catch (const CircuitException &) {
159 result = Interval::all();
160 }
161 break;
162 case gate_rv: {
163 auto spec = parse_distribution_spec(gc.getExtra(g));
164 if (!spec) break;
165 // Natural support per family (Normal ℝ, Uniform [a,b], Exp/Erlang [0,∞)).
166 const DistSupport s = makeDistribution(*spec)->support();
167 result = {s.lo, s.hi};
168 break;
169 }
170 case gate_arith: {
171 auto op = static_cast<provsql_arith_op>(gc.getInfos(g).first);
172 const auto &wires = gc.getWires(g);
173 if (wires.empty()) break;
174 Interval first = intervalOf(gc, wires[0], cache);
175 switch (op) {
177 result = first;
178 for (std::size_t i = 1; i < wires.size(); ++i)
179 result = add(result, intervalOf(gc, wires[i], cache));
180 break;
182 result = first;
183 for (std::size_t i = 1; i < wires.size(); ++i)
184 result = mul(result, intervalOf(gc, wires[i], cache));
185 break;
187 if (wires.size() != 2) break;
188 result = sub(first, intervalOf(gc, wires[1], cache));
189 break;
191 if (wires.size() != 2) break;
192 result = divInt(first, intervalOf(gc, wires[1], cache));
193 break;
195 /* Truncation toward zero is monotone: truncate the bounds. */
196 if (wires.size() != 2) break;
197 result = divInt(first, intervalOf(gc, wires[1], cache));
198 if (std::isfinite(result.lo)) result.lo = std::trunc(result.lo);
199 if (std::isfinite(result.hi)) result.hi = std::trunc(result.hi);
200 break;
202 if (wires.size() != 1) break;
203 result = neg(first);
204 break;
206 /* max/min are monotone in each argument, so the support propagates
207 * directly and soundly (independent of any correlation between the
208 * children): [max(lo_i), max(hi_i)] resp. [min(lo_i), min(hi_i)]. */
209 result = first;
210 for (std::size_t i = 1; i < wires.size(); ++i) {
211 Interval o = intervalOf(gc, wires[i], cache);
212 result = { std::max(result.lo, o.lo), std::max(result.hi, o.hi) };
213 }
214 break;
216 result = first;
217 for (std::size_t i = 1; i < wires.size(); ++i) {
218 Interval o = intervalOf(gc, wires[i], cache);
219 result = { std::min(result.lo, o.lo), std::min(result.hi, o.hi) };
220 }
221 break;
223 if (wires.size() != 2) break;
224 result = powInt(first, intervalOf(gc, wires[1], cache));
225 break;
226 case PROVSQL_ARITH_LN:
227 if (wires.size() != 1) break;
228 result = lnInt(first);
229 break;
231 if (wires.size() != 1) break;
232 result = expInt(first);
233 break;
235 /* Rounding is monotone, so the bounds carry over; the digits, if a
236 * second wire gives them, are a constant we do not read here, and
237 * widening by one unit in the last place of the bounds is sound for
238 * any of them. */
239 result = { std::floor(first.lo), std::ceil(first.hi) };
240 break;
242 if (wires.size() != 1) break;
243 result = { std::floor(first.lo), std::floor(first.hi) };
244 break;
246 if (wires.size() != 1) break;
247 result = { std::ceil(first.lo), std::ceil(first.hi) };
248 break;
250 /* The value in its own type: a double carries it as it is, and
251 * rounding to a real is monotone, so the interval follows. */
252 if (wires.size() != 1) break;
253 result = first;
254 break;
256 if (wires.size() != 1) break;
257 result = { static_cast<double>(static_cast<float>(first.lo)),
258 static_cast<double>(static_cast<float>(first.hi)) };
259 break;
261 /* |x| is not monotone: an interval straddling zero has 0 as its
262 * least value, and the greatest is the farther endpoint. */
263 if (wires.size() != 1) break;
264 if (first.lo >= 0.0)
265 result = first;
266 else if (first.hi <= 0.0)
267 result = { -first.hi, -first.lo };
268 else
269 result = { 0.0, std::max(-first.lo, first.hi) };
270 break;
272 /* Continuous percentile over interleaved [ind, x, ...] wires:
273 * every draw interpolates within the present values, so the hull
274 * of the value wires' supports is a sound bound regardless of
275 * which subset is present. */
276 if (wires.size() < 2 || wires.size() % 2 != 0) break;
277 result = intervalOf(gc, wires[1], cache);
278 for (std::size_t i = 3; i < wires.size(); i += 2) {
279 Interval o = intervalOf(gc, wires[i], cache);
280 result = { std::min(result.lo, o.lo), std::max(result.hi, o.hi) };
281 }
282 break;
283 }
284 break;
285 }
286 case gate_semimod: {
287 /* HAVING-style constant wrapper: semimod(gate_one, value). The
288 * semiring action of gate_one (always true) on a scalar leaves
289 * the scalar unchanged in every world, so the interval of the
290 * semimod equals the interval of its value child. Other
291 * semimod shapes (non-trivial k_gate) keep the conservative
292 * all-real default here; @c support_intervalOf widens it to
293 * a sampling-time support for rv_* callers. */
294 const auto &wires = gc.getWires(g);
295 if (wires.size() == 2 && gc.getGateType(wires[0]) == gate_one)
296 result = intervalOf(gc, wires[1], cache);
297 break;
298 }
299 case gate_mixture: {
300 /* Support of a mixture is the union of its branch supports.
301 * Two shapes:
302 * - Classic 3-wire [p_token, x_token, y_token]: the Bernoulli
303 * is a Boolean leaf and contributes nothing to the scalar
304 * interval.
305 * - Categorical N-wire [key, mul_1, ..., mul_n]: each mulinput
306 * carries its outcome value in extra; the support is the
307 * [min, max] of those values. */
308 const auto &wires = gc.getWires(g);
309 if (gc.isCategoricalMixture(g)) {
310 double lo = std::numeric_limits<double>::infinity();
311 double hi = -std::numeric_limits<double>::infinity();
312 bool any = false;
313 for (std::size_t i = 1; i < wires.size(); ++i) {
314 double v;
315 try { v = parseDoubleStrict(gc.getExtra(wires[i])); }
316 catch (const CircuitException &) { any = false; break; }
317 lo = std::min(lo, v);
318 hi = std::max(hi, v);
319 any = true;
320 }
321 if (any) result = {lo, hi};
322 } else if (wires.size() == 3) {
323 Interval ix = intervalOf(gc, wires[1], cache);
324 Interval iy = intervalOf(gc, wires[2], cache);
325 result = {std::min(ix.lo, iy.lo), std::max(ix.hi, iy.hi)};
326 }
327 break;
328 }
329 case gate_case: {
330 /* Guarded selection [g_1, v_1, ..., g_k, v_k, default]: the result is
331 * always one of the value branches, so the support is the union of the
332 * values' supports (the guards only decide which one). Values are the
333 * odd-indexed wires plus the final default. */
334 const auto &wires = gc.getWires(g);
335 if (wires.empty()) break;
336 result = intervalOf(gc, wires.back(), cache); /* default */
337 for (std::size_t i = 1; i + 1 < wires.size(); i += 2) {
338 Interval o = intervalOf(gc, wires[i], cache);
339 result = { std::min(result.lo, o.lo), std::max(result.hi, o.hi) };
340 }
341 break;
342 }
343 default:
344 /* gate_agg is intentionally not handled here -- the empty-subset
345 * NULL semantics make a flat interval misleading, so the
346 * runRangeCheck loop dispatches agg-bearing cmps to a separate
347 * decider that knows the asymmetry between sound FALSE and
348 * unsound TRUE decisions for SUM / MIN / MAX. All other gate
349 * types fall through to the all-real default. */
350 break;
351 }
352
353 cache[g] = result;
354 return result;
355}
356
357/**
358 * @brief Decide a @c gate_cmp from the interval of @c (lhs - rhs).
359 *
360 * Returns @c NaN when the comparator cannot be decided from interval
361 * bounds alone (e.g. the difference straddles zero, or the comparator
362 * is @c = / @c <> on overlapping continuous supports &ndash; both of
363 * which need a CDF, which a downstream analytic pass can supply).
364 * Otherwise returns the certain probability @c 0.0 or @c 1.0.
365 */
366double decideCmp(const Interval &diff, ComparisonOperator op)
367{
368 switch (op) {
370 if (diff.hi < 0.0) return 1.0;
371 if (diff.lo >= 0.0) return 0.0;
372 break;
374 if (diff.hi <= 0.0) return 1.0;
375 if (diff.lo > 0.0) return 0.0;
376 break;
378 if (diff.lo > 0.0) return 1.0;
379 if (diff.hi <= 0.0) return 0.0;
380 break;
382 if (diff.lo >= 0.0) return 1.0;
383 if (diff.hi < 0.0) return 0.0;
384 break;
386 /* Disjoint supports ⇒ certainly false. Overlapping supports of
387 * continuous RVs would have probability zero in the measure-
388 * theoretic sense, but the interval pass alone cannot tell
389 * whether either side is continuous; leave that to a downstream
390 * analytic-CDF pass when one is available. */
391 if (diff.hi < 0.0 || diff.lo > 0.0) return 0.0;
392 break;
394 if (diff.hi < 0.0 || diff.lo > 0.0) return 1.0;
395 break;
396 }
397 return std::numeric_limits<double>::quiet_NaN();
398}
399
400/**
401 * @brief Decide a @c gate_cmp where one side is a @c gate_agg, the
402 * other is a scalar constant.
403 *
404 * Computes a value-interval for the aggregate from its semimod
405 * children's per-row values, then folds the comparator like the
406 * non-agg path &ndash; but accepts only FALSE decisions, never
407 * TRUE. The reason is structural to ProvSQL's HAVING semantics:
408 * the per-aggregator subset enumerators in @c subset.cpp
409 * (@c count_enum, @c sum_dp, @c enumerate_exhaustive) all skip
410 * the empty subset, matching SQL's "no group, no HAVING" rule.
411 * So a HAVING cmp's value is the OR over the @em non-empty subsets
412 * where the predicate holds.
413 *
414 * - When no non-empty subset satisfies the predicate (the bound is
415 * strictly disjoint from the threshold on the right side of the
416 * comparator), the cmp value is exactly @c 0 = @c gate_zero.
417 * FALSE decision: sound.
418 * - When every non-empty subset satisfies the predicate, the cmp
419 * value equals "the group is non-empty" &ndash; the OR over the
420 * children's k_gates &ndash; which is a non-constant Boolean
421 * expression, @em not @c gate_one. Returning TRUE here would
422 * replace the cmp with @c gate_one and over-count probability
423 * mass from the empty world (where the group does not exist),
424 * so TRUE decisions are blocked uniformly across all aggregators.
425 *
426 * Aggregators we don't bound (@c AVG, @c AND, @c OR, @c CHOOSE,
427 * @c ARRAY_AGG, @c NONE) fall through to undecidable.
428 *
429 * @return @c 0.0 if decided to FALSE, @c NaN otherwise.
430 */
431double decideAggVsConstCmp(const GenericCircuit &gc, gate_t agg_gate,
432 ComparisonOperator op, double const_val,
433 bool agg_on_lhs,
434 bool *out_always_true = nullptr)
435{
436 AggregationOperator aop = getAggregationOperator(gc.getInfos(agg_gate).first);
437
438 /* Extract per-child scalar values from the semimod children. */
439 std::vector<double> values;
440 for (gate_t child : gc.getWires(agg_gate)) {
441 if (gc.getGateType(child) != gate_semimod)
442 return std::numeric_limits<double>::quiet_NaN();
443 const auto &sw = gc.getWires(child);
444 if (sw.size() != 2)
445 return std::numeric_limits<double>::quiet_NaN();
446 gate_t value_gate = sw[1];
448 return std::numeric_limits<double>::quiet_NaN();
449 try {
450 values.push_back(parseDoubleStrict(gc.getExtra(value_gate)));
451 } catch (const CircuitException &) {
452 return std::numeric_limits<double>::quiet_NaN();
453 }
454 }
455
456 Interval val_interval = Interval::all();
457
458 switch (aop) {
460 val_interval = {0.0, static_cast<double>(values.size())};
461 break;
463 double sum_neg = 0.0, sum_pos = 0.0;
464 for (double v : values) {
465 if (v < 0.0) sum_neg += v;
466 else sum_pos += v;
467 }
468 val_interval = {std::min(0.0, sum_neg), std::max(0.0, sum_pos)};
469 break;
470 }
473 if (values.empty())
474 return std::numeric_limits<double>::quiet_NaN();
475 val_interval = {*std::min_element(values.begin(), values.end()),
476 *std::max_element(values.begin(), values.end())};
477 break;
478 default:
479 /* AVG / AND / OR / CHOOSE / ARRAY_AGG / NONE: not decidable
480 * with this pass. */
481 return std::numeric_limits<double>::quiet_NaN();
482 }
483
484 Interval lhs = agg_on_lhs ? val_interval : Interval::point(const_val);
485 Interval rhs = agg_on_lhs ? Interval::point(const_val) : val_interval;
486 Interval diff = sub(lhs, rhs);
487 double p = decideCmp(diff, op);
488
489 /* Only FALSE decisions are universally sound (see doc comment).
490 * gate_one would over-credit the empty subset, which provsql_having
491 * deliberately excludes from valid worlds; the safe TRUE rewrite
492 * is "the group is non-empty" = OR over the agg's K-gates, and is
493 * sound only in absorptive semirings. The TRUE signal is therefore
494 * reported via @p out_always_true rather than the return value;
495 * the universal load-time @c runRangeCheck caller ignores it,
496 * while the probability-side @c runHavingAlwaysTrueRewriter caller
497 * acts on it. */
498 if (p == 0.0) return 0.0;
499 if (p == 1.0 && out_always_true != nullptr) *out_always_true = true;
500 return std::numeric_limits<double>::quiet_NaN();
501}
502
503/**
504 * @brief Try to extract a scalar constant from a cmp's child.
505 *
506 * Recognises two shapes:
507 * - bare @c gate_value: parse its @c extra as a double;
508 * - HAVING-style @c gate_semimod with @c k=gate_one and
509 * @c value=gate_value: parse the value's extra.
510 *
511 * Returns @c NaN on any other shape.
512 */
513double extractScalarConst(const GenericCircuit &gc, gate_t g)
514{
515 auto t = gc.getGateType(g);
516 if (t == gate_value) {
517 try { return parseDoubleStrict(gc.getExtra(g)); }
518 catch (const CircuitException &) {
519 return std::numeric_limits<double>::quiet_NaN();
520 }
521 }
522 if (t == gate_semimod) {
523 const auto &w = gc.getWires(g);
524 if (w.size() != 2) return std::numeric_limits<double>::quiet_NaN();
525 if (gc.getGateType(w[0]) != gate_one)
526 return std::numeric_limits<double>::quiet_NaN();
527 if (gc.getGateType(w[1]) != gate_value)
528 return std::numeric_limits<double>::quiet_NaN();
529 try { return parseDoubleStrict(gc.getExtra(w[1])); }
530 catch (const CircuitException &) {
531 return std::numeric_limits<double>::quiet_NaN();
532 }
533 }
534 return std::numeric_limits<double>::quiet_NaN();
535}
536
537/**
538 * @brief Flip the sides of a comparison operator.
539 *
540 * @c (a op b) is equivalent to @c (b flip(op) a). Used to normalise
541 * a cmp so the random-variable side is always on the left.
542 */
544{
545 switch (op) {
552 }
553 return op;
554}
555
556/**
557 * @brief Interpret a @c gate_cmp as a per-RV constraint @c rv op c.
558 *
559 * Returns @c true and fills @p rv_out, @p op_out, @p const_out when
560 * exactly one side of the cmp is a @c gate_rv and the other a
561 * @c gate_value with a parseable scalar; @c false otherwise (both
562 * sides are RVs, both constants, an @c arith subtree appears, etc.).
563 *
564 * Strict-vs-non-strict inequalities are preserved as the operator;
565 * the caller decides whether to treat the boundary as inclusive
566 * (continuous distributions: measure-zero, irrelevant for
567 * feasibility verdicts).
568 */
569bool asRvVsConstCmp(const GenericCircuit &gc, gate_t cmp_gate,
570 gate_t &rv_out, ComparisonOperator &op_out,
571 double &const_out)
572{
573 bool ok = false;
574 ComparisonOperator op = cmpOpFromOid(gc.getInfos(cmp_gate).first, ok);
575 if (!ok) return false;
576 const auto &wires = gc.getWires(cmp_gate);
577 if (wires.size() != 2) return false;
578
579 /* Recognise scalar-vs-constant cmps where the scalar side is a
580 * bare gate_rv (the original use case for the per-cmp resolution
581 * pass) or a gate_mixture (so the conditioning walker can extract
582 * intervals on mixture / categorical variables – value-vs-value
583 * cmps are folded upstream by RangeCheck before they reach this
584 * walker). Dirac (gate_value) is never the scalar side of a
585 * non-trivial cmp at this point; the value-vs-value pair would have
586 * been resolved upstream. */
587 auto isScalarRv = [](gate_type t) {
588 return t == gate_rv || t == gate_mixture;
589 };
590 auto t0 = gc.getGateType(wires[0]);
591 auto t1 = gc.getGateType(wires[1]);
592 if (isScalarRv(t0) && t1 == gate_value) {
593 try { const_out = parseDoubleStrict(gc.getExtra(wires[1])); }
594 catch (const CircuitException &) { return false; }
595 rv_out = wires[0];
596 op_out = op;
597 return true;
598 }
599 if (t0 == gate_value && isScalarRv(t1)) {
600 try { const_out = parseDoubleStrict(gc.getExtra(wires[0])); }
601 catch (const CircuitException &) { return false; }
602 rv_out = wires[1];
603 op_out = flipCmpOp(op);
604 return true;
605 }
606 return false;
607}
608
609/**
610 * @brief Apply a single @c rv-op-constant constraint to a running
611 * interval for the RV.
612 *
613 * Strict vs non-strict inequalities collapse onto the same closed
614 * interval: continuous distributions assign zero mass to the
615 * boundary, so the joint-feasibility verdict is unchanged whether
616 * we use @c < or @c <=. @c <> (NE) cannot be represented as a
617 * single interval and is left to the per-cmp pass.
618 */
619Interval intersectRvConstraint(Interval current, ComparisonOperator op,
620 double c)
621{
622 switch (op) {
625 current.hi = std::min(current.hi, c);
626 break;
629 current.lo = std::max(current.lo, c);
630 break;
632 current.lo = std::max(current.lo, c);
633 current.hi = std::min(current.hi, c);
634 break;
636 /* Cannot represent the complement of a point as a single
637 * interval; leave the running interval unchanged. */
638 break;
639 }
640 return current;
641}
642
643bool intervalEmpty(Interval i) { return i.lo > i.hi; }
644
645/**
646 * @brief Walk an AND-conjunct tree collecting per-RV interval
647 * constraints from its @c gate_cmp leaves.
648 *
649 * Shared between @c isAndJointlyInfeasible (which checks for an empty
650 * intersection) and the public @c collectRvConstraints / conditional
651 * @c compute_support paths. Descends through @c gate_times,
652 * collecting every @c gate_cmp interpretable as `rv op const` and
653 * intersecting its constraint into a running interval for that RV.
654 *
655 * @p complete is set to @c true on entry and cleared if the walk
656 * encounters any structure other than the AND-friendly set
657 * (@c gate_times, @c gate_cmp, @c gate_input, @c gate_one,
658 * @c gate_zero) whose footprint *might* constrain an RV
659 * (i.e. excluding bare Bernoulli factors). Callers that need a
660 * tight bound (the closed-form moment shortcut) must check it; the
661 * support intersection caller can use the result unconditionally
662 * because dropping a disjunctive factor only loosens the interval,
663 * which is sound for a superset bound on the conditional support.
664 *
665 * Cmps that do not interpret as `rv op const` (RV vs RV, arith on
666 * either side, agg…) are silently ignored; they belong to the
667 * conditioning event but don't constrain a single RV's interval.
668 */
669void walkAndConjunctIntervals(
670 const GenericCircuit &gc, gate_t root,
671 std::unordered_map<gate_t, Interval> &rv_intervals,
672 std::unordered_map<gate_t, Interval> &support_cache,
673 bool &complete)
674{
675 std::unordered_set<gate_t> seen;
676 std::stack<gate_t> stk;
677 stk.push(root);
678 complete = true;
679
680 while (!stk.empty()) {
681 gate_t g = stk.top(); stk.pop();
682 if (!seen.insert(g).second) continue;
683
684 auto t = gc.getGateType(g);
685 if (t == gate_cmp) {
686 gate_t rv = static_cast<gate_t>(0);
688 double c = 0.0;
689 if (!asRvVsConstCmp(gc, g, rv, op, c)) {
690 /* Cmp shape we don't interpret (RV vs RV, arith involved).
691 * Conservatively mark the walk incomplete: this cmp belongs
692 * to the event AND could constrain an RV in a way we can't
693 * fold into a single interval. */
694 complete = false;
695 continue;
696 }
697 auto it = rv_intervals.find(rv);
698 Interval current = (it == rv_intervals.end())
699 ? intervalOf(gc, rv, support_cache)
700 : it->second;
701 current = intersectRvConstraint(current, op, c);
702 rv_intervals[rv] = current;
703 continue; /* never descend into a cmp's operands */
704 }
705 if (t == gate_times || t == gate_delta || g == root) {
706 /* gate_delta wraps a single child as the δ-semiring identity on
707 * Booleans, so the AND-conjunct walker is sound to descend
708 * through it -- the wrapper carries no constraint of its own.
709 * Skipping the descent would mark the walk incomplete and force
710 * the moment caller to fall back to MC even when the inner
711 * cmps are decidable closed-form. */
712 for (gate_t c : gc.getWires(g)) stk.push(c);
713 continue;
714 }
715 if (t == gate_input || t == gate_update || t == gate_one ||
716 t == gate_zero) {
717 /* Bernoulli leaf / constants: shift P(event), don't truncate
718 * any continuous RV. Skipping is sound and the walk stays
719 * complete. */
720 continue;
721 }
722 /* gate_plus (OR), gate_monus (set diff), gate_arith, gate_rv, ...:
723 * could affect an RV's conditional distribution in ways that
724 * don't reduce to an interval intersection. Mark the walk
725 * incomplete so a moment closed-form caller falls through to MC. */
726 complete = false;
727 }
728}
729
730/**
731 * @brief Walk @p root's AND-conjunct cmps and decide whether the
732 * conjunction is jointly infeasible by per-RV interval
733 * intersection.
734 *
735 * For every @c gate_cmp reachable through a chain of @c gate_times
736 * starting at @p root, that is interpretable as @c rv-op-constant,
737 * intersect the constraint with the running interval for that RV
738 * (initialised to the RV's distribution support). As soon as any
739 * RV's interval becomes empty, the AND is infeasible.
740 *
741 * Descends only through @c gate_times: @c gate_plus is OR (the
742 * disjuncts could individually be feasible even when each is a
743 * narrow constraint on the RV, so they do not contribute to the
744 * conjunction's infeasibility), @c gate_monus is set difference
745 * (likewise), and other gate types break the AND chain.
746 *
747 * Cmps that this pass cannot interpret (RV vs RV, arith on either
748 * side, agg…) are simply ignored: skipping them is sound &ndash; we
749 * just have fewer constraints, so we never falsely declare
750 * infeasibility we cannot prove.
751 */
752bool isAndJointlyInfeasible(const GenericCircuit &gc, gate_t root)
753{
754 std::unordered_map<gate_t, Interval> rv_intervals;
755 std::unordered_map<gate_t, Interval> support_cache;
756 bool complete;
757 walkAndConjunctIntervals(gc, root, rv_intervals, support_cache, complete);
758 for (const auto &kv : rv_intervals) {
759 if (intervalEmpty(kv.second)) return true;
760 }
761 return false;
762}
763
764/**
765 * @brief Memoised recursive predicate: does @p g's sub-circuit
766 * produce a continuous random variable (no point-mass /
767 * Dirac component)?
768 *
769 * Used to widen the EQ / NE = 0 / 1 shortcut at the cmp resolution
770 * site below the bare-@c gate_rv test, so multi-gate composites like
771 * <tt>Exp(0.4) + Exp(0.3) = c</tt> (heterogeneous-rate exponential
772 * sum, no closed-form Erlang fold) or
773 * <tt>mixture(p, Normal, Uniform) = c</tt> (Bernoulli mixture over
774 * two continuous arms) also resolve at load time. Without this the
775 * cmp falls through to AnalyticEvaluator (which returns NaN for
776 * EQ / NE) and then to the MC marginalisation, which in finite
777 * precision estimates @c P(X = Y) at 0 anyway -- but costs
778 * @c provsql.rv_mc_samples iterations to do so.
779 *
780 * Recursion:
781 * - @c gate_rv -> true (Normal / Uniform / Exp / Erlang all have
782 * continuous densities, no point masses).
783 * - @c gate_value -> false (Dirac at the literal).
784 * - @c gate_arith -> true iff every wire has only-continuous
785 * support. Sums, products, negations, divisions of continuous
786 * RVs stay continuous in distribution; a @c gate_value sibling
787 * poisons the result (e.g. @c X + 2 is continuous, but
788 * @c X * 0 = 0 has a Dirac at zero -- handled by the existing
789 * constant-fold pre-pass, but defensive here).
790 * - @c gate_mixture, Bernoulli 3-wire <tt>[p, X, Y]</tt> -> true
791 * iff X and Y are both continuous; the Boolean @c p only chooses
792 * an arm, so it does not affect the support type.
793 * - @c gate_mixture, categorical
794 * <tt>[key, mul_1, ..., mul_n]</tt> -> false (point masses at
795 * each mulinput's outcome value).
796 * - Any other gate type -> false (defensive: gate_plus / gate_times
797 * / gate_cmp / gate_agg are not continuous-RV containers).
798 *
799 * The cache is keyed on @c gate_t and may be shared across multiple
800 * cmp gates inside a single @c runRangeCheck invocation.
801 */
802bool hasOnlyContinuousSupport(const GenericCircuit &gc, gate_t g,
803 std::unordered_map<gate_t, bool> &cache)
804{
805 auto it = cache.find(g);
806 if (it != cache.end()) return it->second;
807 /* Memoise pessimistically before recursing so a malformed cyclic
808 * sub-circuit (shouldn't happen on well-formed input) returns
809 * @c false rather than blowing the stack. */
810 cache[g] = false;
811
812 bool result = false;
813 auto t = gc.getGateType(g);
814 switch (t) {
815 case gate_rv: {
816 /* A discrete family (Poisson, Binomial) has point masses, so a point
817 * event X = c carries positive mass and must NOT take the continuous
818 * EQ/NE measure-zero shortcut. isDiscrete() is the authoritative
819 * per-family flag (parse_distribution_template handles a latent /
820 * parametric leaf too); a malformed spec falls back to continuous. */
821 auto tmpl = parse_distribution_template(gc.getExtra(g));
822 result = !(tmpl && tmpl->family->factory(0.0, 0.0)->isDiscrete());
823 break;
824 }
825 case gate_value:
826 result = false;
827 break;
828 case gate_arith: {
829 result = true;
830 for (gate_t w : gc.getWires(g)) {
831 if (!hasOnlyContinuousSupport(gc, w, cache)) { result = false; break; }
832 }
833 break;
834 }
835 case gate_mixture: {
836 if (gc.isCategoricalMixture(g)) { result = false; break; }
837 const auto &w = gc.getWires(g);
838 if (w.size() != 3) { result = false; break; }
839 result = hasOnlyContinuousSupport(gc, w[1], cache)
840 && hasOnlyContinuousSupport(gc, w[2], cache);
841 break;
842 }
843 default:
844 result = false;
845 break;
846 }
847
848 cache[g] = result;
849 return result;
850}
851
852/**
853 * @brief Recursive collection of the @c gate_rv and @c gate_input
854 * leaves reachable from @p g.
855 *
856 * The result is a sub-circuit's "random-source footprint": two
857 * sub-circuits are independent iff their random-source sets are
858 * disjoint. Used to gate the exact-EQ Dirac sum-product below: the
859 * factoring @c P(X = Y) = Σ_v @c P(X=v)·P(Y=v) is only valid when
860 * @c X and @c Y are independent, otherwise the per-row coupling
861 * (e.g. two mixtures sharing a Bernoulli @c p_token) breaks the
862 * factoring and the sum-product silently produces the wrong
863 * probability.
864 *
865 * Descent rules: @c gate_arith and @c gate_mixture descend into all
866 * children (Bernoulli @c p_token, categorical key, mulinputs all
867 * contribute to the random footprint). @c gate_value is a
868 * deterministic literal and contributes no random source. Other
869 * gate types (Boolean / agg / etc.) don't appear under a continuous
870 * cmp side in well-formed circuits; defensively, they contribute
871 * nothing.
872 */
873const std::unordered_set<gate_t> &
874collectRandomLeaves(const GenericCircuit &gc, gate_t g,
875 std::unordered_map<gate_t, std::unordered_set<gate_t>> &cache)
876{
877 auto it = cache.find(g);
878 if (it != cache.end()) return it->second;
879 /* Insert an empty entry first so a recursive call on a cyclic
880 * sub-circuit returns early. std::unordered_map insertion does
881 * not invalidate references to existing elements, but it MAY
882 * rehash on growth (invalidating ALL references, including the
883 * one we're about to capture). Build the result locally, then
884 * write it back in one shot at the end. */
885 cache.emplace(g, std::unordered_set<gate_t>{});
886
887 std::unordered_set<gate_t> out;
888 auto t = gc.getGateType(g);
889 if (t == gate_rv || t == gate_input) {
890 out.insert(g);
891 } else if (t == gate_arith || t == gate_mixture) {
892 for (gate_t w : gc.getWires(g)) {
893 const auto &child = collectRandomLeaves(gc, w, cache);
894 out.insert(child.begin(), child.end());
895 }
896 }
897
898 /* Overwrite the placeholder; locate by find() to avoid a fresh
899 * insertion that could rehash and invalidate other iterators in
900 * upstream frames. */
901 auto fit = cache.find(g);
902 fit->second = std::move(out);
903 return fit->second;
904}
905
906using DiracMap = std::unordered_map<double, double>;
907using DiracMapOpt = std::optional<DiracMap>;
908
909/**
910 * @brief Recursive extraction of @p g's Dirac mass map (value -> mass).
911 *
912 * Returns @c std::nullopt when the sub-circuit's discrete component
913 * is not statically extractable (e.g. an opaque @c gate_arith over
914 * mixtures, a Bernoulli mixture whose @c p_token is a compound
915 * Boolean, etc.). When the sub-circuit is purely continuous the
916 * map is well-defined but empty (no Diracs, no masses).
917 *
918 * Used by the exact EQ shortcut below: for independent @c X, @c Y
919 * with extractable mass maps @c M_X, @c M_Y:
920 * <tt>P(X = Y) = Σ_{v ∈ M_X ∩ M_Y} M_X[v] · M_Y[v]</tt>. Continuous
921 * components contribute zero by measure-zero arguments (Dirac vs
922 * continuous and continuous vs continuous), so they need not appear
923 * in the sum.
924 *
925 * Shape rules:
926 * - @c gate_value:v: a Dirac at the literal with mass @c 1.
927 * - @c gate_rv: continuous in every supported family, empty map.
928 * - categorical @c gate_mixture <tt>[key, mul_1, ..., mul_n]</tt>:
929 * sum @c getProb(mul_i) into @c map[parseDouble(extra(mul_i))].
930 * Multiple mulinputs at the same outcome (which the constructor
931 * doesn't produce but is sound to handle) merge masses.
932 * - Bernoulli @c gate_mixture <tt>[p_token, X, Y]</tt> with
933 * @c p_token a bare @c gate_input: pull @c π = @c getProb(p_token)
934 * and recurse into X, Y to get @c M_X, @c M_Y; result is
935 * <tt>π·M_X[v] + (1-π)·M_Y[v]</tt> per outcome value. Compound
936 * Boolean @c p_tokens (whose probability would have to come from
937 * a recursive @c probability_evaluate call) bail.
938 * - Anything else: @c std::nullopt.
939 */
940DiracMapOpt
941collectDiracMassMap(const GenericCircuit &gc, gate_t g,
942 std::unordered_map<gate_t, DiracMapOpt> &cache)
943{
944 auto it = cache.find(g);
945 if (it != cache.end()) return it->second;
946 /* Pessimistic cycle guard, same reasoning as @c collectRandomLeaves. */
947 cache.emplace(g, std::nullopt);
948
949 DiracMapOpt result;
950 auto t = gc.getGateType(g);
951 switch (t) {
952 case gate_value: {
953 try {
954 DiracMap m;
955 m[parseDoubleStrict(gc.getExtra(g))] = 1.0;
956 result = std::move(m);
957 } catch (const CircuitException &) {
958 /* unparseable extra: bail */
959 }
960 break;
961 }
962 case gate_rv: {
963 /* A CONTINUOUS leaf has no point masses (an empty map is exact: the
964 * Dirac sum-product then contributes zero, correct by measure zero).
965 * A DISCRETE leaf (Poisson, Binomial) DOES have point masses, but
966 * they are not statically enumerable (infinite / large support), so
967 * decline (nullopt) -- claiming an empty map here would make the
968 * sum-product read "no overlap" and wrongly fold X = c to false.
969 * isDiscrete() is the authoritative per-family flag. */
970 auto tmpl = parse_distribution_template(gc.getExtra(g));
971 if (tmpl && tmpl->family->factory(0.0, 0.0)->isDiscrete())
972 result = std::nullopt;
973 else
974 result = DiracMap{};
975 break;
976 }
977 case gate_mixture: {
978 const auto &w = gc.getWires(g);
979 if (gc.isCategoricalMixture(g)) {
980 DiracMap m;
981 bool ok = true;
982 for (std::size_t i = 1; i < w.size(); ++i) {
983 double v;
984 try { v = parseDoubleStrict(gc.getExtra(w[i])); }
985 catch (const CircuitException &) { ok = false; break; }
986 const double p = gc.getProb(w[i]);
987 if (!std::isfinite(p) || p < 0.0 || p > 1.0) { ok = false; break; }
988 m[v] += p;
989 }
990 if (ok) result = std::move(m);
991 } else if (w.size() == 3
992 && gc.getGateType(w[0]) == gate_input) {
993 const double pi = gc.getProb(w[0]);
994 if (std::isfinite(pi) && pi >= 0.0 && pi <= 1.0) {
995 auto mx = collectDiracMassMap(gc, w[1], cache);
996 auto my = collectDiracMassMap(gc, w[2], cache);
997 if (mx && my) {
998 DiracMap m;
999 for (const auto &[v, mass] : *mx) m[v] += pi * mass;
1000 for (const auto &[v, mass] : *my) m[v] += (1.0 - pi) * mass;
1001 result = std::move(m);
1002 }
1003 }
1004 }
1005 break;
1006 }
1007 default:
1008 break;
1009 }
1010
1011 auto fit = cache.find(g);
1012 fit->second = result;
1013 return result;
1014}
1015
1016} // namespace
1017
1019{
1020 std::unordered_map<gate_t, Interval> cache;
1021 /* Shared across all cmp gates in this @c runRangeCheck invocation.
1022 * Keyed on gate_t and immutable across cmp iterations because
1023 * resolving one cmp only changes the cmp's own gate type, not
1024 * the sub-circuit underneath @c wires[0..1] of other cmps. */
1025 std::unordered_map<gate_t, bool> continuous_support_cache;
1026 std::unordered_map<gate_t, DiracMapOpt> dirac_cache;
1027 std::unordered_map<gate_t, std::unordered_set<gate_t>> leaf_cache;
1028 unsigned resolved = 0;
1029
1030 /* Snapshot the cmp gate ids before we start mutating: in-place
1031 * resolution turns a @c gate_cmp into a @c gate_input, but
1032 * @c getNbGates only grows, never shrinks, so iterating by index
1033 * over the original count is safe. We re-check the type at each
1034 * step to skip already-resolved slots. */
1035 const auto nb = gc.getNbGates();
1036 std::vector<gate_t> cmps;
1037 for (std::size_t i = 0; i < nb; ++i) {
1038 auto g = static_cast<gate_t>(i);
1039 if (gc.getGateType(g) == gate_cmp)
1040 cmps.push_back(g);
1041 }
1042
1043 for (gate_t c : cmps) {
1044 if (gc.getGateType(c) != gate_cmp) continue; /* defensive */
1045
1046 bool ok = false;
1047 ComparisonOperator op = cmpOpFromOid(gc.getInfos(c).first, ok);
1048 if (!ok) continue;
1049
1050 const auto &wires = gc.getWires(c);
1051 if (wires.size() != 2) continue;
1052
1053 /* Identity shortcut: when both sides of the cmp are the same
1054 * gate (same UUID), the sampler's per-iteration memoisation
1055 * guarantees both reads return identical values, so the
1056 * comparator collapses to a constant. Universal across gate
1057 * types and semirings; runs first so neither the continuous
1058 * EQ/NE shortcut nor the interval-based path needs an explicit
1059 * @c lhs != rhs guard. */
1060 if (wires[0] == wires[1]) {
1061 double p = std::numeric_limits<double>::quiet_NaN();
1062 switch (op) {
1066 p = 1.0; break;
1070 p = 0.0; break;
1071 }
1072 gc.resolveCmpToBernoulli(c, p);
1073 ++resolved;
1074 continue;
1075 }
1076
1077 /* Continuous EQ / NE shortcut: P(X = c) = 0 and P(X != c) = 1
1078 * exactly when at least one side has a continuous distribution
1079 * (point equality has measure zero under any continuous
1080 * distribution). Universal across semirings: the gate_zero /
1081 * gate_one rewrite is meaningful in every semiring (not just
1082 * probability), so the resolution belongs here rather than in
1083 * AnalyticEvaluator.
1084 *
1085 * @c hasOnlyContinuousSupport widens the test beyond a bare
1086 * @c gate_rv leaf: heterogeneous-rate exponential sums, products
1087 * of independent continuous RVs, and Bernoulli mixtures over
1088 * two continuous arms all qualify because their distribution
1089 * has no point-mass component. Categorical mixtures (point
1090 * masses at each outcome value) and pure-deterministic
1091 * @c gate_value sub-circuits do NOT qualify and fall through to
1092 * the agg / interval / AnalyticEvaluator paths.
1093 *
1094 * The @c wires[0] == @c wires[1] case is already handled by the
1095 * identity shortcut above. */
1096 if (op == ComparisonOperator::EQ ||
1097 op == ComparisonOperator::NE) {
1098 bool lhs_continuous = hasOnlyContinuousSupport(gc, wires[0],
1099 continuous_support_cache);
1100 bool rhs_continuous = hasOnlyContinuousSupport(gc, wires[1],
1101 continuous_support_cache);
1102 if (lhs_continuous || rhs_continuous) {
1103 double p = (op == ComparisonOperator::EQ) ? 0.0 : 1.0;
1104 gc.resolveCmpToBernoulli(c, p);
1105 ++resolved;
1106 continue;
1107 }
1108
1109 /* Exact Dirac sum-product. When both sides have extractable
1110 * @c (value -> mass) maps AND the two sub-circuits are
1111 * independent (random-leaf footprints disjoint), the
1112 * convolution at zero of @c (X - Y) has support exactly on
1113 * @c Dirac(X) ∩ Dirac(Y) with mass
1114 * <tt>M_X(v) · M_Y(v)</tt> per overlapping value; the
1115 * continuous and continuous-vs-Dirac contributions vanish by
1116 * measure zero. This generalises the bare-disjoint case to
1117 * any pair of statically-known discrete distributions:
1118 * <tt>P(categorical(a) = categorical(b))</tt> with overlapping
1119 * outcomes, mixtures with @c as_random branches, etc.
1120 *
1121 * The independence test is essential: two mixtures sharing a
1122 * Bernoulli @c p_token are correlated and the sum-product
1123 * factoring breaks (the actual @c P(X=Y) cannot be recovered
1124 * from the marginals alone). @c collectRandomLeaves'
1125 * footprint-disjoint check is the gate.
1126 *
1127 * When both maps are empty (purely continuous on both sides)
1128 * the existing branch above already fired, so the sum-product
1129 * path here only runs for at-least-one-discrete shapes. */
1130 auto m_l = collectDiracMassMap(gc, wires[0], dirac_cache);
1131 auto m_r = collectDiracMassMap(gc, wires[1], dirac_cache);
1132 if (m_l && m_r) {
1133 const auto &leaves_l = collectRandomLeaves(gc, wires[0], leaf_cache);
1134 const auto &leaves_r = collectRandomLeaves(gc, wires[1], leaf_cache);
1135 bool independent = true;
1136 for (gate_t leaf : leaves_l) {
1137 if (leaves_r.count(leaf)) { independent = false; break; }
1138 }
1139 if (independent) {
1140 double p_eq = 0.0;
1141 /* Iterate over the smaller map to keep the sum at
1142 * O(min(|M_l|, |M_r|)) lookups. */
1143 const DiracMap *small = (m_l->size() <= m_r->size()) ? &*m_l : &*m_r;
1144 const DiracMap *large = (m_l->size() <= m_r->size()) ? &*m_r : &*m_l;
1145 for (const auto &[v, mass] : *small) {
1146 auto fit = large->find(v);
1147 if (fit != large->end()) p_eq += mass * fit->second;
1148 }
1149 /* Clamp into @c [0, 1] defensively: floating-point summation
1150 * of masses (each in [0, 1]) might overshoot by an ULP, and
1151 * @c resolveCmpToBernoulli requires a strict probability. */
1152 if (p_eq < 0.0) p_eq = 0.0;
1153 if (p_eq > 1.0) p_eq = 1.0;
1154 double p = (op == ComparisonOperator::EQ) ? p_eq : 1.0 - p_eq;
1155 gc.resolveCmpToBernoulli(c, p);
1156 ++resolved;
1157 continue;
1158 }
1159 }
1160 }
1161
1162 /* HAVING-style cmp: agg on one side, scalar constant on the
1163 * other. Decide via the agg-aware path which is cheaper than
1164 * intervalOf + decideCmp and which knows the empty-subset NULL
1165 * semantics for SUM / MIN / MAX (see decideAggVsConstCmp). */
1166 bool lhs_is_agg = gc.getGateType(wires[0]) == gate_agg;
1167 bool rhs_is_agg = gc.getGateType(wires[1]) == gate_agg;
1168 if (lhs_is_agg != rhs_is_agg) {
1169 gate_t agg_side = lhs_is_agg ? wires[0] : wires[1];
1170 gate_t const_side = lhs_is_agg ? wires[1] : wires[0];
1171 double const_val = extractScalarConst(gc, const_side);
1172 if (!std::isnan(const_val)) {
1173 double p = decideAggVsConstCmp(gc, agg_side, op, const_val,
1174 lhs_is_agg);
1175 if (!std::isnan(p)) {
1176 gc.resolveCmpToBernoulli(c, p);
1177 ++resolved;
1178 continue;
1179 }
1180 }
1181 }
1182
1183 /* Interval-based path for non-agg cmps (RV, gate_arith, value). */
1184 Interval lhs = intervalOf(gc, wires[0], cache);
1185 Interval rhs = intervalOf(gc, wires[1], cache);
1186 /* Skip if both sides are unbounded; @c decideCmp would never
1187 * return a decision and the work is wasted. */
1188 if (lhs.isAll() && rhs.isAll()) continue;
1189
1190 Interval diff = sub(lhs, rhs);
1191 double p = decideCmp(diff, op);
1192 if (!std::isnan(p)) {
1193 gc.resolveCmpToBernoulli(c, p);
1194 ++resolved;
1195 }
1196 }
1197
1198 /* Joint-conjunction pass: walk every @c gate_times and check
1199 * whether its AND-conjunct cmps, viewed together, constrain some
1200 * shared RV to an empty interval. Catches the joint-infeasibility
1201 * case the per-cmp pass above cannot see (each cmp individually
1202 * leaves a non-empty range, but their intersection is empty).
1203 *
1204 * Snapshot the gate_times indices first: @c resolveGateToZero
1205 * mutates the type, so iterating the live vector while resolving
1206 * would skip slots. The post-snapshot type re-check guards against
1207 * a @c gate_times that the per-cmp pass somehow already collapsed
1208 * (currently not possible, but cheap insurance for future passes). */
1209 const auto nb_after = gc.getNbGates();
1210 std::vector<gate_t> times_gates;
1211 for (std::size_t i = 0; i < nb_after; ++i) {
1212 auto g = static_cast<gate_t>(i);
1213 if (gc.getGateType(g) == gate_times)
1214 times_gates.push_back(g);
1215 }
1216 for (gate_t t : times_gates) {
1217 if (gc.getGateType(t) != gate_times) continue; /* defensive */
1218 if (isAndJointlyInfeasible(gc, t)) {
1219 gc.resolveGateToZero(t);
1220 ++resolved;
1221 }
1222 }
1223
1224 return resolved;
1225}
1226
1227/**
1228 * @brief Probability-side pre-pass: rewrite HAVING-style @c gate_cmp
1229 * gates that are provably TRUE on the agg's value-interval
1230 * into an OR over the agg's per-row K-gates.
1231 *
1232 * Companion to @c runCountCmpEvaluator's Poisson-binomial pre-pass:
1233 * where that one resolves @c COUNT op C to a closed-form Bernoulli,
1234 * this one catches the always-true sub-case (e.g. @c COUNT <= K with
1235 * @c K >= N inputs, or any aggregator whose value-interval entirely
1236 * satisfies the predicate) and replaces the cmp with @c gate_plus
1237 * over the agg's K-gates -- the "group is non-empty" indicator.
1238 *
1239 * Why a separate pass: @c runRangeCheck deliberately blocks TRUE
1240 * decisions because @c gate_one is universally unsound for HAVING
1241 * (it would credit the empty world). The safe TRUE rewrite
1242 * "OR of K-gates" requires absorptive @c gate_plus semantics
1243 * (probability, Boolean, formula, why, which, max-min, max-max), so
1244 * the pass is restricted to the probability-evaluate path where
1245 * absorption is guaranteed by the downstream BoolExpr translation.
1246 *
1247 * Fires regardless of @c provsql.cmp_probability_evaluation: when
1248 * the Poisson-binomial path is disabled (developer A/B testing),
1249 * this lighter shortcut still catches the always-true case and
1250 * spares the d-DNNF compiler the 2^N-clause DNF that
1251 * @c provsql_having's @c enumerate_valid_worlds would otherwise emit.
1252 *
1253 * Same matching contract as @c decideAggVsConstCmp for the agg side:
1254 * cmp wires must be {gate_agg, scalar-const-encoded-as-semimod}, the
1255 * agg's children must all be @c gate_semimod, and the agg kind must
1256 * be one of COUNT / SUM / MIN / MAX (the only kinds with an
1257 * interval). Mismatches leave the cmp untouched.
1258 *
1259 * @param gc Circuit to mutate in place.
1260 * @return Number of comparators rewritten to gate_plus.
1261 */
1262/**
1263 * @brief Decide every comparison whose two sides are the very same gate.
1264 *
1265 * Two aggregates that read the same rows and contribute the same value on
1266 * each are ONE gate, the circuit being hash-consed: @c count(v) and
1267 * @c count(w) over a group both count its rows. The comparison is then
1268 * settled by the operator alone, whatever the data -- but settling it TRUE
1269 * is not @c gate_one. A tautology over a group still says that the group is
1270 * there, which is the OR over the k-gates; @c gate_one would credit the world
1271 * where the group is empty and SQL returns no row at all. That is the
1272 * over-credit the doc comment on @c decideAggVsConstCmp describes, reached by
1273 * another road.
1274 *
1275 * Runs at the FRONT of the resolution pipeline: this is a structural fact
1276 * about the circuit, not a probabilistic one, and the value simplifier that
1277 * runs early would otherwise fold the comparison away as a constant and lose
1278 * the group with it.
1279 */
1281{
1282 unsigned resolved = 0;
1283 const auto nb = gc.getNbGates();
1284 std::vector<gate_t> cmps;
1285 for (std::size_t i = 0; i < nb; ++i) {
1286 auto g = static_cast<gate_t>(i);
1287 if (gc.getGateType(g) == gate_cmp)
1288 cmps.push_back(g);
1289 }
1290
1291 for (gate_t c : cmps) {
1292 if (gc.getGateType(c) != gate_cmp) continue; /* defensive */
1293 bool ok = false;
1294 ComparisonOperator op = cmpOpFromOid(gc.getInfos(c).first, ok);
1295 if (!ok) continue;
1296 const auto &wires = gc.getWires(c);
1297 if (wires.size() != 2) continue;
1298 const bool lhs_is_agg = gc.getGateType(wires[0]) == gate_agg;
1299 /* The very same gate on both sides. Two aggregates that read the same
1300 * rows and contribute the same value on each are ONE gate, the circuit
1301 * being hash-consed: count(v) and count(w) over a group both count its
1302 * rows, so the comparison is decided by the operator alone, whatever the
1303 * data. Deciding it true must not give gate_one, though -- a tautology
1304 * over a group still says that the group is there, which is the OR over
1305 * the k-gates. gate_one would credit the world where the group is empty,
1306 * where SQL returns no row at all; that is the same over-credit the doc
1307 * comment on decideAggVsConstCmp describes, reached by another road. */
1308 if (!(wires[0] == wires[1] && lhs_is_agg)) continue;
1309 {
1310 const bool reflexive_true = op == ComparisonOperator::EQ ||
1311 op == ComparisonOperator::LE ||
1313 gate_t agg = wires[0];
1314 std::vector<gate_t> ks;
1315 bool shape_ok = true;
1316
1317 /* Every contribution a value: the aggregate then has a value on any
1318 * non-empty group, so the comparison is not the UNKNOWN of a NULL
1319 * aggregate, which SQL filters out. */
1320 for (gate_t ch : gc.getWires(agg)) {
1321 if (gc.getGateType(ch) != gate_semimod) { shape_ok = false; break; }
1322 const auto &sw = gc.getWires(ch);
1323 if (sw.size() != 2 || gc.getGateType(sw[1]) != gate_value) {
1324 shape_ok = false;
1325 break;
1326 }
1327 ks.push_back(sw[0]);
1328 }
1329 if (!shape_ok) continue;
1330
1331 if (!reflexive_true) {
1332 gc.resolveGateToZero(c); /* x < x, x > x, x <> x: in no world */
1333 ++resolved;
1334 continue;
1335 }
1336 /* A scalar COUNT has a value over no row too (it is 0), so its
1337 * tautology holds in every world, the empty one included. */
1338 if ((gc.getInfos(agg).second & PROVSQL_AGG_SCALAR_FLAG) != 0 &&
1339 getAggregationOperator(gc.getInfos(agg).first) ==
1341 gc.resolveCmpToBernoulli(c, 1.0);
1342 ++resolved;
1343 continue;
1344 }
1345 if (ks.empty()) continue;
1346 gc.resolveCmpToPlusOfKGates(c, ks);
1347 ++resolved;
1348 continue;
1349 }
1350
1351 }
1352 return resolved;
1353}
1354
1356{
1357 unsigned resolved = 0;
1358 const auto nb = gc.getNbGates();
1359
1360 std::vector<gate_t> cmps;
1361 cmps.reserve(nb / 8); /* rough guess */
1362 for (std::size_t i = 0; i < nb; ++i) {
1363 auto g = static_cast<gate_t>(i);
1364 if (gc.getGateType(g) == gate_cmp)
1365 cmps.push_back(g);
1366 }
1367
1368 for (gate_t c : cmps) {
1369 if (gc.getGateType(c) != gate_cmp) continue; /* defensive */
1370
1371 bool ok = false;
1372 ComparisonOperator op = cmpOpFromOid(gc.getInfos(c).first, ok);
1373 if (!ok) continue;
1374
1375 const auto &wires = gc.getWires(c);
1376 if (wires.size() != 2) continue;
1377
1378 bool lhs_is_agg = gc.getGateType(wires[0]) == gate_agg;
1379 bool rhs_is_agg = gc.getGateType(wires[1]) == gate_agg;
1380
1381 if (lhs_is_agg == rhs_is_agg) continue; /* both agg or neither */
1382
1383 gate_t agg_side = lhs_is_agg ? wires[0] : wires[1];
1384 gate_t const_side = lhs_is_agg ? wires[1] : wires[0];
1385
1386 double const_val = extractScalarConst(gc, const_side);
1387 if (std::isnan(const_val)) continue;
1388
1389 bool always_true = false;
1390 double p = decideAggVsConstCmp(gc, agg_side, op, const_val,
1391 lhs_is_agg, &always_true);
1392 if (!always_true) {
1393 /* p might be 0.0 (already handled by runRangeCheck at load time
1394 * if simplify_on_load is on); skip either way. */
1395 (void)p;
1396 continue;
1397 }
1398
1399 /* Scalar aggregation (no GROUP BY): the single result row always exists, so
1400 * a tautological predicate is gate_one -- probability 1, including the
1401 * empty-input world -- but only where the aggregate has a value there.
1402 * count(*) does (it is 0, and count >= 0 holds); sum, min, max and avg are
1403 * NULL over no row, the comparison is then unknown and the row is filtered
1404 * out, so the empty world must be excluded exactly as it is for a group.
1405 * The "group is non-empty" rewrite below does that, and is what the doc
1406 * comment on decideAggVsConstCmp calls the empty-world over-credit only
1407 * where the aggregate is defined on the empty input. */
1408 if ((gc.getInfos(agg_side).second & PROVSQL_AGG_SCALAR_FLAG) != 0 &&
1409 getAggregationOperator(gc.getInfos(agg_side).first) ==
1411 gc.resolveCmpToBernoulli(c, 1.0);
1412 ++resolved;
1413 continue;
1414 }
1415
1416 /* Gather the per-row K-gates from the agg's semimod children. */
1417 std::vector<gate_t> ks;
1418 bool shape_ok = true;
1419 ks.reserve(gc.getWires(agg_side).size());
1420 for (gate_t ch : gc.getWires(agg_side)) {
1421 if (gc.getGateType(ch) != gate_semimod) { shape_ok = false; break; }
1422 const auto &sw = gc.getWires(ch);
1423 if (sw.size() != 2) { shape_ok = false; break; }
1424 ks.push_back(sw[0]); /* K side; M side is sw[1] = gate_value */
1425 }
1426 if (!shape_ok || ks.empty()) continue;
1427
1428 gc.resolveCmpToPlusOfKGates(c, ks);
1429 ++resolved;
1430 }
1431
1432 return resolved;
1433}
1434
1435namespace {
1436
1437/* Sampling-time support for aggregation gates. Unlike @c intervalOf
1438 * (shared with the @c runRangeCheck cmp dispatcher, which must stay
1439 * conservative on agg to respect SQL NULL semantics), this widens
1440 * gate_agg and non-trivial gate_semimod to the actual range of
1441 * scalar MC samples @c MonteCarloSampler::evalScalar can produce.
1442 * Used only by @c compute_support, called from rv_support /
1443 * rv_histogram / rv_moment fallbacks, so the cmp decider remains
1444 * untouched.
1445 *
1446 * Empty-group convention (see test/sql/continuous_aggregation §5):
1447 * COUNT and SUM yield 0; MIN / MAX / AVG yield NaN. NaN sits
1448 * outside any real interval, so callers binning samples drop those
1449 * worlds automatically; the moment averagers in
1450 * Expectation::mc_raw_moment also skip them.
1451 */
1452Interval aggSupportOf(const GenericCircuit &gc, gate_t root,
1453 std::unordered_map<gate_t, Interval> &cache)
1454{
1455 const auto type = gc.getGateType(root);
1456 if (type == gate_semimod) {
1457 const auto &wires = gc.getWires(root);
1458 if (wires.size() != 2) return Interval::all();
1459 Interval vi = intervalOf(gc, wires[1], cache);
1460 if (gc.getGateType(wires[0]) == gate_one) return vi;
1461 /* Boolean k child: per-iteration scalar is value · 1_{k fires},
1462 * so the support is the union of {0} and the value's range. */
1463 return Interval{std::min(0.0, vi.lo), std::max(0.0, vi.hi)};
1464 }
1465 if (type != gate_agg) return intervalOf(gc, root, cache);
1466
1467 const auto &wires = gc.getWires(root);
1469 getAggregationOperator(gc.getInfos(root).first);
1470
1471 std::vector<gate_t> sm_children;
1472 sm_children.reserve(wires.size());
1473 for (gate_t c : wires)
1474 if (gc.getGateType(c) == gate_semimod) sm_children.push_back(c);
1475
1476 auto child_value_iv = [&](gate_t sm) -> Interval {
1477 const auto &sw = gc.getWires(sm);
1478 if (sw.size() != 2) return Interval::all();
1479 return intervalOf(gc, sw[1], cache);
1480 };
1481 auto child_always_fires = [&](gate_t sm) -> bool {
1482 const auto &sw = gc.getWires(sm);
1483 return sw.size() == 2 && gc.getGateType(sw[0]) == gate_one;
1484 };
1485
1486 const auto inf = std::numeric_limits<double>::infinity();
1487 switch (op) {
1489 /* [0, n_rows]. Each semimod contributes 0 or 1 to the count. */
1490 return Interval{0.0, static_cast<double>(sm_children.size())};
1492 /* Per row, contribution is value if k fires, else 0; sum the
1493 * per-row support intervals. Always-firing rows contribute
1494 * their value interval verbatim; possibly-firing rows
1495 * contribute [min(0, v.lo), max(0, v.hi)]. */
1496 double lo = 0.0, hi = 0.0;
1497 for (gate_t sm : sm_children) {
1498 Interval vi = child_value_iv(sm);
1499 if (vi.isAll()) return Interval::all();
1500 if (child_always_fires(sm)) {
1501 lo += vi.lo;
1502 hi += vi.hi;
1503 } else {
1504 lo += std::min(0.0, vi.lo);
1505 hi += std::max(0.0, vi.hi);
1506 }
1507 }
1508 return Interval{lo, hi};
1509 }
1512 /* MIN / MAX of values where k_i fires. The actual MIN (or MAX)
1513 * is some value from one of the firing rows, so the support is
1514 * the union of the children's value intervals. Empty-group
1515 * worlds finalise to NaN, which sits outside the real
1516 * interval. */
1517 if (sm_children.empty()) return Interval::all();
1518 double lo = inf;
1519 double hi = -inf;
1520 for (gate_t sm : sm_children) {
1521 Interval vi = child_value_iv(sm);
1522 if (vi.isAll()) return Interval::all();
1523 lo = std::min(lo, vi.lo);
1524 hi = std::max(hi, vi.hi);
1525 }
1526 if (lo > hi) return Interval::all();
1527 return Interval{lo, hi};
1528 }
1529 default:
1530 /* AVG: ratio depends on the world's row count. AND / OR /
1531 * CHOOSE / ARRAY_AGG: not numeric carriers rv_* surfaces.
1532 * Keep the conservative all-real default. */
1533 return Interval::all();
1534 }
1535}
1536
1537} // namespace
1538
1539std::pair<double, double>
1541 std::optional<gate_t> event_root)
1542{
1543 std::unordered_map<gate_t, Interval> cache;
1544 Interval iv = aggSupportOf(gc, root, cache);
1545
1546 /* Conditional path: intersect with the event's AND-conjunct
1547 * constraints on @p root. Walks event_root collecting `rv op c`
1548 * cmps; non-target constraints are ignored (they affect P(event)
1549 * but not the truncation of root's distribution). Even if the
1550 * walk is "incomplete" (gate_plus / gate_monus / arith encountered)
1551 * the result is sound: we're computing a SUPERSET bound on the
1552 * conditional support, and the unconditional support is already a
1553 * superset, so the intersection of the collected constraints with
1554 * the unconditional is also a superset. */
1555 if (event_root.has_value()) {
1556 std::unordered_map<gate_t, Interval> rv_intervals;
1557 bool complete;
1558 walkAndConjunctIntervals(gc, *event_root, rv_intervals, cache, complete);
1559 auto it = rv_intervals.find(root);
1560 if (it != rv_intervals.end()) {
1561 iv.lo = std::max(iv.lo, it->second.lo);
1562 iv.hi = std::min(iv.hi, it->second.hi);
1563 /* Defensively clamp to avoid an inverted interval if a buggy
1564 * walker produced one; should not happen but cheap. */
1565 if (iv.lo > iv.hi) iv.lo = iv.hi;
1566 }
1567 }
1568
1569 return {iv.lo, iv.hi};
1570}
1571
1572std::optional<std::pair<double, double>>
1574 gate_t target_rv)
1575{
1576 std::unordered_map<gate_t, Interval> rv_intervals;
1577 std::unordered_map<gate_t, Interval> support_cache;
1578 bool complete;
1579 walkAndConjunctIntervals(gc, event_root, rv_intervals, support_cache,
1580 complete);
1581 if (!complete) return std::nullopt;
1582 /* If the walk found no cmp constraining target_rv, the conditional
1583 * support is the unconditional support (the event is independent
1584 * of target_rv along the recognised structure). Returning the
1585 * unconditional interval lets the moment closed-form path
1586 * short-circuit to the unconditional moment, matching the
1587 * mathematical truth. */
1588 auto it = rv_intervals.find(target_rv);
1589 Interval iv;
1590 if (it != rv_intervals.end()) {
1591 iv = it->second;
1592 /* Intersect with the RV's own support to be safe (event may
1593 * over-constrain past the support, e.g. `Exp(λ) < -1`). */
1594 Interval base = intervalOf(gc, target_rv, support_cache);
1595 iv.lo = std::max(iv.lo, base.lo);
1596 iv.hi = std::min(iv.hi, base.hi);
1597 if (iv.lo > iv.hi) iv.lo = iv.hi;
1598 } else {
1599 iv = intervalOf(gc, target_rv, support_cache);
1600 }
1601 return std::make_pair(iv.lo, iv.hi);
1602}
1603
1604/**
1605 * @brief Parse a @c gate_value's @c extra as a finite @c float8.
1606 *
1607 * Sibling of @c extract_constant_string in @c having_semantics.cpp but
1608 * parsing a double, with a const @c GenericCircuit ref (used in the
1609 * closed-form shape detector path). Bails on @c NaN / @c ±Infinity so a downstream
1610 * stem renderer never sees a non-finite x coordinate.
1611 */
1613 double &out)
1614{
1615 if (gc.getGateType(x) != gate_value) return false;
1616 const std::string &s = gc.getExtra(x);
1617 if (s.empty()) return false;
1618 try {
1619 size_t idx = 0;
1620 double v = std::stod(s, &idx);
1621 if (idx != s.size() || !std::isfinite(v)) return false;
1622 out = v;
1623 return true;
1624 } catch (...) {
1625 return false;
1626 }
1627}
1628
1629/** @brief Same parsing applied to a mulinput's outcome label (categorical). */
1631 double &out)
1632{
1633 if (gc.getGateType(mul) != gate_mulinput) return false;
1634 const std::string &s = gc.getExtra(mul);
1635 if (s.empty()) return false;
1636 try {
1637 size_t idx = 0;
1638 double v = std::stod(s, &idx);
1639 if (idx != s.size() || !std::isfinite(v)) return false;
1640 out = v;
1641 return true;
1642 } catch (...) {
1643 return false;
1644 }
1645}
1646
1647std::optional<TruncatedSingleRv>
1649 std::optional<gate_t> event_root)
1650{
1651 if (gc.getGateType(root) != gate_rv) return std::nullopt;
1652 auto spec = parse_distribution_spec(gc.getExtra(root));
1653 if (!spec) return std::nullopt;
1654
1655 /* Natural support per family. Normal is unbounded both sides;
1656 * Uniform sits exactly on its parameters; Exp / Erlang on
1657 * [0, +inf). Used both as the unconditional case and as the
1658 * intersection seed for collectRvConstraints (which already
1659 * intersects internally, but the bare-natural case still needs
1660 * a baseline). */
1661 const DistSupport nat_support = makeDistribution(*spec)->support();
1662 double nat_lo = nat_support.lo;
1663 double nat_hi = nat_support.hi;
1664
1665 /* Unconditional path: return natural support, mark untruncated. */
1666 if (!event_root.has_value()
1667 || gc.getGateType(*event_root) == gate_one) {
1668 return TruncatedSingleRv{*spec, nat_lo, nat_hi, /*truncated=*/false};
1669 }
1670
1671 /* Infeasible event resolved upstream by RangeCheck: the cmp was
1672 * folded to gate_zero, the conditional distribution is undefined.
1673 * @c collectRvConstraints would silently fall back to the natural
1674 * support here (its walker skips gate_zero like gate_one), so we
1675 * have to detect this explicitly. */
1676 if (gc.getGateType(*event_root) == gate_zero) return std::nullopt;
1677
1678 auto iv = collectRvConstraints(gc, *event_root, root);
1679 if (!iv.has_value()) return std::nullopt;
1680 if (!(iv->first < iv->second)) return std::nullopt;
1681
1682 return TruncatedSingleRv{*spec, iv->first, iv->second, /*truncated=*/true};
1683}
1684
1686 std::optional<gate_t> event_root)
1687{
1688 if (!event_root.has_value()) return false;
1689 const auto et = gc.getGateType(*event_root);
1690 if (et == gate_one) return false;
1691 /* RangeCheck folded the event to false upstream – universal
1692 * signal, independent of root gate type (a constant scalar
1693 * value paired with an impossible cmp lands here too). */
1694 if (et == gate_zero) return true;
1695 /* Walk the event's AND-conjuncts; an empty intersection with the
1696 * RV's natural support is the second infeasibility signal that
1697 * @c matchTruncatedSingleRv collapses into @c std::nullopt. Only
1698 * applicable when the root is itself a bare gate_rv that the
1699 * walker recognises. */
1700 if (gc.getGateType(root) != gate_rv) return false;
1701 auto iv = collectRvConstraints(gc, *event_root, root);
1702 if (!iv.has_value()) return false;
1703 return !(iv->first < iv->second);
1704}
1705
1706/**
1707 * @brief Unconditional probability mass of a shape over the
1708 * interval @c [lo, hi].
1709 *
1710 * @c TruncatedSingleRv arms supplied here must carry
1711 * @c truncated == @c false (the unconditional shape); the helper
1712 * uses the natural support to compute the CDF endpoints, so calling
1713 * with an already-truncated input would double-truncate.
1714 *
1715 * Recursive: a Bernoulli mixture's mass is the Bernoulli-weighted
1716 * combination of its arms' masses. Categorical mass is the sum of
1717 * outcome masses falling in the interval. Dirac mass is 1 iff the
1718 * Dirac value sits in the interval, else 0. Returns @c std::nullopt
1719 * when a leaf's spec defeats the closed-form CDF (e.g. non-integer
1720 * Erlang shape – @c cdfAt returns NaN there).
1721 */
1722static std::optional<double>
1723shape_mass(const ClosedFormShape &s, double lo, double hi)
1724{
1725 return std::visit([&](const auto &v) -> std::optional<double> {
1726 using T = std::decay_t<decltype(v)>;
1727 if constexpr (std::is_same_v<T, TruncatedSingleRv>) {
1728 const double a = std::max(lo, v.lo);
1729 const double b = std::min(hi, v.hi);
1730 if (!(a < b)) return 0.0;
1731 const double cl = std::isfinite(a) ? cdfAt(v.spec, a) : 0.0;
1732 const double ch = std::isfinite(b) ? cdfAt(v.spec, b) : 1.0;
1733 if (std::isnan(cl) || std::isnan(ch)) return std::nullopt;
1734 return ch - cl;
1735 } else if constexpr (std::is_same_v<T, DiracShape>) {
1736 return (v.value >= lo && v.value <= hi) ? 1.0 : 0.0;
1737 } else if constexpr (std::is_same_v<T, CategoricalShape>) {
1738 double m = 0.0;
1739 for (const auto &pr : v.outcomes)
1740 if (pr.first >= lo && pr.first <= hi) m += pr.second;
1741 return m;
1742 } else if constexpr (std::is_same_v<T, BernoulliMixtureShape>) {
1743 auto L = shape_mass(*v.left, lo, hi);
1744 auto R = shape_mass(*v.right, lo, hi);
1745 if (!L || !R) return std::nullopt;
1746 return v.p * (*L) + (1.0 - v.p) * (*R);
1747 }
1748 return std::nullopt;
1749 }, s);
1750}
1751
1752/**
1753 * @brief Conditional shape after truncating the underlying variable
1754 * to @c [lo, hi].
1755 *
1756 * Bare-RV arm: intersects its natural / current truncation with
1757 * @c [lo, hi] and marks the result truncated so downstream
1758 * @c shape_pdf renormalises by the truncated CDF. Dirac: keep iff
1759 * value ∈ interval, otherwise nullopt (infeasible). Categorical:
1760 * keep outcomes in interval, renormalise masses. Bernoulli mixture:
1761 * recursively truncate each arm and reweight the Bernoulli by the
1762 * ratio of arm masses (the standard
1763 * @f$ \pi' = \pi Z_L / (\pi Z_L + (1-\pi) Z_R) @f$ update); a
1764 * fully-eliminated arm degenerates to the surviving one. Returns
1765 * @c nullopt when the truncated shape has zero mass (caller can
1766 * raise infeasibility).
1767 */
1768static std::optional<ClosedFormShape>
1769truncateShape(const ClosedFormShape &s, double lo, double hi)
1770{
1771 return std::visit([&](const auto &v) -> std::optional<ClosedFormShape> {
1772 using T = std::decay_t<decltype(v)>;
1773 if constexpr (std::is_same_v<T, TruncatedSingleRv>) {
1774 const double a = std::max(lo, v.lo);
1775 const double b = std::min(hi, v.hi);
1776 if (!(a < b)) return std::nullopt;
1777 return ClosedFormShape{TruncatedSingleRv{v.spec, a, b, /*trunc=*/true}};
1778 } else if constexpr (std::is_same_v<T, DiracShape>) {
1779 if (v.value < lo || v.value > hi) return std::nullopt;
1780 return ClosedFormShape{v};
1781 } else if constexpr (std::is_same_v<T, CategoricalShape>) {
1782 CategoricalShape out;
1783 double total = 0.0;
1784 for (const auto &pr : v.outcomes) {
1785 if (pr.first >= lo && pr.first <= hi) {
1786 out.outcomes.emplace_back(pr.first, pr.second);
1787 total += pr.second;
1788 }
1789 }
1790 if (out.outcomes.empty() || !(total > 0.0)) return std::nullopt;
1791 for (auto &pr : out.outcomes) pr.second /= total;
1792 return ClosedFormShape{std::move(out)};
1793 } else if constexpr (std::is_same_v<T, BernoulliMixtureShape>) {
1794 auto mL = shape_mass(*v.left, lo, hi);
1795 auto mR = shape_mass(*v.right, lo, hi);
1796 if (!mL || !mR) return std::nullopt;
1797 const double pL = v.p * (*mL);
1798 const double pR = (1.0 - v.p) * (*mR);
1799 const double Z = pL + pR;
1800 if (!(Z > 0.0)) return std::nullopt;
1801 auto Lt = truncateShape(*v.left, lo, hi);
1802 auto Rt = truncateShape(*v.right, lo, hi);
1803 /* Either arm eliminated by the truncation collapses to the
1804 * surviving arm (its mass was already 0 in shape_mass, so the
1805 * reweighted p_arm is 1). */
1806 if (!Lt && !Rt) return std::nullopt;
1807 if (!Lt) return Rt;
1808 if (!Rt) return Lt;
1810 m.p = pL / Z;
1811 m.left = std::make_shared<ClosedFormShape>(std::move(*Lt));
1812 m.right = std::make_shared<ClosedFormShape>(std::move(*Rt));
1813 return ClosedFormShape{std::move(m)};
1814 }
1815 return std::nullopt;
1816 }, s);
1817}
1818
1819std::optional<ClosedFormShape>
1821 std::optional<gate_t> event_root)
1822{
1823 /* Test "event is trivial true": either absent, or resolved to
1824 * gate_one by load-time simplification. */
1825 const bool event_trivial = !event_root.has_value()
1826 || gc.getGateType(*event_root) == gate_one;
1827
1828 /* Bare gate_rv root: delegate to the existing single-RV matcher
1829 * so the truncation logic (collectRvConstraints) is the single
1830 * source of truth across the closed-form-shape surface. */
1831 if (gc.getGateType(root) == gate_rv) {
1832 /* Conjugate observe-evidence: the posterior is itself a bare
1833 * distribution of the prior's family, so the shape is the
1834 * (untruncated) posterior -- exact pdf/CDF for the histogram and
1835 * curve renderers. Declines to the truncation matcher on any
1836 * mismatch (whose AND-conjunct walker treats a gate_observe as an
1837 * uninterpretable factor and declines in turn). */
1838 if (!event_trivial)
1839 if (auto post = conjugatePosterior(gc, root, *event_root)) {
1840 const DistSupport sup = makeDistribution(*post)->support();
1841 return ClosedFormShape{
1842 TruncatedSingleRv{*post, sup.lo, sup.hi, /*truncated=*/false}};
1843 }
1844 auto m = matchTruncatedSingleRv(gc, root, event_root);
1845 if (!m) return std::nullopt;
1846 return ClosedFormShape{*m};
1847 }
1848
1849 /* Helper: match the shape unconditionally first, then if the event
1850 * is non-trivial extract an interval via collectRvConstraints and
1851 * apply truncateShape. Used by the Dirac / categorical / mixture
1852 * branches below so all three honour conditioning through the same
1853 * pipeline. */
1854 auto with_optional_truncation =
1855 [&](std::optional<ClosedFormShape> unc)
1856 -> std::optional<ClosedFormShape> {
1857 if (!unc) return std::nullopt;
1858 if (event_trivial) return unc;
1859 auto iv = collectRvConstraints(gc, *event_root, root);
1860 if (!iv.has_value()) return std::nullopt;
1861 if (!(iv->first < iv->second)) return std::nullopt;
1862 return truncateShape(*unc, iv->first, iv->second);
1863 };
1864
1865 /* Dirac point: a gate_value with extra parseable as a finite
1866 * float8 (the underlying form of as_random(c)). Conditioning on
1867 * a constant is normally folded upstream by RangeCheck to
1868 * gate_one / gate_zero, but a probabilistic event whose footprint
1869 * doesn't constrain the constant lands here untouched (the cmp
1870 * walker returns the unconditional support); truncateShape then
1871 * keeps the Dirac iff its value falls in the recognised interval. */
1872 if (gc.getGateType(root) == gate_value) {
1873 double v;
1874 if (!extract_finite_double(gc, root, v)) return std::nullopt;
1875 return with_optional_truncation(ClosedFormShape{DiracShape{v}});
1876 }
1877
1878 /* gate_mixture: either the explicit categorical form
1879 * (isCategoricalMixture) or the classic Bernoulli triple
1880 * [p_token, x_token, y_token]. */
1881 if (gc.getGateType(root) == gate_mixture) {
1882 const auto &w = gc.getWires(root);
1883
1884 if (gc.isCategoricalMixture(root)) {
1886 cs.outcomes.reserve(w.size() - 1);
1887 for (std::size_t i = 1; i < w.size(); ++i) {
1888 double v;
1889 if (!extract_mulinput_value(gc, w[i], v)) return std::nullopt;
1890 double p = gc.getProb(w[i]);
1891 if (!std::isfinite(p) || p < 0.0 || p > 1.0) return std::nullopt;
1892 cs.outcomes.emplace_back(v, p);
1893 }
1894 if (cs.outcomes.empty()) return std::nullopt;
1895 return with_optional_truncation(ClosedFormShape{std::move(cs)});
1896 }
1897
1898 /* Classic Bernoulli mixture: 3 wires, [p_token, x_token, y_token]
1899 * with p_token a bare gate_input; compound Boolean p bails (the
1900 * generic path would need a probability-over-Boolean-circuit
1901 * pre-pass we deliberately do not run here). */
1902 if (w.size() != 3) return std::nullopt;
1903 if (gc.getGateType(w[0]) != gate_input) return std::nullopt;
1904 double p = gc.getProb(w[0]);
1905 if (!std::isfinite(p) || p < 0.0 || p > 1.0) return std::nullopt;
1906
1907 auto left = matchClosedFormDistribution(gc, w[1], std::nullopt);
1908 auto right = matchClosedFormDistribution(gc, w[2], std::nullopt);
1909 if (!left || !right) return std::nullopt;
1910
1912 m.p = p;
1913 m.left = std::make_shared<ClosedFormShape>(std::move(*left));
1914 m.right = std::make_shared<ClosedFormShape>(std::move(*right));
1915 return with_optional_truncation(ClosedFormShape{std::move(m)});
1916 }
1917
1918 return std::nullopt;
1919}
1920
1921} // namespace provsql
1922
1923extern "C" {
1924
1925/**
1926 * @brief SQL: rv_support(token uuid, prov uuid, OUT lo float8, OUT hi float8)
1927 *
1928 * Loads the persisted circuit rooted at @p token, intersects with the
1929 * AND-conjunct cmps in @p prov constraining @p token, and returns the
1930 * resulting @c [lo, hi] support interval. When @p prov resolves to
1931 * @c gate_one (the unconditional default after load-time
1932 * simplification), the conditional path is skipped and the bare
1933 * unconditional support of @p token is returned.
1934 *
1935 * @c -Infinity / @c +Infinity float8 represent unbounded ends (e.g.
1936 * the support of a normal RV is @c [-Infinity, +Infinity]).
1937 */
1938Datum rv_support(PG_FUNCTION_ARGS)
1939{
1940 try {
1941 pg_uuid_t *token = PG_GETARG_UUID_P(0);
1942 pg_uuid_t *prov = PG_GETARG_UUID_P(1);
1943
1944 gate_t root_gate, event_gate;
1945 auto gc = getJointCircuit(*token, *prov, root_gate, event_gate);
1946
1947 /* gate_one as event-side means the conditioning is the trivial
1948 * "always true" event (either the user passed gate_one() directly
1949 * or load-time simplification collapsed the event to it). Take
1950 * the unconditional path. */
1951 std::optional<gate_t> event_opt;
1952 if (gc.getGateType(event_gate) != gate_one)
1953 event_opt = event_gate;
1954
1955 /* A stored "X | C" arrives here as a conditioned root: peel it to the
1956 * bare target and fold the condition into the event, so the support is
1957 * the conditional (truncated) one rather than the unconditional. */
1958 root_gate = provsql::lift_conditioning(gc, root_gate, event_opt);
1959
1960 auto iv = provsql::compute_support(gc, root_gate, event_opt);
1961
1962 TupleDesc tupdesc;
1963 Datum values[2];
1964 bool nulls[2] = {false, false};
1965
1966 if (get_call_result_type(fcinfo, NULL, &tupdesc) != TYPEFUNC_COMPOSITE)
1967 provsql_error("rv_support: expected composite return type");
1968 tupdesc = BlessTupleDesc(tupdesc);
1969
1970 values[0] = Float8GetDatum(iv.first);
1971 values[1] = Float8GetDatum(iv.second);
1972
1973 PG_RETURN_DATUM(HeapTupleGetDatum(heap_form_tuple(tupdesc, values, nulls)));
1974 } catch (const std::exception &e) {
1975 provsql_error("rv_support: %s", e.what());
1976 } catch (...) {
1977 provsql_error("rv_support: unknown exception");
1978 }
1979 PG_RETURN_NULL();
1980}
1981
1982} // extern "C"
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.
Typed aggregation value, operator, and aggregator abstractions.
AggregationOperator
SQL aggregation functions tracked by ProvSQL.
Definition Aggregation.h:51
@ MAX
MAX → input type.
Definition Aggregation.h:55
@ COUNT
COUNT(*) or COUNT(expr) → integer.
Definition Aggregation.h:52
@ SUM
SUM → integer or float.
Definition Aggregation.h:53
@ MIN
MIN → input type.
Definition Aggregation.h:54
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
Closed-form CDF resolution for trivial gate_cmp shapes.
static CircuitCache cache
Process-local singleton circuit gate cache.
GenericCircuit getJointCircuit(const std::vector< pg_uuid_t > &tokens, std::vector< gate_t > &gates)
Multi-root variant of getJointCircuit.
Build in-memory circuits from the mmap-backed persistent store.
gate_t
Strongly-typed gate identifier.
Definition Circuit.h:49
Exact conjugate-prior posteriors for observe-evidence circuits.
Per-family polymorphic view over a continuous gate_rv distribution (§F.1 class hierarchy).
Analytical expectation / variance / moment evaluator over RV circuits.
Continuous random-variable helpers (distribution parsing, moments).
Datum rv_support(PG_FUNCTION_ARGS)
SQL: rv_support(token uuid, prov uuid, OUT lo float8, OUT hi float8).
Support-based bound check for continuous-RV comparators.
iterator end()
Past-the-end iterator for the cache.
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.
void resolveGateToZero(gate_t g)
Replace an arbitrary gate (typically gate_times) by gate_zero.
void resolveCmpToPlusOfKGates(gate_t g, const std::vector< gate_t > &ks)
Replace a gate_cmp by a gate_plus over the given per-row K-gates (the OR of the agg's row-presence in...
bool isCategoricalMixture(gate_t g) const
Test whether g is a categorical-form gate_mixture (the explicit provsql.categorical output).
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.
void resolveCmpToBernoulli(gate_t g, double p)
Replace a gate_cmp by a constant Boolean leaf (gate_one for p == 1, gate_zero for p == 0) or by a Ber...
std::pair< unsigned, unsigned > getInfos(gate_t g) const
Return the integer annotation pair for gate g.
static pg_uuid_t value_gate(const char *prefix, const char *suffix, const char *extra)
The value gate of the text str, at the address named name ("value" followed by str,...
std::pair< double, double > compute_support(const GenericCircuit &gc, gate_t root, std::optional< gate_t > event_root)
Compute the [lo, hi] support interval of a scalar sub-circuit rooted at root.
static std::optional< ClosedFormShape > truncateShape(const ClosedFormShape &s, double lo, double hi)
Conditional shape after truncating the underlying variable to [lo, hi].
std::optional< ClosedFormShape > matchClosedFormDistribution(const GenericCircuit &gc, gate_t root, std::optional< gate_t > event_root)
Detect any of the closed-form shapes supported by rv_analytical_curves.
std::variant< TruncatedSingleRv, DiracShape, CategoricalShape, BernoulliMixtureShape > ClosedFormShape
One of the closed-form shapes the analytical-curves payload can render: bare RV (continuous PDF/CDF),...
Definition RangeCheck.h:201
gate_t lift_conditioning(GenericCircuit &gc, gate_t root, std::optional< gate_t > &event_opt)
Lift conditioning out of a scalar arithmetic expression.
static bool extract_mulinput_value(const GenericCircuit &gc, gate_t mul, double &out)
Same parsing applied to a mulinput's outcome label (categorical).
static bool extract_finite_double(const GenericCircuit &gc, gate_t x, double &out)
Parse a gate_value's extra as a finite float8.
unsigned runReflexiveCmpRewriter(GenericCircuit &gc)
Probability-side pre-pass: rewrite HAVING-style gate_cmp gates that are provably TRUE on the agg's va...
double parseDoubleStrict(const std::string &s)
Strictly parse s as a double.
unsigned runRangeCheck(GenericCircuit &gc)
Run the support-based pruning pass over gc.
bool eventIsProvablyInfeasible(const GenericCircuit &gc, gate_t root, std::optional< gate_t > event_root)
True iff the conditioning event is provably infeasible for a bare gate_rv root.
std::unique_ptr< Distribution > makeDistribution(const DistributionSpec &spec)
Construct the per-family Distribution for a parsed spec.
static std::optional< double > shape_mass(const ClosedFormShape &s, double lo, double hi)
Unconditional probability mass of a shape over the interval [lo, hi].
std::optional< DistributionSpec > conjugatePosterior(const GenericCircuit &gc, gate_t target, gate_t evidence)
The exact posterior of target given evidence, as a resolved distribution spec, when the circuit match...
std::optional< DistributionSpec > parse_distribution_spec(const std::string &s)
Parse the on-disk text encoding of a gate_rv distribution.
std::optional< std::pair< double, double > > collectRvConstraints(const GenericCircuit &gc, gate_t event_root, gate_t target_rv)
Walk event_root collecting rv op c constraints on target_rv.
std::optional< DistributionTemplate > parse_distribution_template(const std::string &s)
Parse the on-disk text encoding of a gate_rv distribution, keeping wired (token) parameters as wire r...
std::optional< TruncatedSingleRv > matchTruncatedSingleRv(const GenericCircuit &gc, gate_t root, std::optional< gate_t > event_root)
Detect a closed-form, optionally-truncated single-RV shape.
double cdfAt(const DistributionSpec &d, double c)
Closed-form CDF for a basic continuous distribution.
unsigned runHavingAlwaysTrueRewriter(GenericCircuit &gc)
Uniform error-reporting macros for ProvSQL.
#define provsql_error(fmt,...)
Report a fatal ProvSQL error and abort the current 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
@ gate_rv
Continuous random-variable leaf (extra encodes distribution).
@ gate_case
N-ary guarded selection over scalar (RV) children: wires are [guard_1, value_1, .....
@ gate_mixture
Probabilistic mixture: three wires [p_token (gate_input Bernoulli), x_token, y_token]; samples x when...
@ 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...
C++ utility functions for UUID manipulation.
Bernoulli mixture (gate_mixture with the [p_token, x_token, y_token] shape).
Definition RangeCheck.h:218
std::shared_ptr< ClosedFormShape > right
Definition RangeCheck.h:221
std::shared_ptr< ClosedFormShape > left
Definition RangeCheck.h:220
Categorical distribution over a finite outcome set.
Definition RangeCheck.h:189
std::vector< std::pair< double, double > > outcomes
(value, mass) pairs
Definition RangeCheck.h:190
Point mass at a finite scalar value (a gate_value root, or an as_random(c) leaf surfaced as a gate_va...
Definition RangeCheck.h:173
A closed support interval [lo, hi] (±infinity for unbounded).
Detection result for a closed-form, optionally-truncated single-RV shape.
Definition RangeCheck.h:102