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
49 changes: 43 additions & 6 deletions src/substrait/type_inference.py
Original file line number Diff line number Diff line change
Expand Up @@ -645,12 +645,49 @@ def infer_expression_type(
if reference_type == "direct_reference":
segment = expression.selection.direct_reference

segment_reference_type = segment.WhichOneof("reference_type")

if segment_reference_type == "struct_field":
return schema.types[segment.struct_field.field]
else:
raise Exception(f"Unknown reference_type {reference_type}")
result = stt.Type(struct=schema)
while True:
kind = segment.WhichOneof("reference_type")
if kind == "struct_field":
if result.WhichOneof("kind") != "struct":
raise ValueError("Struct field child requires a struct")
child = segment.struct_field
if not 0 <= child.field < len(result.struct.types):
raise IndexError(
f"Struct field index {child.field} is out of range"
)
result = result.struct.types[child.field]
elif kind == "list_element":
if result.WhichOneof("kind") != "list":
raise ValueError("List element reference requires a list")
child = segment.list_element
result = result.list.type
elif kind == "map_key":
if result.WhichOneof("kind") != "map":
raise ValueError("Map key reference requires a map")
child = segment.map_key
key_type = stt.Type()
key_type.CopyFrom(infer_literal_type(child.map_key))
key_kind = key_type.WhichOneof("kind")
if key_kind == result.map.key.WhichOneof("kind"):
detail = getattr(key_type, key_kind)
detail.nullability = getattr(
result.map.key, key_kind
).nullability
if child.map_key.WhichOneof("literal_type") != "null":
detail.type_variation_reference = (
child.map_key.type_variation_reference
)
if key_type != result.map.key:
raise ValueError(
"Map key literal type does not match map key type"
)
result = result.map.value
Comment thread
coderabbitai[bot] marked this conversation as resolved.
else:
raise Exception(f"Unknown reference_type {kind}")
if not child.HasField("child"):
return result
segment = child.child
else:
raise Exception(f"Unknown reference_type {reference_type}")

Expand Down
21 changes: 21 additions & 0 deletions tests/dataframe/test_frame.py
Original file line number Diff line number Diff line change
Expand Up @@ -79,6 +79,27 @@ def test_filter_select_matches_builder():
assert fluent.SerializeToString() == raw.SerializeToString()


@pytest.mark.parametrize("kind", ["list", "map"])
def test_filter_collection_access_is_null(kind):
from substrait.builders.type import list as list_type
from substrait.builders.type import map as map_type
from substrait.type_inference import infer_plan_schema

value_type = string()
collection = (
list_type(value_type) if kind == "list" else map_type(string(), value_type)
)
ns = named_struct(
names=["values"], struct=struct(types=[collection], nullable=False)
)
df = sub.read_named_table("t", ns)
access = (
sub.col("values")[0] if kind == "list" else sub.col("values").map_key("key")
)
plan = df.filter(access.is_null()).select(access.alias("value")).to_plan()
assert infer_plan_schema(plan).struct.types[0] == value_type


def test_with_columns_named_appends_projection():
fluent = people_df().with_columns(bonus=sub.col("age") + 1).to_plan()
# ProjectRel appends: output has original columns + the new one.
Expand Down
211 changes: 211 additions & 0 deletions tests/test_type_inference.py
Original file line number Diff line number Diff line change
Expand Up @@ -983,6 +983,217 @@ def test_infer_expression_type_selection():
assert result == expected


@pytest.mark.parametrize("root", ["row", "outer_steps", "outer_anchor", "lambda"])
@pytest.mark.parametrize("depth", [1, 2])
@pytest.mark.parametrize("nullable", [False, True])
def test_infer_expression_type_nested_struct_selection(root, depth, nullable):
from substrait.type_inference import (
_outer_anchor_binding,
lambda_scope,
outer_schemas,
)

expected = stt.Type(string=stt.Type.String(nullability=_NULL if nullable else _REQ))
selected = expected
segment = stalg.Expression.ReferenceSegment(
struct_field=stalg.Expression.ReferenceSegment.StructField(field=1)
)
for level in range(depth):
selected = stt.Type(
struct=stt.Type.Struct(types=[struct.types[0], selected], nullability=_REQ)
)
segment = stalg.Expression.ReferenceSegment(
struct_field=stalg.Expression.ReferenceSegment.StructField(
field=0 if level == depth - 1 else 1, child=segment
)
)
row = stt.Type.Struct(types=[selected], nullability=_REQ)
ref = stalg.Expression.FieldReference(direct_reference=segment)
if root == "row":
ref.root_reference.SetInParent()
actual = infer_expression_type(stalg.Expression(selection=ref), row)
elif root == "lambda":
ref.lambda_parameter_reference.steps_out = 0
with lambda_scope(row):
actual = infer_expression_type(stalg.Expression(selection=ref), struct)
elif root == "outer_anchor":
ref.outer_reference.rel_reference = 5
with _outer_anchor_binding(5, row):
actual = infer_expression_type(stalg.Expression(selection=ref), struct)
else:
ref.outer_reference.steps_out = 1
token = outer_schemas.set((stt.NamedStruct(struct=row),))
try:
actual = infer_expression_type(stalg.Expression(selection=ref), struct)
finally:
outer_schemas.reset(token)
assert actual == expected


@pytest.mark.parametrize("index", [-1, 3])
def test_nested_struct_selection_rejects_out_of_range_field(index):
row = stt.Type.Struct(types=[stt.Type(struct=struct)], nullability=_REQ)
expr = _field_reference(0)
expr.selection.direct_reference.struct_field.child.struct_field.field = index
with pytest.raises(IndexError, match="Struct field index .* is out of range"):
infer_expression_type(expr, row)


def test_nested_struct_selection_rejects_child_on_scalar():
expr = _field_reference(0)
expr.selection.direct_reference.struct_field.child.struct_field.field = 0
with pytest.raises(ValueError, match="Struct field child requires a struct"):
infer_expression_type(expr, struct)


@pytest.mark.parametrize("kind", ["list_element", "map_key"])
@pytest.mark.parametrize("nullable", [False, True])
def test_nested_collection_selection(kind, nullable):
expr = _field_reference(0)
child = expr.selection.direct_reference.struct_field.child
expected = stt.Type(string=stt.Type.String(nullability=_NULL if nullable else _REQ))
if kind == "list_element":
child.list_element.offset = -1
selected = stt.Type(list=stt.Type.List(type=expected, nullability=_REQ))
else:
child.map_key.map_key.string = "key"
selected = stt.Type(
map=stt.Type.Map(
key=stt.Type(string=stt.Type.String(nullability=_REQ)),
value=expected,
nullability=_REQ,
)
)
row = stt.Type.Struct(types=[selected], nullability=_REQ)
assert infer_expression_type(expr, row) == expected


@pytest.mark.parametrize("literal_nullable", [False, True])
@pytest.mark.parametrize("compatible", [False, True])
def test_map_key_reference_validates_literal_type(literal_nullable, compatible):
expr = _field_reference(0)
key = expr.selection.direct_reference.struct_field.child.map_key.map_key
key.nullable = literal_nullable
if compatible:
key.i32 = 7
else:
key.string = "wrong"
expected = stt.Type(string=stt.Type.String(nullability=_NULL))
row = stt.Type.Struct(
types=[
stt.Type(
map=stt.Type.Map(
key=stt.Type(i32=stt.Type.I32(nullability=_REQ)),
value=expected,
nullability=_REQ,
)
)
],
nullability=_REQ,
)
before = row.SerializeToString(), expr.SerializeToString()
if compatible:
assert infer_expression_type(expr, row) == expected
else:
with pytest.raises(ValueError, match="Map key literal type"):
infer_expression_type(expr, row)
assert before == (row.SerializeToString(), expr.SerializeToString())


@pytest.mark.parametrize("scale", [2, 3])
def test_map_key_reference_validates_decimal_parameters(scale):
expr = _field_reference(0)
key = expr.selection.direct_reference.struct_field.child.map_key.map_key
key.decimal.precision = 10
key.decimal.scale = scale
key.decimal.value = (1000).to_bytes(16, "little", signed=True)
row = stt.Type.Struct(
types=[
stt.Type(
map=stt.Type.Map(
key=stt.Type(
decimal=stt.Type.Decimal(
precision=10, scale=2, nullability=_REQ
)
),
value=stt.Type(i32=stt.Type.I32(nullability=_REQ)),
nullability=_REQ,
)
)
],
nullability=_REQ,
)
if scale == 2:
assert infer_expression_type(expr, row) == row.types[0].map.value
else:
with pytest.raises(ValueError, match="Map key literal type"):
infer_expression_type(expr, row)


@pytest.mark.parametrize("variation", [0, 1])
def test_map_key_reference_validates_type_variation(variation):
expr = _field_reference(0)
literal = expr.selection.direct_reference.struct_field.child.map_key.map_key
literal.i32 = 7
literal.type_variation_reference = variation
row = stt.Type.Struct(
types=[
stt.Type(
map=stt.Type.Map(
key=stt.Type(
i32=stt.Type.I32(nullability=_REQ, type_variation_reference=1)
),
value=stt.Type(i32=stt.Type.I32(nullability=_REQ)),
nullability=_REQ,
)
)
],
nullability=_REQ,
)
if variation == 1:
assert infer_expression_type(expr, row) == row.types[0].map.value
else:
with pytest.raises(ValueError, match="Map key literal type"):
infer_expression_type(expr, row)


def test_nested_collection_selection_continues_through_struct():
expected = stt.Type(string=stt.Type.String(nullability=_NULL))
value = stt.Type(struct=stt.Type.Struct(types=[expected], nullability=_REQ))
mapping = stt.Type(
map=stt.Type.Map(
key=stt.Type(string=stt.Type.String(nullability=_REQ)),
value=value,
nullability=_REQ,
)
)
row = stt.Type.Struct(
types=[stt.Type(list=stt.Type.List(type=mapping, nullability=_REQ))],
nullability=_REQ,
)
expr = _field_reference(0)
element = expr.selection.direct_reference.struct_field.child.list_element
element.offset = 0
key = element.child.map_key
key.map_key.string = "key"
key.child.struct_field.field = 0
assert infer_expression_type(expr, row) == expected


@pytest.mark.parametrize(
"kind, message",
[
("list_element", "List element reference requires a list"),
("map_key", "Map key reference requires a map"),
],
)
def test_nested_collection_selection_rejects_wrong_container(kind, message):
expr = _field_reference(0)
getattr(expr.selection.direct_reference.struct_field.child, kind).SetInParent()
with pytest.raises(ValueError, match=message):
infer_expression_type(expr, struct)


def test_infer_expression_type_window_function():
"""Test infer_expression_type with a window function expression."""
expr = stalg.Expression(
Expand Down
Loading