From 5ee89d0797e1f7d363014cbe4afd4156a7b940cf Mon Sep 17 00:00:00 2001 From: Esteban Zimanyi Date: Sun, 11 Oct 2026 01:08:38 +0200 Subject: [PATCH] State what a wrapper binds for each number of arguments its tests of PG_NARGS() decide A wrapper can choose what it reads by any test of PG_NARGS(), and by the else of such a test: Temporal_app_tinst_transfn starts maxdist from -1.0 and maxt from NULL and, under if (PG_NARGS() > 3), reads maxt from argument 3 under if (PG_NARGS() == 4) and, in its else, maxdist from argument 3 and maxt from argument 4. _guards of parser/boundargs.py states for each statement the numbers of arguments it runs for: every conjunct testing PG_NARGS() with ==, !=, <, <=, > or >= narrows them, the else of a test of PG_NARGS() alone runs for the others, nested tests intersect, and comments are read as blanks. A local starting from a literal is that literal for a signature of a number of arguments none of its later assignments runs for (_guarded_default), and _caller_index reads, for a signature, only the assignments a call of its number of arguments runs. appendInstantTransition(tbool, tbool, text, interval) therefore binds maxdist to -1.0, and its argument 3 left to a NULL default passes NULL for maxt, where the catalog stated no bind and paired argument 3 with maxdist. The floor of .github/workflows/pytest.yml rises to 536, the suite collecting the two tests of ArityBranchTests, which read the wrapper of Temporal_app_tinst_transfn. Witness: over MEOS-API ffa1f6f the seven signatures appendInstantTransition(T, T, text, interval) state no boundArgs, so a binding pairs their three SQL arguments after the state with the four parameters interp, maxdist and maxt and finds no overload; tpcpoint's states nullDefaultBinds {"3": {"maxdist": "-1.0"}} for its interval argument. Measured over MobilityDB master 3a8e6dba09, run.py on upstream ffa1f6f and on this branch with the same headers: the two catalogs differ in temporal_app_tinst_transfn alone, the seven signatures of four arguments binding maxdist to -1.0, tpcpoint's argument 3 passing NULL for maxt, and the nine of five arguments stating that argument 4 passes NULL for maxt. The suite collects 536 tests, 451 passing and 85 skipped without a catalog in output/. Why: a binding calls appendInstantAgg with an interpolation and a gap interval as PostgreSQL does only when the catalog states the distance the wrapper passes for that call. --- .github/workflows/pytest.yml | 2 +- parser/boundargs.py | 215 +++++++++++++++++++++++++---------- tests/test_boundargs.py | 101 ++++++++++++++++ 3 files changed, 258 insertions(+), 60 deletions(-) diff --git a/.github/workflows/pytest.yml b/.github/workflows/pytest.yml index 7a3cc40..496ed3a 100644 --- a/.github/workflows/pytest.yml +++ b/.github/workflows/pytest.yml @@ -96,7 +96,7 @@ jobs: # carries, or a change to them is not exercised until after it merges. # Consumers use the action; this repository owns the rules. - name: Refuse a skip, and a suite that shrank - run: tools/check-test-outcome.py "$RUNNER_TEMP/pytest.log" --min-tests 534 + run: tools/check-test-outcome.py "$RUNNER_TEMP/pytest.log" --min-tests 536 # The rules earn their place by refusing a log that carries what they # name. Both fixtures are written here rather than tracked, and the diff --git a/parser/boundargs.py b/parser/boundargs.py index fce7410..c77011f 100644 --- a/parser/boundargs.py +++ b/parser/boundargs.py @@ -38,8 +38,12 @@ ``OUT_DEFAULT_DECIMAL_DIGITS`` and reads argument 1 only under ``if (PG_NARGS() > 1 && ! PG_ARGISNULL(1))``, so ``asEWKT(th3index)`` calls ``tspatial_as_ewkt(temp, OUT_DEFAULT_DECIMAL_DIGITS)``. Such a local is recorded on -each SQL signature stating at most ``k`` arguments, where it is the literal the MEOS -call reads; a signature stating argument ``k`` passes the caller's value. +each SQL signature of a number of arguments for which the wrapper never replaces it, +where it is the literal the MEOS call reads; a signature of another number passes the +caller's value. A test of ``PG_NARGS()`` decides the numbers each statement runs for, +an ``else`` running for the others: ``Temporal_app_tinst_transfn`` reads ``maxt`` from +argument 3 under ``PG_NARGS() == 4`` and ``maxdist`` from argument 3 in its ``else``, so +``appendInstantTransition(tint, tint, text, interval)`` binds ``maxdist`` to ``-1.0``. A wrapper can reach the functions its ``@csqlfn`` tags name through an internal generic that none of them is: ``Numset_shift`` calls ``numset_shift_scale(s, shift, 0, true, @@ -58,6 +62,7 @@ """ from __future__ import annotations +import functools import re from pathlib import Path @@ -171,11 +176,11 @@ def extract_helpers(mdb_src: str | Path) -> dict[str, tuple[str, list[str]]]: def _delegated(body: str, helpers: dict[str, tuple[str, list[str]]]): - """``(helper_body, {helper_param: literal}, {helper_param: (k, literal)})`` when + """``(helper_body, {helper_param: literal}, {helper_param: (arities, literal)})`` when ``body`` delegates to a shared helper, passing the call info through: the helper parameters the wrapper binds to literals, and those it binds to a local it reads from - argument ``k`` only when the call carries it (#_guarded_default), the literal the - local starts from being what a signature omitting argument ``k`` passes on.""" + an argument only for the numbers of arguments ``arities`` (#_guarded_default), the + literal the local starts from being what a signature of another number passes on.""" for m in _DELEG.finditer(body): entry = helpers.get(m.group("name")) if entry is None: @@ -238,41 +243,126 @@ def _literal(arg: str) -> str | None: return None -def _guards(body: str) -> list[tuple[int, int, int]]: - """``(k, start, end)`` for every statement ``body`` runs only under - ``if (PG_NARGS() > k ...)``: the span of the brace block or of the single statement - the test guards.""" - out: list[tuple[int, int, int]] = [] - for m in re.finditer(r"\bif\s*\(", body): - depth, i = 0, m.end() - 1 - for i in range(m.end() - 1, len(body)): - if body[i] == "(": - depth += 1 - elif body[i] == ")": - depth -= 1 - if depth == 0: - break - g = re.match(r"\s*PG_NARGS\s*\(\s*\)\s*>\s*(\d+)\s*(?:&&|$)", - body[m.end():i]) - if not g: +# The numbers of arguments a SQL function can be called with: PostgreSQL's FUNC_MAX_ARGS. +_ARITIES = frozenset(range(101)) +# A test of the number of arguments of the call: `PG_NARGS() == 4`. +_NARGS_TEST = re.compile(r"PG_NARGS\s*\(\s*\)\s*(==|!=|>=|<=|>|<)\s*(\d+)") +_NARGS_OPS = {"==": lambda n, k: n == k, "!=": lambda n, k: n != k, + ">=": lambda n, k: n >= k, "<=": lambda n, k: n <= k, + ">": lambda n, k: n > k, "<": lambda n, k: n < k} + + +def _arities(op: str, k: int) -> frozenset: + """The numbers of arguments for which ``PG_NARGS() k`` holds.""" + return frozenset(n for n in _ARITIES if _NARGS_OPS[op](n, k)) + + +def _split_top(expr: str, op: str) -> list[str]: + """``expr`` split on the operator ``op`` where it sits outside every parenthesis.""" + out, depth, start, i = [], 0, 0, 0 + while i < len(expr): + c = expr[i] + if c == "(": + depth += 1 + elif c == ")": + depth -= 1 + elif depth == 0 and expr.startswith(op, i): + out.append(expr[start:i]) + start = i + len(op) + i = start continue - j = i + 1 - while j < len(body) and body[j].isspace(): - j += 1 - if body.startswith("{", j): - end = j + len(_body(body, j)) + 2 - else: - end = body.find(";", j) + 1 - out.append((int(g.group(1)), j, end)) + i += 1 + out.append(expr[start:]) return out -def _guarded_default(body: str, var: str) -> tuple[int, str] | None: - """``(k, literal)`` when local ``var`` starts from a literal and every later assignment - of it sits under ``if (PG_NARGS() > k ...)``. A SQL signature omitting argument ``k`` - never runs those assignments, so the MEOS call reads the literal, as a SQL DEFAULT - would supply it: ``Tspatial_as_text_common`` starts ``dbl_dig_for_wkt`` from - ``OUT_DEFAULT_DECIMAL_DIGITS`` and reads argument 1 only when the call carries it.""" +def _close_paren(text: str, open_pos: int) -> int: + """The index of the parenthesis closing the one at ``open_pos``.""" + depth = 0 + for i in range(open_pos, len(text)): + if text[i] == "(": + depth += 1 + elif text[i] == ")": + depth -= 1 + if depth == 0: + return i + return len(text) - 1 + + +def _skip_space(text: str, i: int) -> int: + while i < len(text) and text[i].isspace(): + i += 1 + return i + + +def _statement_end(body: str, j: int) -> int: + """The end of the statement starting at ``j``: a brace block, an ``if`` with its + ``else``, or a single statement up to its semicolon.""" + if body.startswith("{", j): + return j + len(_body(body, j)) + 2 + m = re.match(r"if\s*\(", body[j:]) + if m: + end = _statement_end(body, _skip_space(body, _close_paren(body, j + m.end() - 1) + 1)) + e = _skip_space(body, end) + if re.match(r"else\b", body[e:]): + return _statement_end(body, _skip_space(body, e + 4)) + return end + return body.find(";", j) + 1 + + +@functools.lru_cache(maxsize=None) +def _guards(body: str) -> tuple[tuple[frozenset, int, int], ...]: + """``(arities, start, end)`` for every statement ``body`` runs only for the numbers of + arguments ``arities``: the span of the brace block or of the single statement an + ``if`` guards, the test of which states ``PG_NARGS()`` in a conjunct, and of the + ``else`` of a test that states nothing but ``PG_NARGS()``, which runs for the other + numbers. ``Temporal_app_tinst_transfn`` reads ``maxt`` from argument 3 under + ``if (PG_NARGS() == 4)`` and ``maxdist`` from argument 3 in its ``else``, both under + ``if (PG_NARGS() > 3)``. Comments are read as the blanks #_COMMENT replaces them by, + one for each character, so ``else /* PG_NARGS() == 5 */`` opens its block after the + comment and every span stays a span of ``body``.""" + body = _COMMENT.sub(lambda c: " " * len(c.group()), body) + out: list[tuple[frozenset, int, int]] = [] + for m in re.finditer(r"\bif\s*\(", body): + close = _close_paren(body, m.end() - 1) + cond = body[m.end():close] + if len(_split_top(cond, "||")) > 1: + continue + conjuncts = [c.strip() for c in _split_top(cond, "&&")] + tests = [_NARGS_TEST.fullmatch(c) for c in conjuncts] + if not any(tests): + continue + runs = _ARITIES + for t in tests: + if t: + runs = runs & _arities(t.group(1), int(t.group(2))) + j = _skip_space(body, close + 1) + end = _statement_end(body, j) + out.append((runs, j, end)) + e = _skip_space(body, end) + if all(tests) and re.match(r"else\b", body[e:]): + k = _skip_space(body, e + 4) + out.append((_ARITIES - runs, k, _statement_end(body, k))) + return tuple(out) + + +def _runs_at(body: str, pos: int) -> frozenset: + """The numbers of arguments for which the statement of ``body`` at ``pos`` runs.""" + runs = _ARITIES + for r, s, e in _guards(body): + if s <= pos < e: + runs = runs & r + return runs + + +def _guarded_default(body: str, var: str) -> tuple[frozenset, str] | None: + """``(arities, literal)`` when local ``var`` starts from a literal and every later + assignment of it runs only for the numbers of arguments ``arities`` (#_guards). A SQL + signature of another number of arguments never runs those assignments, so the MEOS + call reads the literal, as a SQL DEFAULT would supply it: ``Tspatial_as_text_common`` + starts ``dbl_dig_for_wkt`` from ``OUT_DEFAULT_DECIMAL_DIGITS`` and reads argument 1 + only when the call carries it, and ``Temporal_app_tinst_transfn`` starts ``maxdist`` + from ``-1.0`` and reads it only for five arguments.""" v = re.escape(var) if re.search(r"&\s*" + v + r"\b|(?])" + v + r"\s*(?:\+\+|--|(?:[-+*/%&|^]|<<|>>)=)|(?:\+\+|--)\s*" + v + r"\b", @@ -281,18 +371,17 @@ def _guarded_default(body: str, var: str) -> tuple[int, str] | None: assigns = list(re.finditer(r"(?])" + v + r"\s*=(?!=)", body)) if len(assigns) < 2: return None - guards = _guards(body) init = re.match(r"\s*([^;]+?)\s*;", body[assigns[0].end():]) lit = _literal(init.group(1)) if init else None - if lit is None or any(s <= assigns[0].start() < e for _, s, e in guards): + if lit is None or _runs_at(body, assigns[0].start()) != _ARITIES: return None - ks = set() + runs: frozenset = frozenset() for a in assigns[1:]: - hit = [k for k, s, e in guards if s <= a.start() < e] - if not hit: + r = _runs_at(body, a.start()) + if r == _ARITIES: return None - ks.update(hit) - return (ks.pop(), lit) if len(ks) == 1 else None + runs = runs | r + return runs, lit def _wrapper_bound(body: str, func: dict, drift: list, @@ -303,9 +392,9 @@ def _wrapper_bound(body: str, func: dict, drift: list, """The literals wrapper ``body`` binds in its call to ``func['name']``, keyed by ``func``'s parameter name. Empty if the wrapper does not call ``func`` by name. - A local the wrapper reads from argument ``k`` only when the call carries it - (``_guarded_default``) is caller-sourced for a signature stating ``k`` and a literal - for one omitting it; it goes into ``guarded`` as ``{param: (k, literal)}``. When + A local the wrapper reads from an argument only for the numbers of arguments + ``arities`` (``_guarded_default``) is caller-sourced for a signature of one of them and + a literal for any other; it goes into ``guarded`` as ``{param: (arities, literal)}``. When ``body`` is a helper the wrapper delegates to, ``gsubst`` holds the helper parameters such a local reaches (#_delegated), and a call argument naming one of them goes into ``guarded`` the same way. @@ -373,7 +462,8 @@ def _wrapper_bound(body: str, func: dict, drift: list, def _group_bound(body: str, group: list, helpers: dict, drift: list, documented: dict[str, set], generics: dict | None = None): """``(bound, guarded)``: the literals wrapper ``body`` binds, keyed by parameter name, - and the ``{param: (k, literal)}`` it supplies when the call omits argument ``k``, read + and the ``{param: (arities, literal)}`` it supplies for a call of a number of arguments + outside ``arities``, read from its call to whichever member of ``group`` it names (branches such as the RGEO ternary agree, and the first wins), or from its delegation to a shared helper when it names none, or from its call to the generic the members wrap when it does neither -- @@ -515,8 +605,8 @@ def own(w, sig=None): out = {k: v for k, v in bound.items() if k in pnames} if sig is not None: nargs = len(sig.get("args") or ()) - out.update({k: lit for k, (pos, lit) in guarded.items() - if k in pnames and nargs <= pos}) + out.update({k: lit for k, (runs, lit) in guarded.items() + if k in pnames and nargs not in runs}) return out sigs = func.get("sqlSignatures") or [] @@ -624,13 +714,16 @@ def attach_type_derived_args(idl: dict, mdb_src: str | Path, def _null_value(body: str, var: str, k: int, sig: dict, linear: dict) -> str | None: """The literal local ``var`` of wrapper ``body`` holds when SQL argument ``k`` is NULL, - or None when the wrapper states none: the literal it starts from and replaces only under - ``PG_NARGS() > k`` (#_guarded_default), the interpolation of the signature's temporal - type it takes when the call carries no ``k`` (#_type_choice), or the literal of + or None when the wrapper states none: the literal it starts from where every assignment + replacing it for the number of arguments of ``sig`` reads argument ``k`` + (#_guarded_default), the interpolation of the signature's temporal type it takes when + the call carries no ``k`` (#_type_choice), or the literal of ``var = PG_ARGISNULL(k) ? literal : ...``.""" g = _guarded_default(body, var) - if g is not None and g[0] == k: - return g[1] + if g is not None: + nargs = len(sig.get("args") or ()) + if nargs in g[0] and _caller_index(body, var, nargs=nargs) == k: + return g[1] c = _type_choice(body, var) if c is not None and c[0] == k: _, source, i, lin, other = c @@ -689,7 +782,7 @@ def attach_null_default_binds(idl: dict, mdb_src: str | Path, sql_src: str | Pat fed = {} for pos, a in enumerate(args): if pos < len(params) and _IDENT.match(a): - k = _caller_index(body, a) + k = _caller_index(body, a, nargs=len(sig.get("args") or ())) if k is not None: fed.setdefault(k, []).append((params[pos], a)) binds = {} @@ -740,14 +833,17 @@ def _directly_carried(body: str, args: list[str], assigned: set) -> set[int]: return out -def _caller_index(body: str, arg: str, depth: int = 0) -> int | None: +def _caller_index(body: str, arg: str, depth: int = 0, nargs: int | None = None) -> int | None: """The SQL argument the call argument ``arg`` of a wrapper ``body`` carries, or None. An argument reading ``PG_GETARG_(k)`` carries ``k``. A local carries the argument its assignments read directly, a later ``interp = input_interp_string(fcinfo, 1)`` over the default it starts from; a local read from other locals only, as ``instants`` from ``temparr_extract(array, &count)``, carries what those locals carry, as - #_wrapper_bound reads an assigned local as caller-sourced. One argument, else None.""" + #_wrapper_bound reads an assigned local as caller-sourced. One argument, else None. + Given ``nargs``, only the assignments a call of that many arguments runs are read + (#_guards): ``Temporal_app_tinst_transfn`` reads ``maxt`` from argument 3 for four + arguments and from argument 4 for five.""" arg = _CAST.sub("", arg.strip()) found = _direct_indices(arg) if found: @@ -755,7 +851,8 @@ def _caller_index(body: str, arg: str, depth: int = 0) -> int | None: if not _IDENT.match(arg) or depth > 3: return None rhs = [m.group(1) for m in - re.finditer(r"(?])" + re.escape(arg) + r"\s*=(?!=)\s*([^;]+);", body)] + re.finditer(r"(?])" + re.escape(arg) + r"\s*=(?!=)\s*([^;]+);", body) + if nargs is None or nargs in _runs_at(body, m.start())] direct = set().union(*(_direct_indices(r) for r in rhs)) if rhs else set() if direct: return direct.pop() if len(direct) == 1 else None @@ -763,7 +860,7 @@ def _caller_index(body: str, arg: str, depth: int = 0) -> int | None: for r in rhs: for ident in set(re.findall(r"\b[a-z_]\w*\b", r)) - {arg}: if re.search(r"(?])" + re.escape(ident) + r"\s*=(?!=)", body): - k = _caller_index(body, ident, depth + 1) + k = _caller_index(body, ident, depth + 1, nargs) if k is not None: via.add(k) return via.pop() if len(via) == 1 else None diff --git a/tests/test_boundargs.py b/tests/test_boundargs.py index d58c708..2aa0137 100644 --- a/tests/test_boundargs.py +++ b/tests/test_boundargs.py @@ -1292,5 +1292,106 @@ def test_a_strict_function_binds_none(self): self.assertNotIn("nullDefaultBinds", idl["functions"][0]["sqlSignatures"][1]) +ARITY_BRANCH = r''' +Datum +Temporal_app_tinst_transfn(PG_FUNCTION_ARGS) +{ + Temporal *state = PG_ARGISNULL(0) ? NULL : PG_GETARG_TEMPORAL_P(0); + Temporal *inst = PG_GETARG_TEMPORAL_P(1); + interpType interp; + if (PG_NARGS() == 2 || PG_ARGISNULL(2)) + { + MeosType temptype = oid_meostype(get_fn_expr_argtype(fcinfo->flinfo, 1)); + interp = temptype_supports_linear(temptype) ? LINEAR : STEP; + } + else + interp = input_interp_string(fcinfo, 2); + double maxdist = -1.0; + Interval *maxt = NULL; + if (PG_NARGS() > 3) + { + if (PG_NARGS() == 4) + { + if (! PG_ARGISNULL(3)) + maxt = PG_GETARG_INTERVAL_P(3); + } + else /* PG_NARGS() == 5 */ + { + if (! PG_ARGISNULL(3)) + maxdist = PG_GETARG_FLOAT8(3); + if (! PG_ARGISNULL(4)) + maxt = PG_GETARG_INTERVAL_P(4); + } + } + state = temporal_app_tinst_transfn(state, (TInstant *) inst, interp, maxdist, maxt); + PG_RETURN_TEMPORAL_P(state); +} +''' + + +class ArityBranchTests(unittest.TestCase): + """A wrapper choosing what it reads by a test of ``PG_NARGS()`` other than ``> k``, and + by the ``else`` of such a test, read as #GuardedDefaultTests and #NullDefaultBindTests + read one testing ``PG_NARGS() > k``: each signature binds what a call of its number of + arguments leaves to the literal a local starts from, and each NULL default is paired + with the parameter its argument feeds for that number.""" + + SQL = ''' +CREATE FUNCTION appendInstantTransition(tbool, tbool, interp text DEFAULT NULL, + maxt interval DEFAULT NULL) + RETURNS tbool + AS 'MODULE_PATHNAME', 'Temporal_app_tinst_transfn' + LANGUAGE C IMMUTABLE PARALLEL SAFE; +CREATE FUNCTION appendInstantTransition(tint, tint, interp text DEFAULT NULL, + maxdist float DEFAULT NULL, maxt interval DEFAULT NULL) + RETURNS tint + AS 'MODULE_PATHNAME', 'Temporal_app_tinst_transfn' + LANGUAGE C IMMUTABLE PARALLEL SAFE; +''' + + def setUp(self): + self.tmp = tempfile.TemporaryDirectory() + root = Path(self.tmp.name) + (root / "src").mkdir() + (root / "src" / "temporal_aggfuncs.c").write_text(ARITY_BRANCH) + (root / "sql").mkdir() + (root / "sql" / "040_temporal_aggfuncs.in.sql").write_text(self.SQL) + + def tearDown(self): + self.tmp.cleanup() + + def _func(self): + return {"name": "temporal_app_tinst_transfn", "mdbC": "Temporal_app_tinst_transfn", + "sqlfn": "appendInstantTransition", + "params": [{"name": n} for n in + ("state", "inst", "interp", "maxdist", "maxt")], + "sqlSignatures": [ + {"args": ["tbool", "tbool"], "ret": "tbool"}, + {"args": ["tbool", "tbool", "text"], "ret": "tbool"}, + {"args": ["tbool", "tbool", "text", "interval"], "ret": "tbool", + "argDefaults": [None, None, "NULL", "NULL"]}, + {"args": ["tint", "tint", "text", "float", "interval"], "ret": "tint", + "argDefaults": [None, None, "NULL", "NULL", "NULL"]}]} + + def test_each_number_of_arguments_binds_what_it_leaves(self): + idl, _, _ = merge_boundargs({"functions": [self._func()]}, self.tmp.name) + self.assertEqual( + [s.get("boundArgs") for s in idl["functions"][0]["sqlSignatures"]], + [{"maxdist": "-1.0", "maxt": "NULL"}, {"maxdist": "-1.0", "maxt": "NULL"}, + {"maxdist": "-1.0"}, None]) + + def test_a_null_default_feeds_the_parameter_its_number_reads(self): + idl = {"temporalTypes": {"tbool": {"linear": False}, "tint": {"linear": False}}, + "functions": [self._func()]} + idl, _ = attach_null_default_binds(idl, Path(self.tmp.name) / "src", + Path(self.tmp.name) / "sql") + sigs = idl["functions"][0]["sqlSignatures"] + self.assertEqual(sigs[2]["nullDefaultBinds"], + {"2": {"interp": "STEP"}, "3": {"maxt": "NULL"}}) + self.assertEqual(sigs[3]["nullDefaultBinds"], + {"2": {"interp": "STEP"}, "3": {"maxdist": "-1.0"}, + "4": {"maxt": "NULL"}}) + + if __name__ == "__main__": unittest.main()