diff --git a/tools/codegen_jvm.py b/tools/codegen_jvm.py index c49beb9de..f0ed44d81 100644 --- a/tools/codegen_jvm.py +++ b/tools/codegen_jvm.py @@ -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 @@ -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): @@ -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]: @@ -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: @@ -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]: diff --git a/tools/codegen_spark_udfs.py b/tools/codegen_spark_udfs.py index 55a2e126f..bf96e1e02 100644 --- a/tools/codegen_spark_udfs.py +++ b/tools/codegen_spark_udfs.py @@ -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