Skip to content

Commit 0faf93d

Browse files
committed
Support inferring partition from Hive-style path in add_files
1 parent 068aae5 commit 0faf93d

4 files changed

Lines changed: 256 additions & 2 deletions

File tree

‎pyiceberg/io/pyarrow.py‎

Lines changed: 5 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -2921,11 +2921,15 @@ def parquet_file_to_data_file(io: FileIO, table_metadata: TableMetadata, file_pa
29212921
stats_columns=compute_statistics_plan(schema, table_metadata.properties),
29222922
parquet_column_mapping=parquet_path_to_id_mapping(schema),
29232923
)
2924+
partition_spec = table_metadata.spec()
2925+
partition = partition_spec.partition_from_path(file_path, schema)
2926+
if partition is None:
2927+
partition = statistics.partition(partition_spec, table_metadata.schema())
29242928
data_file = DataFile.from_args(
29252929
content=DataFileContent.DATA,
29262930
file_path=file_path,
29272931
file_format=FileFormat.PARQUET,
2928-
partition=statistics.partition(table_metadata.spec(), table_metadata.schema()),
2932+
partition=partition,
29292933
file_size_in_bytes=len(input_file),
29302934
sort_order_id=None,
29312935
spec_id=table_metadata.default_spec_id,

‎pyiceberg/partitioning.py‎

Lines changed: 54 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -22,7 +22,7 @@
2222
from datetime import date, datetime, time
2323
from functools import cached_property, singledispatch
2424
from typing import Annotated, Any, Generic, TypeVar
25-
from urllib.parse import quote_plus
25+
from urllib.parse import quote_plus, unquote_plus
2626

2727
from pydantic import (
2828
BeforeValidator,
@@ -40,6 +40,7 @@
4040
HourTransform,
4141
IdentityTransform,
4242
MonthTransform,
43+
TimeTransform,
4344
Transform,
4445
TruncateTransform,
4546
UnknownTransform,
@@ -49,7 +50,9 @@
4950
)
5051
from pyiceberg.typedef import IcebergBaseModel, Record
5152
from pyiceberg.types import (
53+
BinaryType,
5254
DateType,
55+
FixedType,
5356
IcebergType,
5457
NestedField,
5558
PrimitiveType,
@@ -255,6 +258,56 @@ def partition_to_path(self, data: Record, schema: Schema) -> str:
255258
path = "/".join([field_str + "=" + value_str for field_str, value_str in zip(field_strs, value_strs, strict=True)])
256259
return path
257260

261+
def partition_from_path(self, location: str, schema: Schema) -> Record | None:
262+
"""Infer a partition Record from a Hive-style path (trailing key=value dirs).
263+
264+
Supports identity, bucket, and truncate (non-binary) transforms, since their
265+
to_human_string output is parseable back into the transform's result value.
266+
Returns None for unsupported transforms or a non-Hive-style path, so the
267+
caller can fall back to another inference strategy.
268+
"""
269+
from pyiceberg.conversions import partition_to_py
270+
271+
if self.is_unpartitioned():
272+
return None
273+
274+
for field in self.fields:
275+
if isinstance(field.transform, (TimeTransform, VoidTransform, UnknownTransform)):
276+
return None
277+
if isinstance(field.transform, TruncateTransform):
278+
source_field = schema.find_field(field.source_id)
279+
if isinstance(source_field.field_type, (FixedType, BinaryType)):
280+
# base64-encoded in the path, not decodable back to bytes
281+
return None
282+
283+
segments = [segment for segment in location.split("/") if segment]
284+
285+
partition_size = len(self.fields)
286+
partition_end_index_exclusive = len(segments) - 1 # exclude the file name
287+
partition_start_index = partition_end_index_exclusive - partition_size
288+
if partition_start_index < 0:
289+
return None
290+
291+
partition_segments = segments[partition_start_index:partition_end_index_exclusive]
292+
if not all("=" in segment for segment in partition_segments):
293+
return None
294+
295+
partition_type = self.partition_type(schema)
296+
field_types = partition_type.fields
297+
298+
values = []
299+
for partition_field, segment, field_type in zip(self.fields, partition_segments, field_types, strict=True):
300+
key, _, raw_value = segment.partition("=")
301+
key = unquote_plus(key)
302+
if key != partition_field.name:
303+
return None
304+
305+
value_str = unquote_plus(raw_value)
306+
value = partition_to_py(field_type.field_type, value_str) if value_str else None
307+
values.append(value)
308+
309+
return Record(*values)
310+
258311
def check_compatible(self, schema: Schema, allow_missing_fields: bool = False) -> None:
259312
# if the underlying field is dropped, we cannot check they are compatible -- continue
260313
schema_fields = schema._lazy_id_to_field

‎tests/integration/test_add_files.py‎

Lines changed: 113 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1009,6 +1009,119 @@ def test_add_files_hour_transform(session_catalog: Catalog) -> None:
10091009
tbl.add_files(file_paths=[file_path])
10101010

10111011

1012+
@pytest.mark.integration
1013+
def test_add_files_infers_partition_from_hive_style_path(
1014+
spark: SparkSession, session_catalog: Catalog, format_version: int
1015+
) -> None:
1016+
identifier = f"default.partitioned_table_hive_style_path_v{format_version}"
1017+
1018+
partition_spec = PartitionSpec(
1019+
PartitionField(source_id=4, field_id=1000, transform=IdentityTransform(), name="baz"),
1020+
spec_id=0,
1021+
)
1022+
1023+
tbl = _create_table(session_catalog, identifier, format_version, partition_spec)
1024+
1025+
# File's own "baz" values (900, 901) disagree with the path (baz=123): path must win.
1026+
file_paths = [
1027+
f"s3://warehouse/default/partitioned_table_hive_style_path/v{format_version}/baz=123/test-{i}.parquet" for i in range(2)
1028+
]
1029+
for i, file_path in enumerate(file_paths):
1030+
fo = tbl.io.new_output(file_path)
1031+
with fo.create(overwrite=True) as fos:
1032+
with pq.ParquetWriter(fos, schema=ARROW_SCHEMA) as writer:
1033+
writer.write_table(
1034+
pa.Table.from_pylist(
1035+
[
1036+
{
1037+
"foo": True,
1038+
"bar": "bar_string",
1039+
"baz": 900 + i,
1040+
"qux": date(2024, 3, 7),
1041+
}
1042+
],
1043+
schema=ARROW_SCHEMA,
1044+
)
1045+
)
1046+
1047+
tbl.add_files(file_paths=file_paths)
1048+
1049+
partition_rows = spark.sql(f"SELECT partition, record_count, file_count FROM {identifier}.partitions").collect()
1050+
assert [row.record_count for row in partition_rows] == [2]
1051+
assert [row.file_count for row in partition_rows] == [2]
1052+
assert [row.partition.baz for row in partition_rows] == [123]
1053+
1054+
assert len(tbl.scan().to_arrow()) == 2, "Expected 2 rows"
1055+
1056+
1057+
@pytest.mark.integration
1058+
def test_add_files_infers_bucket_partition_from_hive_style_path(
1059+
spark: SparkSession, session_catalog: Catalog, format_version: int
1060+
) -> None:
1061+
identifier = f"default.partitioned_table_bucket_hive_style_path_v{format_version}"
1062+
1063+
partition_spec = PartitionSpec(
1064+
PartitionField(source_id=4, field_id=1000, transform=BucketTransform(num_buckets=3), name="baz_bucket_3"),
1065+
spec_id=0,
1066+
)
1067+
1068+
tbl = _create_table(session_catalog, identifier, format_version, partition_spec)
1069+
1070+
# Without a Hive-style path this spec fails, see test_add_files_to_bucket_partitioned_table_fails.
1071+
file_path = f"s3://warehouse/default/partitioned_table_bucket_hive_style_path/v{format_version}/baz_bucket_3=1/test.parquet"
1072+
fo = tbl.io.new_output(file_path)
1073+
with fo.create(overwrite=True) as fos:
1074+
with pq.ParquetWriter(fos, schema=ARROW_SCHEMA) as writer:
1075+
writer.write_table(
1076+
pa.Table.from_pylist(
1077+
[
1078+
{"foo": True, "bar": "bar_string", "baz": 0, "qux": date(2024, 3, 7)},
1079+
{"foo": True, "bar": "bar_string", "baz": 1, "qux": date(2024, 3, 7)},
1080+
],
1081+
schema=ARROW_SCHEMA,
1082+
)
1083+
)
1084+
1085+
tbl.add_files(file_paths=[file_path])
1086+
1087+
partition_rows = spark.sql(f"SELECT partition, record_count, file_count FROM {identifier}.partitions").collect()
1088+
assert [row.record_count for row in partition_rows] == [2]
1089+
assert [row.partition.baz_bucket_3 for row in partition_rows] == [1]
1090+
1091+
assert len(tbl.scan().to_arrow()) == 2, "Expected 2 rows"
1092+
1093+
1094+
@pytest.mark.integration
1095+
def test_add_files_falls_back_to_stats_when_path_is_not_hive_style(
1096+
spark: SparkSession, session_catalog: Catalog, format_version: int
1097+
) -> None:
1098+
identifier = f"default.partitioned_table_non_hive_style_path_v{format_version}"
1099+
1100+
partition_spec = PartitionSpec(
1101+
PartitionField(source_id=4, field_id=1000, transform=IdentityTransform(), name="baz"),
1102+
spec_id=0,
1103+
)
1104+
1105+
tbl = _create_table(session_catalog, identifier, format_version, partition_spec)
1106+
1107+
# No "baz=" dir, so falls back to stats-based inference.
1108+
file_path = f"s3://warehouse/default/partitioned_table_non_hive_style_path/v{format_version}/test.parquet"
1109+
fo = tbl.io.new_output(file_path)
1110+
with fo.create(overwrite=True) as fos:
1111+
with pq.ParquetWriter(fos, schema=ARROW_SCHEMA) as writer:
1112+
writer.write_table(
1113+
pa.Table.from_pylist(
1114+
[{"foo": True, "bar": "bar_string", "baz": 123, "qux": date(2024, 3, 7)}],
1115+
schema=ARROW_SCHEMA,
1116+
)
1117+
)
1118+
1119+
tbl.add_files(file_paths=[file_path])
1120+
1121+
partition_rows = spark.sql(f"SELECT partition, record_count FROM {identifier}.partitions").collect()
1122+
assert [row.partition.baz for row in partition_rows] == [123]
1123+
1124+
10121125
@pytest.mark.integration
10131126
def test_add_files_to_branch(spark: SparkSession, session_catalog: Catalog, format_version: int) -> None:
10141127
identifier = f"default.test_add_files_branch_v{format_version}"

‎tests/table/test_partitioning.py‎

Lines changed: 84 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -31,6 +31,7 @@
3131
IdentityTransform,
3232
MonthTransform,
3333
TruncateTransform,
34+
VoidTransform,
3435
YearTransform,
3536
)
3637
from pyiceberg.typedef import Record
@@ -194,6 +195,89 @@ def test_partition_spec_to_path_dropped_source_id() -> None:
194195
assert spec.partition_to_path(record, schema) == "my%23str%25bucket=my%2Bstr/other+str%2Bbucket=%28+%29/my%21int%3Abucket=10"
195196

196197

198+
def test_partition_from_path_identity() -> None:
199+
schema = Schema(
200+
NestedField(field_id=1, name="foo", field_type=StringType(), required=False),
201+
NestedField(field_id=2, name="baz", field_type=IntegerType(), required=True),
202+
)
203+
spec = PartitionSpec(
204+
PartitionField(source_id=1, field_id=1000, transform=IdentityTransform(), name="foo"),
205+
PartitionField(source_id=2, field_id=1001, transform=IdentityTransform(), name="baz"),
206+
spec_id=0,
207+
)
208+
209+
assert spec.partition_from_path("s3://bucket/table/data/foo=hello/baz=123/00000-0.parquet", schema) == Record("hello", 123)
210+
211+
212+
def test_partition_from_path_unpartitioned() -> None:
213+
schema = Schema(NestedField(field_id=1, name="foo", field_type=StringType(), required=False))
214+
assert UNPARTITIONED_PARTITION_SPEC.partition_from_path("s3://bucket/table/data/00000-0.parquet", schema) is None
215+
216+
217+
def test_partition_from_path_not_hive_style() -> None:
218+
schema = Schema(NestedField(field_id=1, name="foo", field_type=StringType(), required=False))
219+
spec = PartitionSpec(PartitionField(source_id=1, field_id=1000, transform=IdentityTransform(), name="foo"), spec_id=0)
220+
221+
assert spec.partition_from_path("s3://bucket/table/data/00000-0.parquet", schema) is None
222+
223+
224+
def test_partition_from_path_field_name_mismatch() -> None:
225+
schema = Schema(NestedField(field_id=1, name="foo", field_type=StringType(), required=False))
226+
spec = PartitionSpec(PartitionField(source_id=1, field_id=1000, transform=IdentityTransform(), name="foo"), spec_id=0)
227+
228+
assert spec.partition_from_path("s3://bucket/table/data/wrong=hello/00000-0.parquet", schema) is None
229+
230+
231+
def test_partition_from_path_bucket_transform() -> None:
232+
schema = Schema(NestedField(field_id=1, name="int", field_type=IntegerType(), required=True))
233+
spec = PartitionSpec(
234+
PartitionField(source_id=1, field_id=1000, transform=BucketTransform(num_buckets=3), name="int_bucket"), spec_id=0
235+
)
236+
237+
assert spec.partition_from_path("s3://bucket/table/data/int_bucket=1/00000-0.parquet", schema) == Record(1)
238+
239+
240+
def test_partition_from_path_truncate_transform() -> None:
241+
schema = Schema(NestedField(field_id=1, name="str", field_type=StringType(), required=False))
242+
spec = PartitionSpec(
243+
PartitionField(source_id=1, field_id=1000, transform=TruncateTransform(width=3), name="str_trunc"), spec_id=0
244+
)
245+
246+
assert spec.partition_from_path("s3://bucket/table/data/str_trunc=abc/00000-0.parquet", schema) == Record("abc")
247+
248+
249+
def test_partition_from_path_truncate_binary_transform_unsupported() -> None:
250+
schema = Schema(NestedField(field_id=1, name="bin", field_type=BinaryType(), required=False))
251+
spec = PartitionSpec(
252+
PartitionField(source_id=1, field_id=1000, transform=TruncateTransform(width=3), name="bin_trunc"), spec_id=0
253+
)
254+
255+
# base64-encoded path value, not decodable back to bytes
256+
assert spec.partition_from_path("s3://bucket/table/data/bin_trunc=YWJj/00000-0.parquet", schema) is None
257+
258+
259+
def test_partition_from_path_time_transform_unsupported() -> None:
260+
schema = Schema(NestedField(field_id=1, name="date", field_type=DateType(), required=False))
261+
spec = PartitionSpec(PartitionField(source_id=1, field_id=1000, transform=MonthTransform(), name="date_month"), spec_id=0)
262+
263+
# calendar string in path, not the raw int
264+
assert spec.partition_from_path("s3://bucket/table/data/date_month=2024-03/00000-0.parquet", schema) is None
265+
266+
267+
def test_partition_from_path_void_transform_unsupported() -> None:
268+
schema = Schema(NestedField(field_id=1, name="foo", field_type=StringType(), required=False))
269+
spec = PartitionSpec(PartitionField(source_id=1, field_id=1000, transform=VoidTransform(), name="foo_null"), spec_id=0)
270+
271+
assert spec.partition_from_path("s3://bucket/table/data/foo_null=null/00000-0.parquet", schema) is None
272+
273+
274+
def test_partition_from_path_url_encoded_value() -> None:
275+
schema = Schema(NestedField(field_id=1, name="foo", field_type=StringType(), required=False))
276+
spec = PartitionSpec(PartitionField(source_id=1, field_id=1000, transform=IdentityTransform(), name="foo"), spec_id=0)
277+
278+
assert spec.partition_from_path("s3://bucket/table/data/foo=a%2Bb/00000-0.parquet", schema) == Record("a+b")
279+
280+
197281
def test_partition_type(table_schema_simple: Schema) -> None:
198282
spec = PartitionSpec(
199283
PartitionField(source_id=1, field_id=1000, transform=TruncateTransform(width=19), name="str_truncate"),

0 commit comments

Comments
 (0)