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
53 changes: 47 additions & 6 deletions parser/boundargs.py
Original file line number Diff line number Diff line change
Expand Up @@ -561,17 +561,57 @@ def _caller_index(body: str, arg: str, depth: int = 0) -> int | None:
return via.pop() if len(via) == 1 else None


def _sql_arg_params(body: str, func: dict) -> list[str] | None:
def extract_meos_bodies(meos_src: str | Path) -> dict[str, str]:
"""``{function: body_text}`` for every documented MEOS definition under ``meos_src``,
read with the definition pattern #extract_param_lists reads."""
from parser.outparam import _FUNC
out: dict[str, str] = {}
for f in sorted(Path(meos_src).rglob("*.c")):
text = f.read_text(errors="ignore")
for m in _FUNC.finditer(text):
out.setdefault(m.group("name"), _COMMENT.sub(" ", _body(text, m.end() - 1)))
return out


def _kernel_args(body: str, func: dict, meos_body: str | None) -> list[str | None] | None:
"""The call arguments of the wrapper ``body`` standing for the parameters of ``func``,
one per parameter, where the two meet at a shared kernel rather than the wrapper calling
``func``, one hop down as #_delegated follows a wrapper into its shared helper:
``Distance_value_set`` calls ``distance_set_value(s, value)`` over the value it reads
first, and ``distance_set_int`` calls the same kernel over ``s`` and its integer ``i``
converted, so ``i`` stands where the wrapper passes ``value``. The kernel is the first
callee of ``meos_body`` the wrapper calls with as many arguments; a parameter named in no
argument of that call, or in two, stands for None."""
if not meos_body:
return None
names = [p["name"] for p in func.get("params", [])]
for m in _CALLEE.finditer(meos_body):
margs = _call_args(meos_body, m.group("name"))
wargs = _call_args(body, m.group("name"))
if not margs or not wargs or len(margs) != len(wargs):
continue
out: list[str | None] = []
for n in names:
pos = [j for j, a in enumerate(margs) if n in re.findall(r"\b[A-Za-z_]\w*\b", a)]
out.append(wargs[pos[0]] if len(pos) == 1 else None)
if any(a is not None for a in out):
return out
return None


def _sql_arg_params(body: str, func: dict, meos_body: str | None = None) -> 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"])
arguments carrying SQL arguments 0 to n-1, one each. They are read off the wrapper's call
of ``func``, else off the kernel the wrapper and ``meos_body``, the body of ``func``,
both call (#_kernel_args)."""
args = _call_args(body, func["name"]) or _kernel_args(body, func, meos_body)
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:
if a is None or a.strip().startswith("&") or _literal(a.strip()) is not None:
continue
k = _caller_index(body, a)
if k is not None:
Expand Down Expand Up @@ -606,6 +646,7 @@ def merge_sql_arg_params(idl: dict, mdb_src: str | Path,
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 {}
bodies = extract_meos_bodies(meos_src) if meos_src else {}
n = 0
for func in idl["functions"]:
primary = func.get("mdbC")
Expand All @@ -614,8 +655,8 @@ def merge_sql_arg_params(idl: dict, mdb_src: str | Path,
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]
orders = [_sql_arg_params(wrappers[w], func, bodies.get(func["name"]))
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]
Expand Down
37 changes: 37 additions & 0 deletions tests/test_boundargs.py
Original file line number Diff line number Diff line change
Expand Up @@ -969,6 +969,21 @@ def test_the_literal_is_attached_then_stripped(self):
Temporal *result = tdistance_tgeo_geo(temp, gs);
"""
TDISTANCE = {"name": "tdistance_tgeo_geo", "params": [{"name": "temp"}, {"name": "gs"}]}
# A commuted wrapper calling the kernel the public function calls, as Distance_value_set of
# mobilitydb/src/temporal/set_ops.c and distance_set_int of meos/src/temporal/set_ops_meos.c do;
# the public function passes its value converted.
KERNEL_WRAPPER = """
Datum value = PG_GETARG_DATUM(0);
Set *s = PG_GETARG_SET_P(1);
Datum result = distance_set_value(s, value);
PG_FREE_IF_COPY(s, 1);
PG_RETURN_DATUM(result);
"""
KERNEL_PUBLIC = """
VALIDATE_INTSET(s, INT_MAX);
return (int) distance_set_value(s, (long) i);
"""
DISTANCE_SET_INT = {"name": "distance_set_int", "params": [{"name": "s"}, {"name": "i"}]}


class SqlArgParamsTests(unittest.TestCase):
Expand All @@ -986,6 +1001,21 @@ def test_the_c_order_is_not_stated(self):
self.assertIsNone(_sql_arg_params(body, {"name": "tdistance_tgeo_geo",
"params": [{"name": "gs"}, {"name": "temp"}]}))

def test_a_shared_kernel_states_the_order(self):
"""#test_a_commuted_wrapper_reads_the_second_parameter_first, through a kernel."""
self.assertEqual(_sql_arg_params(KERNEL_WRAPPER, DISTANCE_SET_INT, KERNEL_PUBLIC),
["i", "s"])

def test_a_shared_kernel_in_the_c_order_is_not_stated(self):
"""#test_the_c_order_is_not_stated, through a kernel."""
body = KERNEL_WRAPPER.replace("PG_GETARG_DATUM(0)", "PG_GETARG_DATUM(1)").replace(
"PG_GETARG_SET_P(1)", "PG_GETARG_SET_P(0)")
self.assertIsNone(_sql_arg_params(body, DISTANCE_SET_INT, KERNEL_PUBLIC))

def test_without_the_public_body_nothing_is_stated(self):
"""#test_a_shared_kernel_states_the_order without the body of the public function."""
self.assertIsNone(_sql_arg_params(KERNEL_WRAPPER, DISTANCE_SET_INT))


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

Expand All @@ -1003,6 +1033,13 @@ def test_the_sequence_constructor_reads_the_interpolation_second(self):
self.assertEqual(self.fns["tsequence_make"]["shape"]["sqlArgParams"],
["instants", "interp", "lower_inc", "upper_inc"])

def test_the_number_first_reads_through_the_shared_kernel(self):
"""#test_the_sequence_constructor_reads_the_interpolation_second, for
nearestApproachDistance(float, tfloat), whose wrapper NAD_number_tnumber calls the
kernel nad_tfloat_float calls."""
sigs = self.fns["nad_tfloat_float"]["sqlSignatures"]
self.assertEqual([s.get("sqlArgParams") for s in sigs], [None, ["d", "temp"]])


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