Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
107 changes: 107 additions & 0 deletions parser/boundargs.py
Original file line number Diff line number Diff line change
Expand Up @@ -521,6 +521,113 @@ def own(w, sig=None):
return idl, n, list(dict.fromkeys(drift))


# The SQL argument a wrapper reads: `PG_GETARG_<T>(k)`, or a helper handed the call info
# with the index, as `input_interp_string(fcinfo, 1)` reads argument 1.
_GETARG = re.compile(r"PG_GETARG_\w+\s*\(\s*(\d+)\s*\)|\bfcinfo\s*,\s*(\d+)\s*\)")
# A leading C cast, `(TInstant **) temparr_extract(...)`.
_CAST = re.compile(r"^\(\s*[\w\s*]+\)\s*")


def _direct_indices(expr: str) -> set[int]:
return {int(a or b) for a, b in _GETARG.findall(expr)}


def _caller_index(body: str, arg: str, depth: int = 0) -> int | None:
"""The SQL argument the call argument ``arg`` of a wrapper ``body`` carries, or None.

An argument reading ``PG_GETARG_<T>(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."""
arg = _CAST.sub("", arg.strip())
found = _direct_indices(arg)
if found:
return found.pop() if len(found) == 1 else None
if not _IDENT.match(arg) or depth > 3:
return None
rhs = [m.group(1) for m in
re.finditer(r"(?<![\w.>])" + re.escape(arg) + r"\s*=(?!=)\s*([^;]+);", body)]
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
via = set()
for r in rhs:
for ident in set(re.findall(r"\b[a-z_]\w*\b", r)) - {arg}:
if re.search(r"(?<![\w.>])" + re.escape(ident) + r"\s*=(?!=)", body):
k = _caller_index(body, ident, depth + 1)
if k is not None:
via.add(k)
return via.pop() if len(via) == 1 else None


def _sql_arg_params(body: str, func: dict) -> list[str] | None:
"""The C parameters of ``func`` in the order of the SQL arguments the wrapper ``body``
reads for them, or None when that is their C order or cannot be read: the call
arguments carrying SQL arguments 0 to n-1, one each."""
args = _call_args(body, func["name"])
if not args:
return None
params = func.get("params", [])
by_k: dict[int, str] = {}
for a, p in zip(args, params):
if a.strip().startswith("&") or _literal(a.strip()) is not None:
continue
k = _caller_index(body, a)
if k is not None:
if k in by_k:
return None
by_k[k] = p["name"]
if not by_k or sorted(by_k) != list(range(len(by_k))):
return None
order = [by_k[k] for k in sorted(by_k)]
in_c = [p["name"] for p in params if p["name"] in order]
return order if order != in_c else None


def merge_sql_arg_params(idl: dict, mdb_src: str | Path,
sql_src: str | Path | None = None,
meos_src: str | Path | None = None) -> tuple[dict, int]:
"""(idl, count): the C parameters of a function in the order of the SQL arguments its
wrapper reads for them, where that is not their C order: ``tgeogpointSeq(tgeogpoint[],
text, boolean, boolean)`` reads the instants, the interpolation and the two inclusions,
while ``tsequence_make`` takes the instants, their count, the two inclusions and then the
interpolation. A binding pairing the SQL arguments with the C parameters by position
pairs them by these names instead.

Each SQL signature is traced to the wrapper whose CREATE FUNCTION states it, as
#merge_boundargs traces its literals: ``shape.sqlArgParams`` holds the order when every
signature reads the same one, else each signature that reads another carries its own
``sqlArgParams`` (``Concat_jsonb_jsonbset`` reads the jsonb first and
``Concat_jsonbset_jsonb`` the set, over one ``concat_jsonbset_jsonb``). Stated for the
exception only, as ``boundArgs`` is: where it is absent, the SQL arguments follow the C
parameters."""
from parser.sqlfn import _meos_to_mdb, _wrapper_sql_sigs
wrappers = extract_wrappers(mdb_src)
m2d = _meos_to_mdb(meos_src) if meos_src else {}
w2sig = _wrapper_sql_sigs(sql_src) if sql_src else {}
n = 0
for func in idl["functions"]:
primary = func.get("mdbC")
if not primary:
continue
ws = [primary] + [w for w in m2d.get(func["name"]) or () if w != primary]
sigs = func.get("sqlSignatures") or []
sig_ws = [_signature_wrapper(func, s, ws, w2sig) or primary for s in sigs] or [primary]
orders = [_sql_arg_params(wrappers[w], func) if w in wrappers else None
for w in sig_ws]
if len({tuple(o or ()) for o in orders}) == 1:
if orders[0]:
func.setdefault("shape", {})["sqlArgParams"] = orders[0]
n += 1
continue
for s, o in zip(sigs, orders):
if o:
s["sqlArgParams"] = o
n += 1
return idl, n


# `#define NAME <literal>`: an object-like macro whose body is one integer, float or
# boolean literal, the form every bound flag and default takes (`#define REST_AT true`,
# `#define OUT_DEFAULT_DECIMAL_DIGITS 15`). A function-like macro has `(` right after its
Expand Down
11 changes: 9 additions & 2 deletions run.py
Original file line number Diff line number Diff line change
Expand Up @@ -15,8 +15,8 @@
from parser.nullable import merge_nullable
from parser.nullresult import attach_null_result
from parser.outparam import extract_param_names, merge_outparams
from parser.boundargs import (attach_call_literals, merge_boundargs, resolve_bound_names,
strip_call_literals)
from parser.boundargs import (attach_call_literals, merge_boundargs, merge_sql_arg_params,
resolve_bound_names, strip_call_literals)
from parser.aggregates import attach_aggregates
from parser.altsqlfn import attach_alt_sql_names
from parser.compositions import attach_compositions
Expand Down Expand Up @@ -318,6 +318,13 @@ def main():
f"(neither caller arg, out-param, nor literal — inspect):", file=sys.stderr)
for fn, pn, reason in ba_drift:
print(f" {fn}({pn}) — {reason}", file=sys.stderr)
# A wrapper can read its SQL arguments in an order other than the C parameters'
# (`tgeogpointSeq` reads the interpolation before the bound inclusions, which
# `tsequence_make` takes first): `shape.sqlArgParams`, or a signature's own where the
# wrappers disagree, names the parameters in SQL argument order.
idl, nsap = merge_sql_arg_params(idl, MDB_SRC, sql_src=SQL_SRC, meos_src=MEOS_SRC)
print(f" SQL argument orders other than the C order (`sqlArgParams`): {nsap}",
file=sys.stderr)
# A bound literal can name a macro of a header the parse does not read (the
# installed headers carry none of temporal/temporal.h); its value comes from the
# source headers, so the catalog states what every bound literal is.
Expand Down
68 changes: 66 additions & 2 deletions tests/test_boundargs.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,8 +12,9 @@
import unittest
from pathlib import Path

from parser.boundargs import (attach_call_literals, extract_call_literals, extract_wrappers,
merge_boundargs, resolve_bound_names, strip_call_literals)
from parser.boundargs import (_sql_arg_params, attach_call_literals, extract_call_literals,
extract_wrappers, merge_boundargs, resolve_bound_names,
strip_call_literals)

# A synthetic MobilityDB wrapper source (mobilitydb/src/**/*.c shape).
SAMPLE = '''
Expand Down Expand Up @@ -940,5 +941,68 @@ def test_the_literal_is_attached_then_stripped(self):
for p in f["params"]))


# The sequence constructor reads the instants, the interpolation and the two inclusions, as
# Tsequence_constructor of mobilitydb/src/temporal/temporal.c does, and passes them in the
# C order of tsequence_make.
CONSTRUCTOR = """
ArrayType *array = PG_GETARG_ARRAYTYPE_P(0);
int count;
TInstant **instants = (TInstant **) temparr_extract(array, &count);
meosType temptype = instants[0]->temptype;
interpType interp = temptype_supports_linear(temptype) ? LINEAR : STEP;
if (PG_NARGS() > 1 && ! PG_ARGISNULL(1))
interp = input_interp_string(fcinfo, 1);
bool lower_inc = true, upper_inc = true;
if (PG_NARGS() > 2 && ! PG_ARGISNULL(2))
lower_inc = PG_GETARG_BOOL(2);
if (PG_NARGS() > 3 && ! PG_ARGISNULL(3))
upper_inc = PG_GETARG_BOOL(3);
TSequence *result = tsequence_make(instants, count, lower_inc, upper_inc,
interp, NORMALIZE);
"""
TSEQUENCE_MAKE = {"name": "tsequence_make", "params": [
{"name": n} for n in ("instants", "count", "lower_inc", "upper_inc", "interp", "normalize")]}
# A commuted wrapper passes its arguments in the other order.
COMMUTED = """
GSERIALIZED *gs = PG_GETARG_GSERIALIZED_P(0);
Temporal *temp = PG_GETARG_TEMPORAL_P(1);
Temporal *result = tdistance_tgeo_geo(temp, gs);
"""
TDISTANCE = {"name": "tdistance_tgeo_geo", "params": [{"name": "temp"}, {"name": "gs"}]}


class SqlArgParamsTests(unittest.TestCase):
"""#_sql_arg_params of parser/boundargs.py over synthetic wrapper bodies."""

def test_the_constructor_reads_its_arguments_in_another_order(self):
self.assertEqual(_sql_arg_params(CONSTRUCTOR, TSEQUENCE_MAKE),
["instants", "interp", "lower_inc", "upper_inc"])

def test_a_commuted_wrapper_reads_the_second_parameter_first(self):
self.assertEqual(_sql_arg_params(COMMUTED, TDISTANCE), ["gs", "temp"])

def test_the_c_order_is_not_stated(self):
body = COMMUTED.replace("tdistance_tgeo_geo(temp, gs)", "tdistance_tgeo_geo(gs, temp)")
self.assertIsNone(_sql_arg_params(body, {"name": "tdistance_tgeo_geo",
"params": [{"name": "gs"}, {"name": "temp"}]}))


IDL = Path(__file__).resolve().parent.parent / "output" / "meos-idl.json"


class SqlArgParamsContractTests(unittest.TestCase):
"""Over the generated catalog: the sequence constructor states the order its wrapper reads."""

def setUp(self):
if not IDL.exists():
self.skipTest(f"{IDL} not generated; run `python run.py` first")
import json
self.fns = {f["name"]: f for f in json.loads(IDL.read_text())["functions"]}

def test_the_sequence_constructor_reads_the_interpolation_second(self):
self.assertEqual(self.fns["tsequence_make"]["shape"]["sqlArgParams"],
["instants", "interp", "lower_inc", "upper_inc"])


if __name__ == "__main__":
unittest.main()
Loading