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
35 changes: 30 additions & 5 deletions src/substrait/type_inference.py
Original file line number Diff line number Diff line change
Expand Up @@ -1163,16 +1163,41 @@ def infer_rel_schema(rel: stalg.Rel, *, registry=None, subtrees=()) -> stt.Type.
"expand switching field has no duplicate expressions; its "
"output type cannot be inferred"
)
# All duplicates of a switching field share one type; the first
# determines the output column type.
field_types.append(
# Duplicates share a type class but may differ in nullability.
Comment thread
nielspardon marked this conversation as resolved.
# The output is nullable if any duplicate is nullable.
duplicate_types = [
infer_expression_type(
duplicates[0],
duplicate,
parent_schema,
registry=registry,
subtrees=subtrees,
)
)
for duplicate in duplicates
]
bound = [
t for t in duplicate_types if t.WhichOneof("kind") != "unbound"
]
unbound = [
t for t in duplicate_types if t.WhichOneof("kind") == "unbound"
]
if len({t.WhichOneof("kind") for t in bound}) > 1:
raise ValueError(
"expand switching field duplicates must share one type class"
)
if any(
_field_nullability(t) == stt.Type.NULLABILITY_NULLABLE
for t in bound
):
# Nullable whatever an unbound placeholder later binds to.
output_type = _with_field_nullability(
bound[0], stt.Type.NULLABILITY_NULLABLE
)
elif unbound:
# Otherwise the nullability depends on the placeholder.
output_type = unbound[0]
else:
output_type = bound[0]
field_types.append(output_type)
# Expand appends an i32 column with the index of the duplicate the row
# is derived from.
field_types.append(
Expand Down
119 changes: 119 additions & 0 deletions tests/builders/plan/test_expand.py
Original file line number Diff line number Diff line change
@@ -1,3 +1,4 @@
import pytest
import substrait.algebra_pb2 as stalg
import substrait.plan_pb2 as stp
import substrait.type_pb2 as stt
Expand Down Expand Up @@ -86,6 +87,124 @@ def test_expand_schema_inference():
assert kinds == ["string", "fp64", "i32"]


@pytest.mark.parametrize(
"nullabilities",
[
pytest.param((False, True), id="nullable-last"),
pytest.param((True, False), id="nullable-first"),
pytest.param((False, False), id="all-required"),
pytest.param((True, True), id="all-nullable"),
pytest.param((False, False, True), id="nullable-third"),
pytest.param((False,), id="single-required"),
pytest.param((True,), id="single-nullable"),
],
)
def test_expand_switching_field_nullability(nullabilities):
# algebra.proto: a switching field is nullable if any duplicate is nullable.
names = [f"value_{i}" for i in range(len(nullabilities))]
input_schema = stt.NamedStruct(
names=["region", *names],
struct=stt.Type.Struct(
types=[string(nullable=False)] + [fp64(nullable=n) for n in nullabilities],
nullability=stt.Type.NULLABILITY_REQUIRED,
),
)
plan = expand(
read_named_table("sales", input_schema),
fields=[
("consistent", column("region")),
("switching", [column(name) for name in names]),
],
names=["region", "value", "idx"],
)(None)
original = stp.Plan()
original.CopyFrom(plan)

schema = infer_plan_schema(plan)

assert schema.struct.types[0] == string(nullable=False)
assert schema.struct.types[1] == fp64(nullable=any(nullabilities))
assert plan == original


def test_expand_switching_field_preserves_unbound_type():
# A partially bound plan may carry a placeholder with no nullability.
unbound = stt.Type(unbound=stt.Type.Unbound())
input_schema = stt.NamedStruct(
names=["u"],
struct=stt.Type.Struct(
types=[unbound], nullability=stt.Type.NULLABILITY_REQUIRED
),
)
plan = expand(
read_named_table("partial", input_schema),
fields=[("switching", [column("u"), column("u")])],
names=["value", "idx"],
)(None)

assert infer_plan_schema(plan).struct.types[0] == unbound


def _switching_plan(types):
names = [f"value_{i}" for i in range(len(types))]
schema = stt.NamedStruct(
names=names,
struct=stt.Type.Struct(types=types, nullability=stt.Type.NULLABILITY_REQUIRED),
)
return expand(
read_named_table("partial", schema),
fields=[("switching", [column(name) for name in names])],
names=["value", "idx"],
)(None)


@pytest.mark.parametrize(
"bound_nullabilities, unbound_position",
[
pytest.param((False,), 0, id="unbound-required"),
pytest.param((False,), 1, id="required-unbound"),
pytest.param((True,), 0, id="unbound-nullable"),
pytest.param((True,), 1, id="nullable-unbound"),
pytest.param((False, True), 0, id="unbound-required-nullable"),
pytest.param((False, True), 1, id="required-unbound-nullable"),
pytest.param((False, True), 2, id="required-nullable-unbound"),
],
)
def test_expand_switching_field_combines_unbound_duplicates(
bound_nullabilities, unbound_position
):
unbound = stt.Type(unbound=stt.Type.Unbound())
types = [fp64(nullable=n) for n in bound_nullabilities]
types.insert(unbound_position, unbound)
plan = _switching_plan(types)
original = plan.SerializeToString()

result = infer_plan_schema(plan).struct.types[0]

# A known nullable duplicate fixes nullability regardless of the placeholder.
# Otherwise the placeholder may still bind to a nullable type.
expected = fp64(nullable=True) if any(bound_nullabilities) else unbound
assert result == expected
assert plan.SerializeToString() == original


@pytest.mark.parametrize("reverse", [False, True])
@pytest.mark.parametrize("unbound_position", [None, 0, 1, 2])
def test_expand_switching_field_rejects_mixed_type_classes(reverse, unbound_position):
types = [fp64(nullable=False), string(nullable=True)]
if reverse:
types.reverse()
if unbound_position is not None:
types.insert(unbound_position, stt.Type(unbound=stt.Type.Unbound()))
plan = _switching_plan(types)
original = plan.SerializeToString()

with pytest.raises(ValueError, match="duplicates must share one type class"):
infer_plan_schema(plan)

assert plan.SerializeToString() == original


def test_expand_empty_switching_field_raises_clear_error():
# An empty switching field has no expression to derive a type from; schema
# inference must raise a clear error rather than an opaque IndexError.
Expand Down
8 changes: 8 additions & 0 deletions tests/dataframe/test_frame.py
Original file line number Diff line number Diff line change
Expand Up @@ -828,6 +828,14 @@ def test_unpivot_schema_inference_allows_chaining():
assert plan.relations[-1].root.input.HasField("filter")


@pytest.mark.parametrize("columns", [["a", "b"], ["b", "a"]])
def test_unpivot_rejects_mixed_type_classes(columns):
frame = sub.read_named_table("mixed", {"a": sub.i32, "b": sub.string})

with pytest.raises(ValueError, match="duplicates must share one type class"):
frame.unpivot(columns).to_plan()


def test_unpivot_requires_on():
with pytest.raises(ValueError, match="at least one column"):
_wide_df().unpivot([], index="region")
Expand Down
Loading