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: 64 additions & 22 deletions tools/codegen_jvm.py
Original file line number Diff line number Diff line change
Expand Up @@ -1306,6 +1306,58 @@ def _signature_overload(m, f, sig, args, ret, jsig, bound, layout):
return ov, why, False


def _default_variants(sargs, sdef, nparams, arity=None):
"""The defaults each shorter call of a signature passes, as #_emit_eval takes them: None
for the call stating every argument, then the Java literals of the trailing SQL defaults
a call leaves out, one more each time, while each is a plain constant (#_default_literal).
An overload a NULL default derives (#_null_default_signatures) answers the call of
``arity`` arguments alone, the prefix of the signature PostgreSQL accepts."""
defaults = (list(sdef) + [None] * len(sargs))[:len(sargs)]
variants = [None]
for k in range(1, len(sargs) + 1):
lits = [_default_literal(sargs[len(sargs) - k + j], defaults[len(sargs) - k + j])
for j in range(k)]
if any(x is None for x in lits):
break
variants.append(lits)
if arity is None:
return variants
return [dv for dv in variants if nparams - (len(dv) if dv else 0) == arity]


def _null_default_signatures(sig, sargs, sdef, sbound):
"""(arity, signature, args, defaults, boundArgs) for each shorter call of ``sig`` that
leaves an argument to a NULL default: that argument is no SQL argument of the call, and
the parameters it feeds take what the wrapper passes for it, the catalog's
``nullDefaultBinds``. A call of ``arity`` arguments leaves every later argument to its
default; a literal one keeps its place, so #_default_variants passes it as it passes
any other, and a NULL one the catalog states no value for ends the shorter calls."""
nd = sig.get('nullDefaultBinds') or {}
defaults = (list(sdef) + [None] * len(sargs))[:len(sargs)]
order = sig.get('sqlArgParams')
out = []
for arity in range(len(sargs) - 1, -1, -1):
tail = range(arity, len(sargs))
if any(defaults[j] is None for j in tail):
break
nulls = [j for j in tail if defaults[j].strip().upper() == 'NULL']
if any(str(j) not in nd for j in nulls):
break
if not nulls:
continue
keep = [j for j in range(len(sargs)) if j not in nulls]
bound = dict(sbound)
for j in nulls:
bound.update(nd[str(j)])
dsig = dict(sig, args=[(sig.get('args') or sargs)[j] for j in keep],
argDefaults=[defaults[j] for j in keep])
if order:
dsig['sqlArgParams'] = [order[j] for j in keep if j < len(order)]
out.append((arity, dsig, [sargs[j] for j in keep], [defaults[j] for j in keep],
bound))
return out


def _sql_overloads(m, cat, jmeos, layout, skipped, keep):
"""Every overload a typed SQL surface registers, as (SQL name, SQL argument types,
argument defaults, eval or None, reason, whether it answers an array of rows): one per
Expand All @@ -1322,11 +1374,17 @@ def _sql_overloads(m, cat, jmeos, layout, skipped, keep):
for (sqlname, sargs, sret, sdef, sbound), sig in zip(m.signatures(f),
f.get('sqlSignatures') or []):
ov, why, isset = _signature_overload(m, f, sig, sargs, sret, jsig, sbound, layout)
yield sqlname, sargs, sdef, ov, why, isset
yield sqlname, sargs, sdef, ov, why, isset, None
for arity, dsig, dargs, ddef, dbound in _null_default_signatures(sig, sargs, sdef,
sbound):
dov, _, disset = _signature_overload(m, f, dsig, dargs, sret, jsig, dbound,
layout)
if dov is not None:
yield sqlname, dargs, ddef, dov, None, disset, arity
for comp in cat.get('compositions') or []:
ov, why = _composition_overload(m, comp, jmeos, layout)
yield (comp['sqlName'], comp['args'], comp.get('argDefaults') or [], ov, why,
ov is not None and bool(comp.get('retSet')))
ov is not None and bool(comp.get('retSet')), None)


def _composition_inputs(m, comp, jmeos, f):
Expand Down Expand Up @@ -2118,21 +2176,13 @@ def run_flink_sql(args):
nset = 0
keep = lambda f: (f.get('sqlfn') and not f.get('sqlfnBackingOnly') # noqa: E731
and f.get('api') != 'internal')
for sqlname, sargs, sdef, ov, why, isset in _sql_overloads(m, cat, jmeos, layout,
for sqlname, sargs, sdef, ov, why, isset, arity in _sql_overloads(m, cat, jmeos, layout,
skipped, keep):
nset += isset
if ov is None:
skipped[why.split('/')[0]] += 1
continue
defaults = (list(sdef) + [None] * len(sargs))[:len(sargs)]
variants = [None]
for k in range(1, len(sargs) + 1):
lits = [_default_literal(sargs[len(sargs) - k + j], defaults[len(sargs) - k + j])
for j in range(k)]
if any(x is None for x in lits):
break
variants.append(lits)
for dv in variants:
for dv in _default_variants(sargs, sdef, len(ov[0]), arity):
shown = ov[0] if dv is None else ov[0][:len(ov[0]) - len(dv)]
key = tuple(t for t, _ in shown)
if key in seen[sqlname]:
Expand Down Expand Up @@ -2818,7 +2868,7 @@ def run_spark_sql(args):
nset = 0
keep = lambda f: (f.get('sqlfn') and not f.get('sqlfnBackingOnly') # noqa: E731
and f.get('api') == 'public')
for sqlname, sargs, sdef, ov, why, isset in _sql_overloads(m, cat, jmeos, layout,
for sqlname, sargs, sdef, ov, why, isset, arity in _sql_overloads(m, cat, jmeos, layout,
skipped, keep):
nset += isset
if ov is None:
Expand All @@ -2828,15 +2878,7 @@ def run_spark_sql(args):
if ret is None:
skipped['spark:ret'] += 1
continue
defaults = (list(sdef) + [None] * len(sargs))[:len(sargs)]
variants = [None]
for k in range(1, len(sargs) + 1):
lits = [_default_literal(sargs[len(sargs) - k + j], defaults[len(sargs) - k + j])
for j in range(k)]
if any(x is None for x in lits):
break
variants.append(lits)
for dv in variants:
for dv in _default_variants(sargs, sdef, len(ov[0]), arity):
shown = ov[0] if dv is None else ov[0][:len(ov[0]) - len(dv)]
key = tuple(t for t, _ in shown)
if key in seen[sqlname]:
Expand Down
7 changes: 6 additions & 1 deletion tools/codegen_spark_udfs.py
Original file line number Diff line number Diff line change
Expand Up @@ -393,7 +393,12 @@ def _omitted(f, p, vis, name):
i = free.index(p["name"])
dflt = (list(s.get("argDefaults") or []) + [None] * len(args))[:len(args)]
if vis <= i < len(args) and all(d is not None for d in dflt[vis:]):
vals.append(dflt[i])
# an argument left to a NULL default passes what the wrapper passes for it,
# the catalog's nullDefaultBinds, as the typed surfaces pass it
# (#_null_default_signatures of codegen_jvm.py)
nd = (s.get("nullDefaultBinds") or {}).get(str(i)) or {}
null = dflt[i].strip().upper() == "NULL"
vals.append(nd.get(p["name"], dflt[i]) if null else dflt[i])
return vals


Expand Down
Loading