diff --git a/parser/boundargs.py b/parser/boundargs.py index 95de933..be97a53 100644 --- a/parser/boundargs.py +++ b/parser/boundargs.py @@ -381,10 +381,10 @@ def _signature_wrapper(func: dict, sig: dict, claimed: list, w2sig: dict) -> str registers it, as ``attach_sqlfn_map`` keeps the first of two wrappers registering the same overload. None when no claimed wrapper states it.""" key = (sig.get("sqlName") or func.get("sqlfn"), tuple(sig.get("args") or ()), - sig.get("ret")) + sig.get("ret"), bool(sig.get("retSet"))) for w in claimed: for s in w2sig.get(w) or (): - if (s["sqlName"], tuple(s["args"]), s["ret"]) == key: + if (s["sqlName"], tuple(s["args"]), s["ret"], s["retSet"]) == key: return w return None diff --git a/parser/sqlfn.py b/parser/sqlfn.py index 988ec80..7bdb030 100644 --- a/parser/sqlfn.py +++ b/parser/sqlfn.py @@ -126,8 +126,10 @@ def _arg_type(decl, vocab): def _create_fn_stmts(text): - """Yield (sqlName, [raw arg decls], returnType|None, wrapper|None) for every + """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`, + PostgreSQL's `proretset`, a function returning any number of such rows. Bounding to the `;` is what stops a `LANGUAGE SQL` default-arg overload (whose own `AS 'SELECT ...'` has no C symbol) from bleeding its RETURNS/AS across the boundary into the next C-backed statement — the cross-statement mis-attribution that produced @@ -143,8 +145,9 @@ def _create_fn_stmts(text): tail = text[i:semi if semi != -1 else len(text)] # ') RETURNS AS ...' wm = _AS_WRAPPER.search(tail) wrapper = wm.group(1) if wm else None - rm = re.match(r"\s*RETURNS\s+(?:SETOF\s+)?(.+?)\s+AS\b", tail, re.I | re.S) - ret = " ".join(rm.group(1).split()) if rm else None + rm = re.match(r"\s*RETURNS\s+(SETOF\s+)?(.+?)\s+AS\b", tail, re.I | re.S) + ret = " ".join(rm.group(2).split()) if rm else None + retset = bool(rm and rm.group(1)) if ret: # PostgreSQL lets the function attributes come in any order, so an # attribute may sit between RETURNS and AS rather than after the body. @@ -153,12 +156,12 @@ def _create_fn_stmts(text): # 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()] - yield sqlname, argdecls, ret, wrapper + yield sqlname, argdecls, ret, wrapper, retset def _wrapper_sql_sigs(sql_src): """MobilityDB-C wrapper name -> list of per-overload SQL signatures - {sqlName, args:[type,...], required, ret}, straight from the CREATE FUNCTION + {sqlName, args:[type,...], required, ret, retSet}, straight from the CREATE FUNCTION statements. The .in.sql CREATE FUNCTION set IS the exact SQL registration surface, so a binding emits ONE registration per signature over the concrete arg types with NO type-scope heuristic — e.g. `minInstant` lands on exactly its four overloads @@ -173,15 +176,15 @@ def _wrapper_sql_sigs(sql_src): stmts, vocab = [], set() for sf in sorted(sql_src.rglob("*.sql")): text = sf.read_text(errors="ignore") - for sqlname, argdecls, ret, wrapper in _create_fn_stmts(text): - stmts.append((sqlname, argdecls, ret, wrapper)) + for sqlname, argdecls, ret, wrapper, retset in _create_fn_stmts(text): + stmts.append((sqlname, argdecls, ret, wrapper, retset)) if ret: vocab.add(ret) # a RETURNS clause is always a type for a in argdecls: bt = _bare_type(a) if bt and " " not in bt: vocab.add(bt) # a single-token arg is always a type - for sqlname, argdecls, ret, wrapper in stmts: + for sqlname, argdecls, ret, wrapper, retset in stmts: if wrapper is None: continue # LANGUAGE SQL / $$ body — no C symbol args = [_arg_type(a, vocab) for a in argdecls] @@ -189,7 +192,7 @@ def _wrapper_sql_sigs(sql_src): required = sum(1 for a in argdecls if not re.search(r"\bDEFAULT\b", a, re.I)) out.setdefault(wrapper, []).append( {"sqlName": sqlname, "args": args, "required": required, - "argDefaults": arg_defaults, "ret": ret}) + "argDefaults": arg_defaults, "ret": ret, "retSet": retset}) return out @@ -428,7 +431,7 @@ def attach_sqlfn_map(idl, meos_src, mdb_src, sql_src=None): wsigs = signatures_for(f["name"], wsigs, scope) scoped = True for s in wsigs: - key = (s["sqlName"], tuple(s["args"]), s["ret"]) + key = (s["sqlName"], tuple(s["args"]), s["ret"], s["retSet"]) if key not in seen: seen.add(key) sigs.append(s) @@ -486,6 +489,8 @@ def attach_sqlfn_map(idl, meos_src, mdb_src, sql_src=None): if not multiname and s["sqlName"] != surface: continue entry = {"args": s["args"], "ret": s["ret"]} + if s["retSet"]: + entry["retSet"] = True if any(d is not None for d in s["argDefaults"]): entry["argDefaults"] = s["argDefaults"] if multiname or s["sqlName"] != f["sqlfn"]: diff --git a/tests/test_sqlfn_setof.py b/tests/test_sqlfn_setof.py new file mode 100644 index 0000000..3ec8d05 --- /dev/null +++ b/tests/test_sqlfn_setof.py @@ -0,0 +1,116 @@ +"""A SQL signature states that it returns a set. + +`RETURNS SETOF integer` returns any number of rows of one integer each, and +`RETURNS integer` exactly one, so the row type alone does not state what a +binding registers: Flink and Spark carry a set-returning signature as one +function returning an array and unfold it into rows. The parser keeps +`SETOF` as the signature's `retSet`, PostgreSQL's `proretset`, beside `ret`, +the type of one row. A signature returning one value carries no `retSet`. + +Plain unittest, no pytest dependency; synthetic sources via a temp dir. +""" +import tempfile +import unittest +from pathlib import Path + +from parser.sqlfn import _create_fn_stmts, _wrapper_sql_sigs, attach_sqlfn_map + +MEOS_C = """ +/** + * @ingroup meos_setspan_accessor + * @brief Return the array of values of an integer set + * @csqlfn #Set_values(), #Set_unnest() + */ +int * +intset_values(const Set *s) +{ +} +""" + +MDB_C = """ +/** + * @brief Return the array of values of a set + * @sqlfn getValues() + */ +Datum +Set_values(PG_FUNCTION_ARGS) +{ +} + +/** + * @brief Return the values of a set as rows + * @sqlfn unnest() + */ +Datum +Set_unnest(PG_FUNCTION_ARGS) +{ +} +""" + +MDB_SQL = """ +CREATE FUNCTION getValues(intset) + RETURNS integer[] + AS 'MODULE_PATHNAME', 'Set_values' + LANGUAGE C IMMUTABLE STRICT PARALLEL SAFE; +CREATE FUNCTION unnest(intset) + RETURNS SETOF integer + AS 'MODULE_PATHNAME', 'Set_unnest' + LANGUAGE C IMMUTABLE STRICT PARALLEL SAFE; +CREATE FUNCTION aDisjointPairs(tgeometry[], tgeometry[], OUT i integer, OUT j integer) + RETURNS setof record + AS 'MODULE_PATHNAME', 'Adisjoint_tgeoarr_tgeoarr' + LANGUAGE C IMMUTABLE STRICT PARALLEL SAFE; +""" + + +def _attach(names): + idl = {"functions": [{"name": n, "api": "public"} for n in names]} + with tempfile.TemporaryDirectory() as d: + meos = Path(d) / "meos" / "src" + mdb = Path(d) / "mdb" + sql = Path(d) / "sql" + for p in (meos, mdb, sql): + p.mkdir(parents=True) + (meos / "temporal").mkdir() + (meos / "temporal" / "meos_catalog.c").write_text("") + (meos / "x.c").write_text(MEOS_C) + (mdb / "y.c").write_text(MDB_C) + (sql / "z.sql").write_text(MDB_SQL) + idl, _, _ = attach_sqlfn_map(idl, str(meos), str(mdb), str(sql)) + return {f["name"]: f for f in idl["functions"]} + + +class SetofStatementTests(unittest.TestCase): + + def test_setof_is_kept_apart_from_the_row_type(self): + stmts = {name: (ret, retset) + for name, _, ret, _, retset in _create_fn_stmts(MDB_SQL)} + self.assertEqual(stmts["unnest"], ("integer", True)) + self.assertEqual(stmts["getValues"], ("integer[]", False)) + + def test_setof_is_read_in_any_case(self): + stmts = {name: (ret, retset) + for name, _, ret, _, retset in _create_fn_stmts(MDB_SQL)} + self.assertEqual(stmts["aDisjointPairs"], ("record", True)) + + def test_every_wrapper_signature_states_its_retset(self): + with tempfile.TemporaryDirectory() as d: + (Path(d) / "x.sql").write_text(MDB_SQL) + sigs = _wrapper_sql_sigs(d) + self.assertTrue(sigs["Set_unnest"][0]["retSet"]) + self.assertFalse(sigs["Set_values"][0]["retSet"]) + + +class SetofCatalogTests(unittest.TestCase): + + def test_the_set_returning_signature_carries_retset(self): + """`intset_values` backs getValues and unnest: only unnest returns rows.""" + f = _attach(["intset_values"])["intset_values"] + by_name = {s.get("sqlName", f["sqlfn"]): s for s in f["sqlSignatures"]} + self.assertEqual(by_name["unnest"], {"args": ["intset"], "ret": "integer", + "retSet": True, "sqlName": "unnest"}) + self.assertNotIn("retSet", by_name["getValues"]) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_sqlfn_wrappers.py b/tests/test_sqlfn_wrappers.py index 1bd6c1e..1370f62 100644 --- a/tests/test_sqlfn_wrappers.py +++ b/tests/test_sqlfn_wrappers.py @@ -152,7 +152,7 @@ def test_an_attribute_between_returns_and_as_is_not_part_of_the_type(self): """PostgreSQL accepts the attributes in any order, so SUPPORT may precede the body — `aTouches(tcbuffer, cbuffer)` is the one place MobilityDB writes it that way.""" - rets = {name: ret for name, _, ret, _ in _create_fn_stmts(MDB_SQL)} + rets = {name: ret for name, _, ret, _, _ in _create_fn_stmts(MDB_SQL)} self.assertEqual(rets["aDwithin"], "boolean") self.assertEqual(rets["eDwithin"], "boolean")