Skip to content

Commit 002837c

Browse files
Gayathri Srividya RajavarapuGayathri Srividya Rajavarapu
authored andcommitted
fix: normalize dictionary types in Arrow scans
1 parent 5da8186 commit 002837c

2 files changed

Lines changed: 78 additions & 1 deletion

File tree

‎pyiceberg/io/pyarrow.py‎

Lines changed: 36 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1205,6 +1205,37 @@ def _pyarrow_schema_ensure_small_types(schema: pa.Schema) -> pa.Schema:
12051205
return visit_pyarrow(schema, _ConvertToSmallTypes())
12061206

12071207

1208+
def _pyarrow_type_ensure_non_dictionary_types(field_type: pa.DataType) -> pa.DataType:
1209+
if pa.types.is_dictionary(field_type):
1210+
return _pyarrow_type_ensure_non_dictionary_types(field_type.value_type)
1211+
elif pa.types.is_struct(field_type):
1212+
return pa.struct([field.with_type(_pyarrow_type_ensure_non_dictionary_types(field.type)) for field in field_type])
1213+
elif pa.types.is_list(field_type):
1214+
return pa.list_(field_type.value_field.with_type(_pyarrow_type_ensure_non_dictionary_types(field_type.value_type)))
1215+
elif pa.types.is_large_list(field_type):
1216+
return pa.large_list(field_type.value_field.with_type(_pyarrow_type_ensure_non_dictionary_types(field_type.value_type)))
1217+
elif pa.types.is_fixed_size_list(field_type):
1218+
return pa.list_(
1219+
field_type.value_field.with_type(_pyarrow_type_ensure_non_dictionary_types(field_type.value_type)),
1220+
field_type.list_size,
1221+
)
1222+
elif pa.types.is_map(field_type):
1223+
return pa.map_(
1224+
field_type.key_field.with_type(_pyarrow_type_ensure_non_dictionary_types(field_type.key_type)),
1225+
field_type.item_field.with_type(_pyarrow_type_ensure_non_dictionary_types(field_type.item_type)),
1226+
keys_sorted=field_type.keys_sorted,
1227+
)
1228+
return field_type
1229+
1230+
1231+
def _pyarrow_table_ensure_non_dictionary_types(table: pa.Table) -> pa.Table:
1232+
schema = pa.schema(
1233+
[field.with_type(_pyarrow_type_ensure_non_dictionary_types(field.type)) for field in table.schema],
1234+
metadata=table.schema.metadata,
1235+
)
1236+
return table.cast(schema) if schema != table.schema else table
1237+
1238+
12081239
@singledispatch
12091240
def visit_pyarrow(obj: pa.DataType | pa.Schema, visitor: PyArrowSchemaVisitor[T]) -> T:
12101241
"""Apply a pyarrow schema visitor to any point within a schema.
@@ -1795,7 +1826,11 @@ def to_table(self, tasks: Iterable[FileScanTask]) -> pa.Table:
17951826
# Note: cannot use pa.Table.from_batches(itertools.chain([first_batch], batches)))
17961827
# as different batches can use different schema's (due to large_ types)
17971828
result = pa.concat_tables(
1798-
(pa.Table.from_batches([batch]) for batch in itertools.chain([first_batch], batches)), promote_options="permissive"
1829+
(
1830+
_pyarrow_table_ensure_non_dictionary_types(pa.Table.from_batches([batch]))
1831+
for batch in itertools.chain([first_batch], batches)
1832+
),
1833+
promote_options="permissive",
17991834
)
18001835

18011836
return result

‎tests/io/test_pyarrow.py‎

Lines changed: 42 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -73,6 +73,7 @@
7373
_ConvertToArrowSchema,
7474
_determine_partitions,
7575
_primitive_to_physical,
76+
_pyarrow_table_ensure_non_dictionary_types,
7677
_read_deletes,
7778
_task_to_record_batches,
7879
_to_requested_schema,
@@ -1301,6 +1302,47 @@ def test_projection_concat_files(schema_int: Schema, file_int: str) -> None:
13011302
assert repr(result_table.schema) == "id: int32"
13021303

13031304

1305+
def test_arrow_scan_to_table_with_mixed_dictionary_and_plain_strings() -> None:
1306+
schema = Schema(NestedField(1, "foo", StringType(), required=False))
1307+
scan = ArrowScan(
1308+
table_metadata=TableMetadataV2(
1309+
location="file://a/b/",
1310+
last_column_id=1,
1311+
format_version=2,
1312+
schemas=[schema],
1313+
partition_specs=[PartitionSpec()],
1314+
),
1315+
io=PyArrowFileIO(),
1316+
projected_schema=schema,
1317+
row_filter=AlwaysTrue(),
1318+
)
1319+
values = pa.array(["a"], type=pa.string())
1320+
batches = iter([pa.record_batch([values], names=["foo"]), pa.record_batch([values.dictionary_encode()], names=["foo"])])
1321+
1322+
with patch.object(scan, "to_record_batches", return_value=batches):
1323+
assert scan.to_table([]).to_pydict() == {"foo": ["a", "a"]}
1324+
1325+
1326+
def test_pyarrow_table_ensure_non_dictionary_types_nested() -> None:
1327+
dictionary_values = pa.array(["a"]).dictionary_encode()
1328+
table = pa.table(
1329+
{
1330+
"struct": pa.StructArray.from_arrays([dictionary_values], names=["value"]),
1331+
"list": pa.ListArray.from_arrays(pa.array([0, 1]), dictionary_values),
1332+
}
1333+
)
1334+
1335+
normalized_table = _pyarrow_table_ensure_non_dictionary_types(table)
1336+
1337+
assert normalized_table.schema == pa.schema(
1338+
[
1339+
pa.field("struct", pa.struct([pa.field("value", pa.string())])),
1340+
pa.field("list", pa.list_(pa.string())),
1341+
]
1342+
)
1343+
assert normalized_table.to_pydict() == {"struct": [{"value": "a"}], "list": [["a"]]}
1344+
1345+
13041346
def test_identity_transform_column_projection(tmp_path: str, catalog: InMemoryCatalog) -> None:
13051347
# Test by adding a non-partitioned data file to a partitioned table, verifying partition value
13061348
# projection from manifest metadata.

0 commit comments

Comments
 (0)