diff --git a/parser/boundargs.py b/parser/boundargs.py index 2a3d537..60c9990 100644 --- a/parser/boundargs.py +++ b/parser/boundargs.py @@ -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: @@ -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") @@ -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] diff --git a/tests/test_boundargs.py b/tests/test_boundargs.py index f7ff88c..1433bc6 100644 --- a/tests/test_boundargs.py +++ b/tests/test_boundargs.py @@ -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): @@ -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" @@ -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()