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
45 changes: 35 additions & 10 deletions tools/codegen_jvm.py
Original file line number Diff line number Diff line change
Expand Up @@ -629,6 +629,13 @@ def __init__(self, cat, jmeos, pkg=SQL_PKG, engine='flink'):
self.jmeos = jmeos
self.fns = cat['functions']
self.by_name = {f['name']: f for f in self.fns}
# The catalog struct layouts, under #struct_layout of codegen_spark_udfs.py, the rule the
# Spark arm sizes them by: the rows of a set-returning signature and a contiguous array
# of structs (#_array_arg) read them, on both engines.
spark = _spark_module()
spark.STRUCTS.update({s['name']: s for s in cat.get('structs') or []})
self.layout = lambda name: (spark.struct_layout(name) # noqa: E731
if name in spark.STRUCTS else None)
self.enc = cat.get('typeEncodings', {})
self.enums = {e['name'] for e in cat.get('enums', [])}
# The value of each macro and enum member a wrapper binds by name.
Expand Down Expand Up @@ -1025,6 +1032,15 @@ def _array_arg(m, elem, a, p, jt, name, temps):
cls = f'{m.pkg}.types.{m.value_class[elem]}'
temps.append((t, f'{name}, {cls}::decode', 'values'))
return f'{cls}[]', t
# A contiguous array of structs (spanset_make reads Span *spans): each element decodes as
# above and its bytes are copied into the C array, each the size #struct_layout of
# codegen_spark_udfs.py gives the catalog struct.
lay = m.layout(_base(el['canonical'])) if c.count('*') == 1 else None
if elem in m.value_class and lay is not None \
and m.sql_cbase.get(elem) == _base(el['canonical']):
cls = f'{m.pkg}.types.{m.value_class[elem]}'
temps.append((t, f'{name}, {cls}::decode, {lay[0]}', 'structs'))
return f'{cls}[]', t
hit = SQL_ARRAY_SCALAR.get((elem, _base(el['c']))) \
or SQL_ARRAY_SCALAR.get((elem, _base(el['canonical'])))
if hit and c.count('*') == 1:
Expand Down Expand Up @@ -1175,6 +1191,7 @@ def _emit_eval(ov, defaults=None):
L.append(' try {')
for t, e, kind in temps:
v = {'value': f'_in.value({e})', 'values': f'_in.values({e})',
'structs': f'_in.structs({e})',
'buffer': f'_in.hold({e})'}[kind]
L.append(f' Pointer {t} = {v};')
L += [f' {s}' for s in body]
Expand Down Expand Up @@ -1860,6 +1877,22 @@ def _value_class_src(m, sql):
return b;
}

/** The values decoded one by one by dec into the contiguous C array of structs of
* the given size MEOS reads, as {@link #values} decodes them into an array of
* pointers: each decoded value is kept for release and its bytes are copied in. */
public <V extends MeosValue> Pointer structs(V[] vs,
java.util.function.Function<V, Pointer> dec, int size) {
Pointer b = hold(buffer(vs.length, size));
for (int i = 0; i < vs.length; i++) {
Pointer p = value(dec.apply(element(vs, i)));
if (p == null) {
throw new IllegalArgumentException("an array element does not decode");
}
p.transferTo(0, b, (long) i * size, size);
}
return b;
}

boolean holds(Pointer r) {
for (int i = 0; i < n; i++) {
if (owned[i] != null && owned[i].address() == r.address()) {
Expand Down Expand Up @@ -2074,11 +2107,7 @@ def run_flink_sql(args):
for sql in m.value_class:
(root / 'types' / f'{m.value_class[sql]}.java').write_text(_value_class_src(m, sql))

# The catalog struct layouts, under #struct_layout of codegen_spark_udfs.py, the rule the
# Spark arm sizes them by.
spark = _spark_module()
spark.STRUCTS.update({s['name']: s for s in cat.get('structs') or []})
layout = lambda name: spark.struct_layout(name) if name in spark.STRUCTS else None # noqa: E731
layout = m.layout

names = defaultdict(list) # SQL name -> eval methods
seen = defaultdict(set)
Expand Down Expand Up @@ -2778,11 +2807,7 @@ def run_spark_sql(args):
for sql in m.value_class:
(root / 'types' / f'{m.value_class[sql]}.java').write_text(_spark_value_class_src(m, sql))

# The catalog struct layouts, under #struct_layout of codegen_spark_udfs.py, as the flink-sql
# engine reads them for a set-returning signature's rows.
spark = _spark_module()
spark.STRUCTS.update({s['name']: s for s in cat.get('structs') or []})
layout = lambda name: spark.struct_layout(name) if name in spark.STRUCTS else None # noqa: E731
layout = m.layout

names = defaultdict(list) # SQL name -> (eval lines, Spark arg types, ret, classes)
seen = defaultdict(set)
Expand Down
8 changes: 8 additions & 0 deletions tools/codegen_spark_udfs.py
Original file line number Diff line number Diff line change
Expand Up @@ -281,6 +281,14 @@ def supported(f):
if r is None:
b = base(f["returnType"]["canonical"])
return ("internal" if b in INTERNAL or b == "__INTERNAL__" else "ret:"+norm(f["returnType"]["canonical"]))
# An array the catalog names in shape.inputArrays is no single value, though a contiguous
# array of structs (spanset_make reads Span *spans) has the C type of one: a UDF decoding
# one value would hand MEOS `count` elements to read past it. The typed SQL surfaces carry
# such an array (#_array_arg of codegen_jvm.py); here it is refused.
arrays = {a["param"] for a in (f.get("shape") or {}).get("inputArrays") or ()}
for p in in_params:
if p["name"] in arrays and (arg_kind(p["canonical"]) or ("",))[0] == "ptr":
return "array:" + norm(p["canonical"])
for p in in_params:
if arg_kind(p["canonical"]) is None:
b = base(p["canonical"])
Expand Down
6 changes: 4 additions & 2 deletions tools/spark-udf-gaps.txt
Original file line number Diff line number Diff line change
Expand Up @@ -52,7 +52,7 @@ cbuffer_as_ewkb array-or-out-param:size_out; unsupported-return:uint8_t *
cbuffer_hash ret:uint32_t
cbufferarr_round unsupported-return:struct Cbuffer **
cbufferarr_to_geom internal
cbufferset_make internal
cbufferset_make array:Cbuffer *
cbufferset_value_n internal
cmp_date_timestamp arg:Timestamp
cmp_timestamp_date arg:Timestamp
Expand Down Expand Up @@ -247,7 +247,7 @@ ne_timestamp_timestamptz arg:Timestamp
ne_timestamptz_timestamp arg:Timestamp
npoint_as_ewkb array-or-out-param:size_out; unsupported-return:uint8_t *
npoint_hash ret:uint32_t
npointset_make internal
npointset_make array:Npoint *
npointset_value_n internal
nsegment_as_ewkb array-or-out-param:size_out; unsupported-return:uint8_t *
overabove_tpcbox_tpcbox arg:TPCBox *
Expand Down Expand Up @@ -428,6 +428,7 @@ setstate_deserialize array-or-out-param:bytes
setstate_serialize array-or-out-param:size_out; unsupported-return:uint8_t *
span_hash ret:uint32_t
spanset_hash ret:uint32_t
spanset_make array:Span *
spanset_spanarr internal
spanset_spans array-or-out-param:count
spanset_split_each_n_spans array-or-out-param:count
Expand Down Expand Up @@ -465,6 +466,7 @@ stbox_tmax arg:TimestampTz *
stbox_tmin arg:TimestampTz *
stbox_to_box3d internal
stbox_to_gbox internal
stboxarr_round array:STBox *
taggstate_deserialize array-or-out-param:bytes; no-encoder:SkipList
taggstate_serialize no-decoder:SkipList; array-or-out-param:size_out; unsupported-return:uint8_t *
tbigint_time_boxes array-or-out-param:count
Expand Down
Loading