diff --git a/parser/boundargs.py b/parser/boundargs.py index 16d9bcb..fce7410 100644 --- a/parser/boundargs.py +++ b/parser/boundargs.py @@ -617,6 +617,92 @@ def attach_type_derived_args(idl: dict, mdb_src: str | Path, return idl, n +# A local a wrapper takes from argument k unless it is NULL: +# `quadbin = PG_ARGISNULL(1) ? 0 : PG_GETARG_QUADBIN(1)`. +_ISNULL_CHOICE = r"(?]){var}\s*=(?!=)\s*PG_ARGISNULL\s*\(\s*(\d+)\s*\)\s*\?\s*([^:;]+?)\s*:" + + +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 + ``var = PG_ARGISNULL(k) ? literal : ...``.""" + g = _guarded_default(body, var) + if g is not None and g[0] == k: + return g[1] + c = _type_choice(body, var) + if c is not None and c[0] == k: + _, source, i, lin, other = c + sargs = sig.get("args") or [] + t = sig.get("ret") if source == "ret" else ( + sargs[i] if i is not None and i < len(sargs) else None) + if t in linear: + return lin if linear[t] else other + m = re.search(_ISNULL_CHOICE.format(var=re.escape(var)), body) + if m and int(m.group(1)) == k: + return _literal(m.group(2).strip()) + return None + + +def attach_null_default_binds(idl: dict, mdb_src: str | Path, sql_src: str | Path, + meos_src: str | Path | None = None) -> tuple[dict, int]: + """(idl, number of positions stated) once each SQL signature states in + ``nullDefaultBinds`` what its wrapper passes for an argument left to a NULL default. + + A SQL argument declared ``DEFAULT NULL`` reaches the wrapper as NULL when the call + leaves it out, and the wrapper chooses what the MEOS function receives: + ``tintSeqSetGaps(tint[], maxt interval DEFAULT NULL, maxdist float DEFAULT NULL)`` + passes ``maxt`` NULL and ``maxdist`` -1.0 (#_null_value). For each such position ``k`` + of a signature whose CREATE FUNCTION is not STRICT (a STRICT one answers NULL without + calling the wrapper), ``nullDefaultBinds[k]`` maps every C parameter argument ``k`` + feeds (#_caller_index) to that value, stated only when the wrapper gives every one of + them a value. Runs once the temporal types are attached; the wrappers are traced as + #merge_boundargs traces them.""" + from parser.sqlfn import _meos_to_mdb, _wrapper_sql_sigs, _wrapper_sql_strict + wrappers = extract_wrappers(mdb_src) + m2d = _meos_to_mdb(meos_src) if meos_src else {} + w2sig = _wrapper_sql_sigs(sql_src) + w2strict = _wrapper_sql_strict(sql_src) + linear = {t: rec.get("linear") for t, rec in (idl.get("temporalTypes") or {}).items()} + 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] + params = [p.get("name") for p in func.get("params", [])] + for sig in func.get("sqlSignatures") or []: + nulls = [k for k, d in enumerate(sig.get("argDefaults") or []) + if d is not None and d.strip().upper() == "NULL"] + if not nulls: + continue + w = _signature_wrapper(func, sig, ws, w2sig) or ws[0] + key = (sig.get("sqlName") or func.get("sqlfn"), tuple(sig.get("args") or ()), + sig.get("ret"), bool(sig.get("retSet"))) + body = wrappers.get(w) + if not body or (w2strict.get(w) or {}).get(key, True): + continue + args = _call_args(body, func["name"]) or _call_args(body, "pg_" + func["name"]) + if not args: + continue + fed = {} + for pos, a in enumerate(args): + if pos < len(params) and _IDENT.match(a): + k = _caller_index(body, a) + if k is not None: + fed.setdefault(k, []).append((params[pos], a)) + binds = {} + for k in nulls: + vals = {p: _null_value(body, a, k, sig, linear) for p, a in fed.get(k, ())} + if vals and all(v is not None for v in vals.values()): + binds[str(k)] = vals + if binds: + sig["nullDefaultBinds"] = binds + n += len(binds) + return idl, n + + # The SQL argument a wrapper reads: `PG_GETARG_(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*\)") diff --git a/parser/sqlfn.py b/parser/sqlfn.py index d24c780..3489f7e 100644 --- a/parser/sqlfn.py +++ b/parser/sqlfn.py @@ -207,7 +207,7 @@ def _strip_sql_comments(text): return "".join(out) -def _create_fn_stmts(text, bodies=False): +def _create_fn_stmts(text, bodies=False, strict=False): """Yield (sqlName, [raw arg decls], returnType|None, wrapper|None, retSet) for every CREATE FUNCTION in `text`, each parsed STATEMENT-BOUNDED (to its terminating `;`). returnType is the type of one returned row; retSet is True for `RETURNS SETOF`, @@ -218,7 +218,9 @@ def _create_fn_stmts(text, bodies=False): garbage return types. wrapper is None for a LANGUAGE SQL / $$ body (no C symbol). With `bodies`, each tuple ends with that body, read from the same `AS` clause #_AS_WRAPPER reads a C symbol from, its quotes undoubled; None for a function - with a C symbol.""" + with a C symbol. With `strict`, each tuple ends with whether the function is STRICT + (`RETURNS NULL ON NULL INPUT`), PostgreSQL's `proisstrict`: a strict function answers + NULL for a NULL argument without calling its wrapper.""" for m in _CREATE_FN.finditer(text): sqlname = m.group(1) i, depth, start = m.end(), 1, m.end() @@ -241,6 +243,11 @@ def _create_fn_stmts(text, bodies=False): # return type `boolean SUPPORT tspatial_supportfn`. Keep only the type. ret = _RET_ATTR.split(ret, maxsplit=1)[0].strip() or ret argdecls = [a for a in _split_top_commas(text[start:arg_close]) if a.strip()] + if strict: + yield (sqlname, argdecls, ret, wrapper, retset, + bool(re.search(r"\bSTRICT\b|\bRETURNS\s+NULL\s+ON\s+NULL\s+INPUT\b", + tail, re.I))) + continue if not bodies: yield sqlname, argdecls, ret, wrapper, retset continue @@ -275,6 +282,27 @@ def _wrapper_sql_sigs(sql_src): return out +def _wrapper_sql_strict(sql_src): + """MobilityDB-C wrapper name -> {(sqlName, args, ret, retSet): STRICT} for each of its + SQL signatures, the signatures #_wrapper_sql_sigs reads, keyed as #_signature_wrapper + matches a catalog signature to the CREATE FUNCTION stating it.""" + out = {} + _, vocab, composites = sql_statements(sql_src) + sql_src = Path(sql_src) + if not sql_src.exists(): + return out + for sf in sorted(sql_src.rglob("*.sql")): + text = _strip_sql_comments(sf.read_text(errors="ignore")) + for sqlname, argdecls, ret, wrapper, retset, strict in _create_fn_stmts(text, + strict=True): + if wrapper is None: + continue + s = sql_signature(sqlname, argdecls, ret, retset, vocab, composites) + out.setdefault(wrapper, {})[ + (s["sqlName"], tuple(s["args"]), s["ret"], s["retSet"])] = strict + return out + + def sql_statements(sql_src): """(statements, type vocabulary, composite types) of the CREATE FUNCTION statements under `sql_src`, each statement as #_create_fn_stmts yields it with its body. The diff --git a/run.py b/run.py index 5ea6f09..62657b2 100644 --- a/run.py +++ b/run.py @@ -15,7 +15,8 @@ from parser.nullable import merge_nullable, extract_param_docs from parser.nullresult import attach_null_result from parser.outparam import extract_param_names, merge_outparams -from parser.boundargs import (attach_call_literals, attach_type_derived_args, merge_boundargs, +from parser.boundargs import (attach_call_literals, attach_null_default_binds, + attach_type_derived_args, merge_boundargs, merge_sql_arg_params, resolve_bound_names, strip_call_literals) from parser.aggregates import attach_aggregates @@ -389,6 +390,11 @@ def main(): idl, ntd = attach_type_derived_args(idl, MDB_SRC, sql_src=SQL_SRC, meos_src=MEOS_SRC) print(f" arguments a wrapper derives from the signature's temporal type: {ntd}", file=sys.stderr) + # State what a wrapper passes for an argument a call leaves to its NULL default, + # `maxdist` -1.0 of `tintSeqSetGaps(tint[], interval)`, on the signatures whose + # CREATE FUNCTION is not STRICT. + idl, nnd = attach_null_default_binds(idl, MDB_SRC, SQL_SRC, meos_src=MEOS_SRC) + print(f" arguments left to a NULL default a wrapper fills: {nnd}", file=sys.stderr) # Keep a SQL signature two public functions claim on the one whose C parameters it # fits, matched by type once the object model and the type relations state the C type diff --git a/tests/test_boundargs.py b/tests/test_boundargs.py index 25da7f5..d58c708 100644 --- a/tests/test_boundargs.py +++ b/tests/test_boundargs.py @@ -12,7 +12,8 @@ import unittest from pathlib import Path -from parser.boundargs import (_sql_arg_params, attach_call_literals, attach_type_derived_args, +from parser.boundargs import (_sql_arg_params, attach_call_literals, attach_null_default_binds, + attach_type_derived_args, extract_call_literals, extract_wrappers, merge_boundargs, resolve_bound_names, strip_call_literals) @@ -1235,5 +1236,61 @@ def test_a_signature_carrying_the_interpolation_binds_none(self): self.assertEqual(n, 3) +class NullDefaultBindTests(unittest.TestCase): + """#attach_null_default_binds states what a wrapper passes for an argument left to its + NULL default, read from the wrapper of #TypeDerivedArgTests, and only for a CREATE + FUNCTION that is not STRICT.""" + + SQL = ''' +CREATE FUNCTION tintSeqSetGaps(tint[], maxt interval DEFAULT NULL, + maxdist float DEFAULT NULL) + RETURNS tint + AS 'MODULE_PATHNAME', 'Tsequenceset_constructor_gaps' + LANGUAGE C IMMUTABLE PARALLEL SAFE; +CREATE FUNCTION tfloatSeqSetGaps(tfloat[], maxt interval DEFAULT NULL, + maxdist float DEFAULT NULL) + RETURNS tfloat + AS 'MODULE_PATHNAME', 'Tsequenceset_constructor_gaps' + LANGUAGE C IMMUTABLE STRICT PARALLEL SAFE; +''' + + def setUp(self): + self.tmp = tempfile.TemporaryDirectory() + root = Path(self.tmp.name) + (root / "src").mkdir() + (root / "src" / "temporal.c").write_text(TYPE_DERIVED) + (root / "sql").mkdir() + (root / "sql" / "022_temporal.in.sql").write_text(self.SQL) + + def tearDown(self): + self.tmp.cleanup() + + def _idl(self): + def sig(name, t): + return {"sqlName": name, "args": [t + "[]", "interval", "float"], "ret": t, + "argDefaults": [None, "NULL", "NULL"]} + return { + "temporalTypes": {"tfloat": {"linear": True}, "tint": {"linear": False}}, + "functions": [ + {"name": "tsequenceset_make_gaps", "mdbC": "Tsequenceset_constructor_gaps", + "params": [{"name": n} for n in + ("instants", "count", "interp", "maxt", "maxdist")], + "sqlSignatures": [sig("tintSeqSetGaps", "tint"), + sig("tfloatSeqSetGaps", "tfloat")]}]} + + def test_each_null_default_takes_the_value_the_wrapper_passes(self): + idl, n = attach_null_default_binds(self._idl(), Path(self.tmp.name) / "src", + Path(self.tmp.name) / "sql") + self.assertEqual(idl["functions"][0]["sqlSignatures"][0]["nullDefaultBinds"], + {"1": {"maxt": "NULL"}, "2": {"maxdist": "-1.0"}}) + self.assertEqual(n, 2) + + def test_a_strict_function_binds_none(self): + # PostgreSQL answers NULL for the NULL argument without calling the wrapper + idl, _ = attach_null_default_binds(self._idl(), Path(self.tmp.name) / "src", + Path(self.tmp.name) / "sql") + self.assertNotIn("nullDefaultBinds", idl["functions"][0]["sqlSignatures"][1]) + + if __name__ == "__main__": unittest.main()