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
86 changes: 86 additions & 0 deletions parser/boundargs.py
Original file line number Diff line number Diff line change
Expand Up @@ -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"(?<![\w.>]){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_<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*\)")
Expand Down
32 changes: 30 additions & 2 deletions parser/sqlfn.py
Original file line number Diff line number Diff line change
Expand Up @@ -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`,
Expand All @@ -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()
Expand All @@ -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
Expand Down Expand Up @@ -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
Expand Down
8 changes: 7 additions & 1 deletion run.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down
59 changes: 58 additions & 1 deletion tests/test_boundargs.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down Expand Up @@ -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()
Loading