diff --git a/dev/provision.py b/dev/provision.py index 846576fb75..ab13346177 100644 --- a/dev/provision.py +++ b/dev/provision.py @@ -440,3 +440,39 @@ AS SELECT number, letter, extra FROM {catalog_name}.default.test_incremental_read """ ) + + +# Format-version 3 fixtures that need the Iceberg Java Table API (nanosecond timestamps, geometry, +# column defaults, unknown) are provisioned by dev/provision_v3.scala inside the Spark container. +def provision_v3_java_fixtures() -> None: + import os + import subprocess + + script = os.path.join(os.path.dirname(os.path.abspath(__file__)), "provision_v3.scala") + container = os.environ.get("PYICEBERG_SPARK_CONTAINER", "pyiceberg-spark") + subprocess.run(["docker", "cp", script, f"{container}:/tmp/provision_v3.scala"], check=True) + result = subprocess.run( + [ + "docker", + "exec", + "-e", + "AWS_REGION=us-east-1", + "-e", + "AWS_ACCESS_KEY_ID=admin", + "-e", + "AWS_SECRET_ACCESS_KEY=password", + container, + "bash", + "-c", + "cd /tmp && $SPARK_HOME/bin/spark-shell --master local[1] --conf spark.ui.enabled=false " + '--driver-java-options "-Duser.home=/tmp" < /tmp/provision_v3.scala', + ], + check=True, + capture_output=True, + text=True, + ) + if "PROVISION_V3_DONE" not in result.stdout: + raise RuntimeError(f"provision_v3.scala did not complete:\n{result.stdout[-4000:]}") + + +provision_v3_java_fixtures() diff --git a/dev/provision_v3.scala b/dev/provision_v3.scala new file mode 100644 index 0000000000..5af6e181c9 --- /dev/null +++ b/dev/provision_v3.scala @@ -0,0 +1,246 @@ +// Licensed to the Apache Software Foundation (ASF) under one +// or more contributor license agreements. See the NOTICE file +// distributed with this work for additional information +// regarding copyright ownership. The ASF licenses this file +// to you under the Apache License, Version 2.0 (the +// "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, +// software distributed under the License is distributed on an +// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +// KIND, either express or implied. See the License for the +// specific language governing permissions and limitations +// under the License. + +// Provisions format-version 3 fixtures that Spark SQL cannot create (nanosecond timestamps, +// geometry/geography, column defaults, unknown) through the Iceberg Java Table API. +// Executed by dev/provision.py inside the Spark container: +// spark-shell --master local[1] < /tmp/provision_v3.scala + +import java.math.BigDecimal +import java.nio.ByteBuffer +import java.time.{LocalDateTime, OffsetDateTime, ZoneOffset} + +import org.apache.iceberg.{DataFiles, FileFormat, PartitionKey, PartitionSpec, Schema, Table} +import org.apache.iceberg.catalog.TableIdentifier +import org.apache.iceberg.data.{GenericAppenderFactory, GenericRecord, InternalRecordWrapper, Record} +import org.apache.iceberg.expressions.Literal +import org.apache.iceberg.rest.RESTCatalog +import org.apache.iceberg.types.Types + +import scala.collection.JavaConverters._ + +val catalog = new RESTCatalog() +catalog.initialize( + "rest", + Map( + "uri" -> "http://rest:8181", + "io-impl" -> "org.apache.iceberg.aws.s3.S3FileIO", + "s3.endpoint" -> "http://object-store:9000", + "warehouse" -> "s3://warehouse/rest/" + ).asJava +) + +val v3 = Map("format-version" -> "3").asJava + +def recreate(name: String, schema: Schema, spec: PartitionSpec = PartitionSpec.unpartitioned()): Table = { + val ident = TableIdentifier.of("default", name) + if (catalog.tableExists(ident)) catalog.dropTable(ident, true) + catalog.createTable(ident, schema, spec, v3) +} + +def appendRecords(table: Table, records: Seq[Record], fileName: String): Unit = { + val spec = table.spec() + val grouped = records.groupBy { r => + val key = new PartitionKey(spec, table.schema()) + key.partition(new InternalRecordWrapper(table.schema().asStruct()).wrap(r)) + key + } + val append = table.newAppend() + grouped.zipWithIndex.foreach { case ((key, recs), i) => + val path = + if (spec.isPartitioned) table.locationProvider().newDataLocation(spec, key, s"$fileName-$i.parquet") + else table.locationProvider().newDataLocation(s"$fileName-$i.parquet") + val out = table.io().newOutputFile(path) + val appender = new GenericAppenderFactory(table.schema(), spec).newAppender(out, FileFormat.PARQUET) + try recs.foreach(appender.add) finally appender.close() + val builder = DataFiles.builder(spec).withInputFile(out.toInputFile()).withMetrics(appender.metrics()).withFormat(FileFormat.PARQUET) + if (spec.isPartitioned) builder.withPartition(key) + append.appendFile(builder.build()) + } + append.commit() +} + +def record(schema: Schema, values: (String, Any)*): Record = { + val rec = GenericRecord.create(schema) + values.foreach { case (k, v) => rec.setField(k, v) } + rec +} + +// --- nanosecond timestamps, unpartitioned (Spark cannot read or write timestamp_ns) --- +{ + val schema = new Schema( + Types.NestedField.required(1, "id", Types.IntegerType.get()), + Types.NestedField.optional(2, "ts_ns", Types.TimestampNanoType.withoutZone()), + Types.NestedField.optional(3, "tstz_ns", Types.TimestampNanoType.withZone()), + Types.NestedField.optional(4, "ts_us", Types.TimestampType.withoutZone()) + ) + val table = recreate("test_v3_ns_timestamps", schema) + appendRecords( + table, + Seq( + record( + schema, + "id" -> Integer.valueOf(1), + "ts_ns" -> LocalDateTime.of(2024, 1, 1, 0, 0, 0, 123456789), + "tstz_ns" -> OffsetDateTime.of(2024, 1, 1, 0, 0, 0, 123456789, ZoneOffset.UTC), + "ts_us" -> LocalDateTime.of(2024, 1, 1, 0, 0, 0, 123456000) + ), + record( + schema, + "id" -> Integer.valueOf(2), + "ts_ns" -> LocalDateTime.of(2024, 2, 2, 0, 0, 0, 1), + "tstz_ns" -> OffsetDateTime.of(2024, 2, 2, 0, 0, 0, 1, ZoneOffset.UTC), + "ts_us" -> LocalDateTime.of(2024, 2, 2, 0, 0, 0, 0) + ), + record(schema, "id" -> Integer.valueOf(3)) + ), + "ns" + ) + println("PROVISIONED test_v3_ns_timestamps") +} + +// --- nanosecond timestamps partitioned by year(ts), month(tstz), bucket(tstz, 4) --- +{ + val schema = new Schema( + Types.NestedField.required(1, "id", Types.IntegerType.get()), + Types.NestedField.optional(2, "ts", Types.TimestampNanoType.withoutZone()), + Types.NestedField.optional(3, "tstz", Types.TimestampNanoType.withZone()) + ) + val spec = PartitionSpec.builderFor(schema).year("ts").month("tstz").bucket("tstz", 4).build() + val table = recreate("test_v3_ns_partitions", schema, spec) + appendRecords( + table, + Seq( + record(schema, "id" -> Integer.valueOf(1), "ts" -> LocalDateTime.of(2023, 12, 31, 23, 59, 59, 999999999), "tstz" -> OffsetDateTime.of(2023, 12, 31, 23, 59, 59, 999999999, ZoneOffset.UTC)), + record(schema, "id" -> Integer.valueOf(2), "ts" -> LocalDateTime.of(2024, 1, 1, 0, 0, 0, 1), "tstz" -> OffsetDateTime.of(2024, 1, 1, 0, 0, 0, 1, ZoneOffset.UTC)), + record(schema, "id" -> Integer.valueOf(3), "ts" -> LocalDateTime.of(2024, 5, 5, 5, 5, 5, 0), "tstz" -> OffsetDateTime.of(2024, 5, 5, 5, 5, 5, 0, ZoneOffset.UTC)) + ), + "nsp" + ) + println("PROVISIONED test_v3_ns_partitions") +} + +// --- geometry / geography: schema only. Iceberg 1.11's generic Parquet writer does not support geo types. --- +{ + val schema = new Schema( + Types.NestedField.required(1, "id", Types.IntegerType.get()), + Types.NestedField.optional(2, "geom", Types.GeometryType.crs84()), + Types.NestedField.optional(3, "geom_srid", Types.GeometryType.of("srid:3857")), + Types.NestedField.optional(4, "geog", Types.GeographyType.crs84()), + Types.NestedField.optional(5, "geog_v", Types.GeographyType.of("srid:4326", org.apache.iceberg.types.EdgeAlgorithm.VINCENTY)) + ) + recreate("test_v3_geo", schema) + println("PROVISIONED test_v3_geo") +} + +// --- column defaults (Spark SQL refuses ADD COLUMN ... DEFAULT); rows inserted through Spark --- +{ + val schema = new Schema( + Types.NestedField.required(1, "id", Types.IntegerType.get()), + Types.NestedField.optional(2, "name", Types.StringType.get()) + ) + val table = recreate("test_v3_defaults", schema) + spark.sql("INSERT INTO rest.default.test_v3_defaults VALUES (1, 'one'), (2, 'two')") + table.refresh() + table + .updateSchema() + .addColumn("color", Types.StringType.get(), "doc", Literal.of("blue")) + .addColumn("qty", Types.IntegerType.get(), "doc", Literal.of(42)) + .addColumn("ratio", Types.DoubleType.get(), "doc", Literal.of(1.5d)) + .addColumn("d", Types.DateType.get(), "doc", Literal.of("2024-03-04").to(Types.DateType.get())) + .addColumn("ts", Types.TimestampType.withZone(), "doc", Literal.of("2024-03-04T05:06:07+00:00").to(Types.TimestampType.withZone())) + .addColumn("dec", Types.DecimalType.of(10, 2), "doc", Literal.of(new BigDecimal("12.34"))) + .addColumn("b", Types.BooleanType.get(), "doc", Literal.of(true)) + .addColumn("bin", Types.BinaryType.get(), "doc", Literal.of(ByteBuffer.wrap(Array[Byte](1, 2)))) + .addColumn("u", Types.UUIDType.get(), "doc", Literal.of("f79c3e09-677c-4bbd-a479-3f349cb785e7").to(Types.UUIDType.get())) + .addRequiredColumn("req", Types.LongType.get(), "doc", Literal.of(7L)) + .commit() + table.refresh() + // write-default changes, initial-default stays "blue" + table.updateSchema().updateColumnDefault("color", Literal.of("green")).commit() + println("PROVISIONED test_v3_defaults") +} + +// --- nested struct default --- +{ + val schema = new Schema( + Types.NestedField.required(1, "id", Types.IntegerType.get()), + Types.NestedField.optional( + 2, + "s", + Types.StructType.of(Types.NestedField.optional(3, "a", Types.IntegerType.get()), Types.NestedField.optional(4, "b", Types.StringType.get())) + ) + ) + val table = recreate("test_v3_nested_defaults", schema) + spark.sql("INSERT INTO rest.default.test_v3_nested_defaults VALUES (1, named_struct('a', 1, 'b', 'x'))") + table.refresh() + table.updateSchema().addColumn("s", "c", Types.IntegerType.get(), "doc", Literal.of(99)).commit() + println("PROVISIONED test_v3_nested_defaults") +} + +// --- unknown type column added after data exists --- +{ + val schema = new Schema(Types.NestedField.required(1, "id", Types.IntegerType.get())) + val table = recreate("test_v3_unknown", schema) + spark.sql("INSERT INTO rest.default.test_v3_unknown VALUES (1), (2)") + table.refresh() + table.updateSchema().addColumn("unk", Types.UnknownType.get()).commit() + println("PROVISIONED test_v3_unknown") +} + +// --- equality deletes (Spark SQL only writes position deletes and deletion vectors) --- +{ + val schema = new Schema( + Types.NestedField.required(1, "id", Types.IntegerType.get()), + Types.NestedField.optional(2, "name", Types.StringType.get()), + Types.NestedField.optional(3, "value", Types.DoubleType.get()) + ) + val table = recreate("test_v3_equality_deletes", schema) + appendRecords( + table, + Seq( + record(schema, "id" -> Integer.valueOf(1), "name" -> "a", "value" -> java.lang.Double.valueOf(1.0)), + record(schema, "id" -> Integer.valueOf(2), "name" -> "b", "value" -> java.lang.Double.valueOf(2.0)), + record(schema, "id" -> Integer.valueOf(3), "name" -> "c", "value" -> java.lang.Double.valueOf(3.0)), + record(schema, "id" -> Integer.valueOf(4), "value" -> java.lang.Double.valueOf(4.0)), + record(schema, "id" -> Integer.valueOf(5), "name" -> "e", "value" -> java.lang.Double.valueOf(5.0)) + ), + "eq-data" + ) + + def writeEqDeletes(fieldIds: Array[Int], deleteSchema: Schema, rows: Seq[Record], fileName: String): org.apache.iceberg.DeleteFile = { + val factory = new GenericAppenderFactory(table.schema(), table.spec(), fieldIds, deleteSchema, null) + val out = table.io().newOutputFile(table.locationProvider().newDataLocation(s"$fileName.parquet")) + val writer = factory.newEqDeleteWriter(org.apache.iceberg.encryption.EncryptedFiles.plainAsEncryptedOutput(out), FileFormat.PARQUET, null) + try rows.foreach(writer.write) finally writer.close() + writer.toDeleteFile() + } + + // delete rows by id, and by name where a null delete value matches a null data value + val idSchema = table.schema().select("id") + val nameSchema = table.schema().select("name") + val byId = writeEqDeletes(Array(1), idSchema, Seq(record(idSchema, "id" -> Integer.valueOf(2)), record(idSchema, "id" -> Integer.valueOf(5))), "eq-delete-id") + val byName = writeEqDeletes(Array(2), nameSchema, Seq(record(nameSchema)), "eq-delete-name") + table.newRowDelta().addDeletes(byId).addDeletes(byName).commit() + + // rows added after the deletes are not affected by them + appendRecords(table, Seq(record(schema, "id" -> Integer.valueOf(2), "name" -> "b2", "value" -> java.lang.Double.valueOf(20.0))), "eq-data-after") + println("PROVISIONED test_v3_equality_deletes") +} + +println("PROVISION_V3_DONE") +System.exit(0) diff --git a/mkdocs/docs/SUMMARY.md b/mkdocs/docs/SUMMARY.md index d268bcc4b0..9f48f3913a 100644 --- a/mkdocs/docs/SUMMARY.md +++ b/mkdocs/docs/SUMMARY.md @@ -26,6 +26,8 @@ - [API](api.md) - [Row Filter Syntax](row-filter-syntax.md) - [Expression DSL](expression-dsl.md) + - [Format version 3](format-version-3.md) + - [Geospatial types](geospatial.md) - [Contributing](contributing.md) - [Community](community.md) - Releases diff --git a/mkdocs/docs/api.md b/mkdocs/docs/api.md index 786223d47e..76671a76ec 100644 --- a/mkdocs/docs/api.md +++ b/mkdocs/docs/api.md @@ -407,6 +407,15 @@ long: [[4.896029,-122.431297,6.0989],[6.56667]] In the case of `tbl.delete(delete_filter="city == 'Groningen'")`, the whole Parquet file will be dropped without checking it contents, since from the Iceberg metadata PyIceberg can derive that all the content in the file matches the predicate. +By default, deletes are copy-on-write: a Parquet file that contains both matching and non-matching rows is rewritten without the matching rows. On a [format version 3](format-version-3.md#merge-on-read-deletes) table with the `write.delete.mode` table property set to `merge-on-read`, the matching rows are instead marked as deleted with a deletion vector, and the data files are left untouched. On v1 and v2 tables, `merge-on-read` falls back to copy-on-write with a warning. + +```python +with tbl.transaction() as transaction: + transaction.set_properties({"write.delete.mode": "merge-on-read"}) + +tbl.delete(delete_filter="city == 'Paris'") +``` + ### Partial overwrites When using the `overwrite` API, you can use an `overwrite_filter` to delete data that matches the filter before appending new data into the table. For example, consider the following Iceberg table: @@ -676,7 +685,7 @@ status: int8 not null snapshot_id: int64 not null sequence_number: int64 not null file_sequence_number: int64 not null -data_file: struct not null, record_count: int64 not null, file_size_in_bytes: int64 not null, column_sizes: map, value_counts: map, null_value_counts: map, nan_value_counts: map, lower_bounds: map, upper_bounds: map, key_metadata: binary, split_offsets: list, equality_ids: list, sort_order_id: int32> not null +data_file: struct not null, record_count: int64 not null, file_size_in_bytes: int64 not null, column_sizes: map, value_counts: map, null_value_counts: map, nan_value_counts: map, lower_bounds: map, upper_bounds: map, key_metadata: binary, split_offsets: list, equality_ids: list, sort_order_id: int32, first_row_id: int64, referenced_data_file: string, content_offset: int64, content_size_in_bytes: int64> not null child 0, content: int8 not null child 1, file_path: string not null child 2, file_format: string not null @@ -713,6 +722,10 @@ data_file: struct child 0, item: int32 child 15, sort_order_id: int32 + child 16, first_row_id: int64 + child 17, referenced_data_file: string + child 18, content_offset: int64 + child 19, content_size_in_bytes: int64 readable_metrics: struct not null, lat: struct not null, long: struct not null> child 0, city: struct not null child 0, column_size: int64 @@ -985,6 +998,10 @@ split_offsets: list equality_ids: list child 0, item: int32 sort_order_id: int32 +first_row_id: int64 +referenced_data_file: string +content_offset: int64 +content_size_in_bytes: int64 readable_metrics: struct not null, lat: struct not null, long: struct not null> child 0, city: struct not null child 0, column_size: int64 @@ -1814,6 +1831,46 @@ The low level API `plan_files` methods returns a set of tasks that provide the f In this case it is up to the engine itself to filter the file itself. Below, `to_arrow()` and `to_duckdb()` that already do this for you. +### Row lineage metadata columns + +On format version 3 tables, the reserved row lineage columns `_row_id` and `_last_updated_sequence_number` can be +requested in `selected_fields`. They are not part of the table schema, and are appended after the table columns in the +order they are requested: + +```python +scan = table.scan(selected_fields=("VendorID", "_row_id", "_last_updated_sequence_number")) +scan.to_arrow() +``` + +A row without a stored value inherits `_row_id` from its data file's `first_row_id` plus the row's position in the file, +and `_last_updated_sequence_number` from the data sequence number of the file. Both are null for rows in data files +added before row lineage was enabled. The columns are read-only: appending a dataframe that contains them is rejected by +the schema compatibility check. + +Rows keep their lineage when `delete()`, `overwrite()` or `upsert()` rewrites a data file (copy-on-write): the copied +rows' `_row_id` and `_last_updated_sequence_number` are stored in the rewritten file. Merge-on-read deletes leave the +data files in place, so row ids are unchanged as well. Rows that `upsert()` updates are written as new rows with new +row ids. Copied rows still count toward the snapshot's `added-rows`, so `next-row-id` advances past them. + +### Variant columns + +On format version 3 tables, `variant` columns are read as a struct of two binaries, `metadata` and `value`, holding +the Variant binary encoding. `pyiceberg.variant` decodes a value to JSON or to Python objects: + +```python +from pyiceberg.variant import to_json, to_python + +for row in table.scan(selected_fields=("id", "v")).to_arrow().to_pylist(): + v = row["v"] + if v is not None: + print(row["id"], to_json(v["metadata"], v["value"]), to_python(v["metadata"], v["value"])) +``` + +`to_python` returns `dict`, `list`, `str`, `int`, `float`, `Decimal`, `bytes`, `date`, `time`, `datetime` or `UUID` +values. Shredded variant files (with a `typed_value` field) cannot be read yet and raise `NotImplementedError` when the +column is projected. Writing variant columns is not supported until PyArrow can annotate the Parquet `VARIANT` +logical type. + ### Apache Arrow diff --git a/mkdocs/docs/configuration.md b/mkdocs/docs/configuration.md index 18f9db973b..c2f7309bb7 100644 --- a/mkdocs/docs/configuration.md +++ b/mkdocs/docs/configuration.md @@ -92,6 +92,7 @@ Iceberg tables support table properties to configure table behavior. | `write.py-location-provider.impl` | String of form `module.ClassName` | null | Optional, [custom `LocationProvider`](configuration.md#loading-a-custom-location-provider) implementation | | `write.data.path` | String pointing to location | `{metadata.location}/data` | Sets the location under which data is written. | | `write.metadata.path` | String pointing to location | `{metadata.location}/metadata` | Sets the location under which metadata is written. | +| `write.delete.mode` | `{copy-on-write,merge-on-read}` | copy-on-write | How `delete()` removes rows from partially matching data files. `merge-on-read` writes [deletion vectors](format-version-3.md#merge-on-read-deletes) on format version 3 tables, and falls back to copy-on-write on v1 and v2 tables. | ### Table behavior options @@ -960,4 +961,4 @@ Previous versions of Java (`<1.4.0`) implementations incorrectly assume the opti ## Nanoseconds Support -PyIceberg currently only supports upto microsecond precision in its TimestampType. PyArrow timestamp types in 's' and 'ms' will be upcast automatically to 'us' precision timestamps on write. Timestamps in 'ns' precision can also be downcast automatically on write if desired. This can be configured by setting the `downcast-ns-timestamp-to-us-on-write` property as "True" in the configuration file, or by setting the `PYICEBERG_DOWNCAST_NS_TIMESTAMP_TO_US_ON_WRITE` environment variable. Refer to the [nanoseconds timestamp proposal document](https://docs.google.com/document/d/1bE1DcEGNzZAMiVJSZ0X1wElKLNkT9kRkk0hDlfkXzvU/edit#heading=h.ibflcctc9i1d) for more details on the long term roadmap for nanoseconds support +Iceberg format version 3 adds the `timestamp_ns` and `timestamptz_ns` types, and PyIceberg reads them at full nanosecond precision. On format version 1 and 2 tables only microsecond precision is available: PyArrow timestamp types in 's' and 'ms' are upcast automatically to 'us' precision on write, and timestamps in 'ns' precision can be downcast automatically on write by setting the `downcast-ns-timestamp-to-us-on-write` property to "True" in the configuration file, or by setting the `PYICEBERG_DOWNCAST_NS_TIMESTAMP_TO_US_ON_WRITE` environment variable. On format version 3 tables nanosecond timestamps are kept as `timestamp_ns` unless this property is set. See [Format version 3](format-version-3.md) for the current state of v3 support. diff --git a/mkdocs/docs/format-version-3.md b/mkdocs/docs/format-version-3.md new file mode 100644 index 0000000000..6aa9ff0e2a --- /dev/null +++ b/mkdocs/docs/format-version-3.md @@ -0,0 +1,116 @@ + + +# Format version 3 + + +The [Iceberg table spec](https://iceberg.apache.org/spec/) version 3 adds new column types, default values, +row lineage, deletion vectors, multi-argument transforms, and table encryption. This page tracks what PyIceberg +supports today. The table is kept in sync with the code; a feature listed as unsupported raises a clear error rather +than silently producing incompatible files. + +The `files`, `data_files`, `delete_files`, `entries` and `all_*` inspect tables include the v3 content file fields +`first_row_id`, `referenced_data_file`, `content_offset` and `content_size_in_bytes`. + +## Support matrix + +| Feature | Read | Write | Notes | +|---|---|---|---| +| v3 table metadata (`next-row-id`, `encryption-keys`) | Yes | Yes | v3 tables can be created on every catalog. `next-row-id` starts at 0. | +| Upgrade a table from v1 or v2 to v3 | n/a | Yes | `upgrade_table_version(3)` sets `next-row-id` to 0 and leaves existing snapshots untouched. | +| Manifests and manifest lists | Yes | Yes | Data and delete manifests; manifest lists assign `first-row-id` to data manifests. | +| Row lineage (`_row_id`, `_last_updated_sequence_number`) | Yes | Yes | Data files inherit `first_row_id` from their manifest, and both columns can be requested in `selected_fields` (see [Row lineage metadata columns](api.md#row-lineage-metadata-columns)). Commits set the snapshot `first-row-id` and `added-rows` and advance `next-row-id`. Row ids and last updated sequence numbers are preserved across copy-on-write rewrites (`delete`, `overwrite`, `upsert`), which store them in the rewritten files, and across merge-on-read deletes; rows changed by `upsert` get new row ids. | +| Deletion vectors (Puffin) | Yes | Yes | Only the blob range named by `content_offset` / `content_size_in_bytes` is read, and the blob length, magic, CRC-32 and cardinality are validated. A DV applies only to its `referenced_data_file` and replaces older position deletes for that file. `delete()` writes DVs when `write.delete.mode` is `merge-on-read` (see [Merge-on-read deletes](#merge-on-read-deletes)); `overwrite()` and `upsert()` are copy-on-write. | +| Position delete files (v2 style) | Yes | No | v3 does not allow new position delete files; deletion vectors replace them. | +| Equality deletes | Yes | No | Equality delete files apply to data files with a lower data sequence number in the same partition, or in every partition when written with an unpartitioned spec. A null delete value matches a null data value. `count()` honours them. | +| `timestamp_ns` / `timestamptz_ns` | Yes | Yes | Values are read and written with nanosecond precision (Parquet INT64 nanoseconds, nanosecond-exact bounds). Row filters and partition pruning work on nanosecond columns. | +| `unknown` | Yes | Yes | Reads as null. Unknown columns must be optional with a null default and are not stored in data files. | +| `variant` | Yes (unshredded) | No | Read as a struct of `metadata` and `value` binaries; decode with `pyiceberg.variant` (see [Variant columns](api.md#variant-columns)). Shredded files raise `NotImplementedError` when the column is read. Writing waits on PyArrow support for the Parquet `VARIANT` logical type. Variant cannot be a partition or sort source. | +| `geometry` / `geography` | Yes | Yes | Type strings are parsed and serialized in the Java/spec form, e.g. `geometry(srid:3857)`. WKB can be appended as binary or GeoArrow columns; with `geoarrow-pyarrow` installed the Parquet `GEOMETRY`/`GEOGRAPHY` logical types are written. Geometry bounds are written as bounding-box points. Only equality and null predicates are supported, and bounds are not used to skip files. See [Geospatial types](geospatial.md). | +| Default values (`initial-default`, `write-default`) | Yes | Yes | `initial-default` is applied on read and `write-default` fills columns missing from appended data, so a required column with a `write-default` may be left out. Adding a column with a default value requires a v3 table. `unknown`, `geometry` and `geography` columns must default to null. | +| Type promotions | Yes | Yes | `unknown` to any primitive type and `date` to `timestamp` / `timestamp_ns` are only allowed on v3 tables; `date` promotion is rejected when the column is partitioned by a transform other than `year`, `month` or `day`. | +| Multi-argument transforms | Yes (as unknown transforms) | No | Fields with several `source-ids` load as unknown transforms and round-trip `source-ids`; their partition values are not used to prune files. | +| Table encryption | Metadata only | No | `encryption-keys` and `key-id` round-trip through metadata. | + +## Creating a v3 table + +Pass `format-version` as a table property. This works on every catalog: + +```python +from pyiceberg.table import TableProperties + +table = catalog.create_table( + identifier="default.events", + schema=schema, + properties={TableProperties.FORMAT_VERSION: "3"}, +) +``` + +An existing v1 or v2 table can be upgraded in place. Rows written before the upgrade have no row ids until the next +commit assigns them: + +```python +with table.transaction() as transaction: + transaction.upgrade_table_version(3) +``` + +Schemas that use v3-only types (`timestamp_ns`, `timestamptz_ns`, `unknown`, `variant`, `geometry`, `geography`) are rejected on +v1 and v2 tables, both on create and when a column is added. + +## Nanosecond timestamps + +On v3 tables, PyArrow `timestamp('ns')` columns map to `timestamp_ns` / `timestamptz_ns` and are read back at full +precision. On v1 and v2 tables nanosecond input must be downcast with the `downcast-ns-timestamp-to-us-on-write` +property (see [Configuration](configuration.md#nanoseconds-support)). + +Row filters on nanosecond columns accept ISO-8601 strings with up to nine fractional digits, integers (nanoseconds +since the epoch), `date` and `datetime` values. Python `datetime` objects only hold microseconds, so use a string or an +integer to filter on sub-microsecond values: + +```python +from pyiceberg.expressions import GreaterThanOrEqual + +table.scan(row_filter=GreaterThanOrEqual("ts_ns", "2024-01-01T00:00:00.123456789")) +table.scan(row_filter=GreaterThanOrEqual("tstz_ns", "2024-01-01T00:00:00.123456789+00:00")) +``` + +As with microsecond timestamps, `timestamptz_ns` filters require a zone offset and `timestamp_ns` filters reject one. + +## Merge-on-read deletes + +With the `write.delete.mode` table property set to `merge-on-read`, `delete()` on a v3 table marks the matching rows +of a data file as deleted in a deletion vector instead of rewriting the file: + +```python +table = catalog.create_table( + identifier="default.events", + schema=schema, + properties={"format-version": "3", "write.delete.mode": "merge-on-read"}, +) +table.delete("id == 2") +``` + +- Data files whose rows all match the filter are still dropped as a whole. +- All deletion vectors of a commit are written to one Puffin file, one `deletion-vector-v1` blob per data file. +- A data file has at most one deletion vector. Deleting more rows from a file that already has one writes a new vector + with the union of the deleted positions and removes the old one in the same commit (`removed-dvs` in the snapshot + summary). Older file-scoped position delete files are merged in and removed as well. +- The snapshot operation is `delete` and no rows are added, so row ids of the remaining rows are unchanged. +- The commit is validated against concurrent commits: it fails with a `ValidationException` if another writer + removed one of the data files or added deletes for it. diff --git a/mkdocs/docs/geospatial.md b/mkdocs/docs/geospatial.md index f7b8433b49..4f3e8eb228 100644 --- a/mkdocs/docs/geospatial.md +++ b/mkdocs/docs/geospatial.md @@ -1,6 +1,6 @@ # Geospatial Types -PyIceberg supports Iceberg v3 geospatial primitive types: `geometry` and `geography`. +PyIceberg models the Iceberg v3 geospatial primitive types `geometry` and `geography`. Support is a work in progress; the [format version 3](format-version-3.md) page lists what works today. ## Overview @@ -14,7 +14,7 @@ Both types store values as WKB (Well-Known Binary) bytes. ## Requirements - Iceberg format version 3 or higher -- Optional: `geoarrow-pyarrow` for GeoArrow extension type metadata and interoperability. Without it, geometry and geography are written as binary in Parquet while the Iceberg schema still preserves the spatial type. Install with `pip install pyiceberg[geoarrow]`. +- Optional: `geoarrow-pyarrow` for GeoArrow extension types and the Parquet `GEOMETRY`/`GEOGRAPHY` logical types. Without it, geometry and geography are written as binary in Parquet while the Iceberg schema still preserves the spatial type. Install with `pip install pyiceberg[geoarrow]`. ## Usage @@ -53,13 +53,14 @@ GeographyType() # Custom CRS GeographyType("EPSG:4326") -# Custom CRS and algorithm -GeographyType("EPSG:4326", "planar") +# Custom CRS and edge-interpolation algorithm: spherical (default), vincenty, thomas, andoyer or karney +GeographyType("EPSG:4326", "vincenty") ``` ### String Type Syntax -Types can also be specified as strings in schema definitions: +Types can also be specified as strings in schema definitions, using the same syntax as the Iceberg spec and the +Java implementation: ```python # Using string type names @@ -67,10 +68,15 @@ NestedField(1, "point", "geometry", required=True) NestedField(2, "region", "geography", required=True) # With parameters -NestedField(3, "location", "geometry('EPSG:4326')", required=True) -NestedField(4, "boundary", "geography('EPSG:4326', 'planar')", required=True) +NestedField(3, "location", "geometry(srid:3857)", required=True) +NestedField(4, "boundary", "geography(srid:4326, vincenty)", required=True) ``` +Type strings are case-insensitive. Types are written to table metadata in the same unquoted form, for example +`geometry(srid:3857)` and `geography(srid:4326, vincenty)`; a default CRS (`OGC:CRS84`) and algorithm (`spherical`) +are omitted. Metadata written by older PyIceberg versions, which quoted the parameters (`geometry('srid:3857')`), +is still parsed. + ## Data Representation Values are represented as WKB (Well-Known Binary) bytes at runtime: @@ -80,15 +86,59 @@ Values are represented as WKB (Well-Known Binary) bytes at runtime: point_wkb = bytes.fromhex("0101000000000000000000000000000000000000") ``` +## Writing + +WKB values can be appended from `binary` or `large_binary` columns, or from GeoArrow `geoarrow.wkb` columns whose CRS +and edge type match the table schema: + +```python +import pyarrow as pa + +table.append(pa.table({"id": pa.array([1], pa.int32()), "location": pa.array([point_wkb], pa.binary())})) +``` + +With `geoarrow-pyarrow` installed, data files are written with the Parquet `GEOMETRY` and `GEOGRAPHY` logical types, +which carry the CRS and, for geography, the edge-interpolation algorithm. The PyArrow Parquet writer only supports the +`spherical` algorithm, so geography columns with any other algorithm are written as plain binary with a warning; the +Iceberg schema still records the algorithm. Without `geoarrow-pyarrow`, all geo columns are written as plain binary. + +Geometry lower and upper bounds are written as the corner points of the column's bounding box, taken from the +Parquet geospatial statistics, including Z and M when every row group has them. The same bounds are read from files +registered with `add_files`. Geography bounds are only written when the Parquet writer provides a bounding box that +does not wrap the antimeridian, which PyArrow currently does not do. Geometry, geography and `unknown` columns must +default to null, so adding one with a `default_value` raises a `ValueError`. + +## Reading Tables Written by Other Engines + +Geometry and geography tables written by the Java implementation (for example through Spark or Sedona) can be +scanned. Parquet `GEOMETRY`/`GEOGRAPHY` columns, and plain binary columns holding WKB, are read according to the +Iceberg schema. With `geoarrow-pyarrow` installed the resulting Arrow columns use the GeoArrow `geoarrow.wkb` +extension type with the column's CRS (and edge type for geography); without it they are `large_binary`. + +```python +from pyiceberg.expressions import EqualTo + +table = catalog.load_table("db.spatial_table") +arrow_table = table.scan(row_filter=EqualTo("location", point_wkb)).to_arrow() +``` + +Row filters on geometry and geography columns support `EqualTo`, `NotEqualTo`, `In`, `NotIn`, `IsNull` and +`NotNull`, which compare WKB bytes. Ordering predicates such as `LessThan` raise a `ValueError`. In the +`readable_metrics` of `table.inspect.files()` and `table.inspect.entries()`, geometry and geography bounds are shown +as raw bytes: little-endian doubles `x:y`, `x:y:z` or `x:y:z:m`, which `pyiceberg.utils.geo.GeospatialBound.from_bytes` +decodes. + ## Current Limitations 1. **WKB/WKT Conversion**: Converting between WKB bytes and WKT strings requires external libraries (like Shapely). PyIceberg does not include this conversion to avoid heavy dependencies. 2. **Spatial Predicates**: Spatial filtering (e.g., ST_Contains, ST_Intersects) is not yet supported for query pushdown. -3. **Bounds Metrics**: Geometry/geography columns do not currently contribute to data file bounds metrics. +3. **Bounds Metrics**: Geometry/geography bounds (bounding boxes) are not used to skip data files. + +4. **Without geoarrow-pyarrow**: When the `geoarrow-pyarrow` package is not installed, geometry and geography columns are stored as binary without the Parquet geospatial logical types, and no geometry bounds are written. The Iceberg schema preserves type information, but other tools reading the Parquet files directly may not recognize them as spatial types. Install with `pip install pyiceberg[geoarrow]` for full GeoArrow support. -4. **Without geoarrow-pyarrow**: When the `geoarrow-pyarrow` package is not installed, geometry and geography columns are stored as binary without GeoArrow extension type metadata. The Iceberg schema preserves type information, but other tools reading the Parquet files directly may not recognize them as spatial types. Install with `pip install pyiceberg[geoarrow]` for full GeoArrow support. +5. **Geography algorithms**: Only `spherical` geography columns are written with the Parquet `GEOGRAPHY` logical type (see [Writing](#writing)). ## Format Version diff --git a/pyiceberg/avro/resolver.py b/pyiceberg/avro/resolver.py index 3f19ae9878..c87af1a097 100644 --- a/pyiceberg/avro/resolver.py +++ b/pyiceberg/avro/resolver.py @@ -105,6 +105,7 @@ TimeType, UnknownType, UUIDType, + VariantType, ) STRUCT_ROOT = -1 @@ -214,6 +215,9 @@ def visit_geography(self, geography_type: "GeographyType") -> Writer: """Geography is written as WKB bytes in Avro.""" return BinaryWriter() + def visit_variant(self, variant_type: VariantType) -> Writer: + raise NotImplementedError("Variant is not supported in Avro files") + CONSTRUCT_WRITER_VISITOR = ConstructWriter() @@ -377,6 +381,9 @@ def visit_geography(self, geography_type: "GeographyType", partner: IcebergType """Geography is written as WKB bytes in Avro.""" return BinaryWriter() + def visit_variant(self, variant_type: VariantType, partner: IcebergType | None) -> Writer: + raise NotImplementedError("Variant is not supported in Avro files") + class ReadSchemaResolver(PrimitiveWithPartnerVisitor[IcebergType, Reader]): __slots__ = ("read_types", "read_enums", "context") @@ -539,6 +546,9 @@ def visit_geography(self, geography_type: "GeographyType", partner: IcebergType """Geography is read as WKB bytes from Avro.""" return BinaryReader() + def visit_variant(self, variant_type: VariantType, partner: IcebergType | None) -> Reader: + raise NotImplementedError("Variant is not supported in Avro files") + class SchemaPartnerAccessor(PartnerAccessor[IcebergType]): def schema_partner(self, partner: IcebergType | None) -> IcebergType | None: diff --git a/pyiceberg/conversions.py b/pyiceberg/conversions.py index 6dd964f436..03f096f41a 100644 --- a/pyiceberg/conversions.py +++ b/pyiceberg/conversions.py @@ -62,6 +62,7 @@ TimeType, UnknownType, UUIDType, + VariantType, strtobool, ) from pyiceberg.utils.datetime import ( @@ -76,11 +77,15 @@ time_str_to_micros, time_to_micros, timestamp_to_micros, + timestamp_to_nanos, timestamptz_to_micros, + timestamptz_to_nanos, to_human_day, to_human_time, to_human_timestamp, + to_human_timestamp_ns, to_human_timestamptz, + to_human_timestamptz_ns, ) from pyiceberg.utils.decimal import decimal_to_bytes, unscaled_to_decimal @@ -327,6 +332,12 @@ def _(_: UnknownType, value: Any) -> None: return None +@to_bytes.register(VariantType) +def _(_: VariantType, value: Any) -> bytes: + """Raise, since the spec does not define a single-value binary serialization for variant.""" + raise ValueError("Variant values have no single-value binary serialization") + + @singledispatch # type: ignore def from_bytes(primitive_type: PrimitiveType, b: bytes) -> L: # type: ignore """Convert bytes to a built-in python value. @@ -404,6 +415,12 @@ def _(type_: UnknownType, buf: bytes) -> None: return None +@from_bytes.register(VariantType) +def _(type_: VariantType, buf: bytes) -> None: + """Raise, since the spec does not define a single-value binary serialization for variant.""" + raise ValueError("Variant values have no single-value binary serialization") + + @singledispatch # type: ignore def to_json(primitive_type: PrimitiveType, val: Any) -> L: # type: ignore """Convert built-in python values into JSON value types. @@ -463,6 +480,22 @@ def _(_: TimestamptzType, val: int | datetime) -> str: return to_human_timestamptz(val) +@to_json.register(TimestampNanoType) +def _(_: TimestampNanoType, val: int | datetime) -> str: + """Python datetime (without timezone) or nanoseconds since epoch serializes into an ISO8601 timestamp with nanoseconds.""" + if isinstance(val, datetime): + val = datetime_to_nanos(val) + return to_human_timestamp_ns(val) + + +@to_json.register(TimestamptzNanoType) +def _(_: TimestamptzNanoType, val: int | datetime) -> str: + """Python datetime (with timezone) or nanoseconds since epoch serializes into an ISO8601 timestamp with nanoseconds.""" + if isinstance(val, datetime): + val = datetime_to_nanos(val) + return to_human_timestamptz_ns(val) + + @to_json.register(FloatType) @to_json.register(DoubleType) def _(_: FloatType | DoubleType, val: float) -> float: @@ -509,6 +542,21 @@ def _(_: UUIDType, val: uuid.UUID) -> str: return str(val) +@to_json.register(UnknownType) +def _(_: UnknownType, val: None) -> None: + """Columns of type unknown always default to null, so only null has a JSON representation.""" + if val is not None: + raise ValueError(f"Columns of type unknown must default to null, got: {val}") + + +@to_json.register(VariantType) +def _(_: VariantType, val: Any) -> None: + """Variant columns must default to null, so the only JSON single value is null.""" + if val is not None: + raise ValueError(f"Variant columns must default to null, got: {val!r}") + return None + + @to_json.register(GeometryType) def _(_: GeometryType, val: bytes) -> str: """Serialize geometry to WKT string per Iceberg spec. @@ -613,6 +661,26 @@ def _(_: TimestamptzType, val: str | int | datetime) -> datetime: return val +@from_json.register(TimestampNanoType) +def _(_: TimestampNanoType, val: str | int | datetime) -> int: + """JSON ISO8601 string into nanoseconds since epoch, since Python datetimes cannot hold nanoseconds.""" + if isinstance(val, str): + return timestamp_to_nanos(val) + if isinstance(val, datetime): + return datetime_to_nanos(val) + return val + + +@from_json.register(TimestamptzNanoType) +def _(_: TimestamptzNanoType, val: str | int | datetime) -> int: + """JSON ISO8601 string into nanoseconds since epoch, since Python datetimes cannot hold nanoseconds.""" + if isinstance(val, str): + return timestamptz_to_nanos(val) + if isinstance(val, datetime): + return datetime_to_nanos(val) + return val + + @from_json.register(FloatType) @from_json.register(DoubleType) def _(_: FloatType | DoubleType, val: float) -> float: @@ -664,10 +732,27 @@ def _(_: UUIDType, val: str | bytes | uuid.UUID) -> uuid.UUID: return val +@from_json.register(UnknownType) +def _(_: UnknownType, val: None) -> None: + """Columns of type unknown always default to null, so only null has a JSON representation.""" + if val is not None: + raise ValueError(f"Columns of type unknown must default to null, got: {val}") + + +@from_json.register(VariantType) +def _(_: VariantType, val: Any) -> None: + """Variant columns must default to null, so the only JSON single value is null.""" + if val is not None: + raise ValueError(f"Variant columns must default to null, got: {val!r}") + return None + + @from_json.register(GeometryType) -def _(_: GeometryType, val: str | bytes) -> bytes: +def _(_: GeometryType, val: str | bytes | None) -> bytes | None: """Convert JSON WKT string into WKB bytes per Iceberg spec. + Geometry columns must default to null, so a null value is returned as-is. + Note: This requires WKT to WKB conversion which is not yet implemented. The Iceberg spec requires geometry values to be represented as WKT strings in JSON, but PyIceberg stores geometry as WKB bytes at runtime. @@ -675,7 +760,7 @@ def _(_: GeometryType, val: str | bytes) -> bytes: Raises: NotImplementedError: WKT to WKB conversion is not yet supported. """ - if isinstance(val, bytes): + if val is None or isinstance(val, bytes): # Already WKB bytes, return as-is return val raise NotImplementedError( @@ -685,9 +770,11 @@ def _(_: GeometryType, val: str | bytes) -> bytes: @from_json.register(GeographyType) -def _(_: GeographyType, val: str | bytes) -> bytes: +def _(_: GeographyType, val: str | bytes | None) -> bytes | None: """Convert JSON WKT string into WKB bytes per Iceberg spec. + Geography columns must default to null, so a null value is returned as-is. + Note: This requires WKT to WKB conversion which is not yet implemented. The Iceberg spec requires geography values to be represented as WKT strings in JSON, but PyIceberg stores geography as WKB bytes at runtime. @@ -695,7 +782,7 @@ def _(_: GeographyType, val: str | bytes) -> bytes: Raises: NotImplementedError: WKT to WKB conversion is not yet supported. """ - if isinstance(val, bytes): + if val is None or isinstance(val, bytes): # Already WKB bytes, return as-is return val raise NotImplementedError( diff --git a/pyiceberg/expressions/__init__.py b/pyiceberg/expressions/__init__.py index ece0db82db..26064dc51e 100644 --- a/pyiceberg/expressions/__init__.py +++ b/pyiceberg/expressions/__init__.py @@ -30,7 +30,7 @@ from pyiceberg.expressions.literals import AboveMax, BelowMin, Literal, literal from pyiceberg.schema import Accessor, Schema from pyiceberg.typedef import IcebergBaseModel, IcebergRootModel, L, LiteralValue, StructProtocol -from pyiceberg.types import DoubleType, FloatType, NestedField +from pyiceberg.types import DoubleType, FloatType, GeographyType, GeometryType, NestedField from pyiceberg.utils.singleton import Singleton @@ -908,7 +908,13 @@ def literal(self) -> LiteralValue: def bind(self, schema: Schema, case_sensitive: bool = True) -> BoundLiteralPredicate: bound_term = self.term.bind(schema, case_sensitive) - lit = self.literal.to(bound_term.ref().field.field_type) + field_type = bound_term.ref().field.field_type + if isinstance(field_type, (GeometryType, GeographyType)) and not isinstance(self, (EqualTo, NotEqualTo)): + raise ValueError( + f"{self.__class__.__name__} is not supported for {field_type} column {bound_term.ref().field.name}: " + "geometry and geography values have no ordering, only (not) equal to, (not) in and (not) null are supported" + ) + lit = self.literal.to(field_type) if isinstance(lit, AboveMax): if isinstance(self, (LessThan, LessThanOrEqual, NotEqualTo)): diff --git a/pyiceberg/expressions/literals.py b/pyiceberg/expressions/literals.py index 61581d9b3c..b50d22b3b5 100644 --- a/pyiceberg/expressions/literals.py +++ b/pyiceberg/expressions/literals.py @@ -41,11 +41,15 @@ DoubleType, FixedType, FloatType, + GeographyType, + GeometryType, IcebergType, IntegerType, LongType, StringType, + TimestampNanoType, TimestampType, + TimestamptzNanoType, TimestamptzType, TimeType, UUIDType, @@ -57,10 +61,15 @@ days_to_date, micros_to_days, micros_to_timestamp, + nanos_to_days, time_str_to_micros, time_to_micros, timestamp_to_micros, + timestamp_to_nanos, timestamptz_to_micros, + timestamptz_to_nanos, + to_human_timestamp_ns, + to_human_timestamptz_ns, ) from pyiceberg.utils.decimal import decimal_to_unscaled, unscaled_to_decimal from pyiceberg.utils.singleton import Singleton @@ -353,6 +362,14 @@ def _(self, _: TimestampType) -> Literal[int]: def _(self, _: TimestamptzType) -> Literal[int]: return TimestampLiteral(self.value) + @to.register(TimestampNanoType) + def _(self, _: TimestampNanoType) -> Literal[int]: + return _nanos_literal(self.value, TimestampNanoLiteral) + + @to.register(TimestamptzNanoType) + def _(self, _: TimestamptzNanoType) -> Literal[int]: + return _nanos_literal(self.value, TimestamptzNanoLiteral) + @to.register(DecimalType) def _(self, type_var: DecimalType) -> Literal[Decimal]: unscaled = Decimal(self.value) @@ -457,6 +474,14 @@ def to(self, type_var: IcebergType) -> Literal: # type: ignore def _(self, _: DateType) -> Literal[int]: return self + @to.register(TimestampNanoType) + def _(self, _: TimestampNanoType) -> Literal[int]: + return _nanos_literal(self.value * NANOS_PER_DAY, TimestampNanoLiteral) + + @to.register(TimestamptzNanoType) + def _(self, _: TimestamptzNanoType) -> Literal[int]: + return _nanos_literal(self.value * NANOS_PER_DAY, TimestamptzNanoLiteral) + class TimeLiteral(Literal[int]): def __init__(self, value: int) -> None: @@ -501,6 +526,79 @@ def _(self, _: TimestamptzType) -> Literal[int]: def _(self, _: DateType) -> Literal[int]: return DateLiteral(micros_to_days(self.value)) + @to.register(TimestampNanoType) + def _(self, _: TimestampNanoType) -> Literal[int]: + return _nanos_literal(self.value * 1_000, TimestampNanoLiteral) + + @to.register(TimestamptzNanoType) + def _(self, _: TimestamptzNanoType) -> Literal[int]: + return _nanos_literal(self.value * 1_000, TimestamptzNanoLiteral) + + +NANOS_PER_DAY = 86_400 * 1_000_000_000 + + +def _nanos_literal(nanos: int, literal_type: type[_TimestampNanoLiteralBase]) -> Literal[int]: + """Wrap nanoseconds from epoch in a nanosecond timestamp literal, or AboveMax/BelowMin when out of the long range.""" + if LongType.max < nanos: + return LongAboveMax() + elif LongType.min > nanos: + return LongBelowMin() + return literal_type(nanos) + + +class _TimestampNanoLiteralBase(Literal[int]): + """Base for nanosecond timestamp literals; the value is nanoseconds from 1970-01-01T00:00:00.""" + + def __init__(self, value: int) -> None: + super().__init__(value, int) + + def increment(self) -> Literal[int]: + return type(self)(self.value + 1) + + def decrement(self) -> Literal[int]: + return type(self)(self.value - 1) + + @singledispatchmethod + def to(self, type_var: IcebergType) -> Literal: # type: ignore + raise TypeError(f"Cannot convert {type(self).__name__} into {type_var}") + + @to.register(TimestampNanoType) + def _(self, _: TimestampNanoType) -> Literal[int]: + return self if isinstance(self, TimestampNanoLiteral) else TimestampNanoLiteral(self.value) + + @to.register(TimestamptzNanoType) + def _(self, _: TimestamptzNanoType) -> Literal[int]: + return self if isinstance(self, TimestamptzNanoLiteral) else TimestamptzNanoLiteral(self.value) + + @to.register(TimestampType) + def _(self, _: TimestampType) -> Literal[int]: + return TimestampLiteral(self.value // 1_000) + + @to.register(TimestamptzType) + def _(self, _: TimestamptzType) -> Literal[int]: + return TimestampLiteral(self.value // 1_000) + + @to.register(DateType) + def _(self, _: DateType) -> Literal[int]: + return DateLiteral(nanos_to_days(self.value)) + + +class TimestampNanoLiteral(_TimestampNanoLiteralBase): + """Literal for timestamp_ns values, in nanoseconds from 1970-01-01T00:00:00.""" + + @model_serializer + def ser_model(self) -> str: + return to_human_timestamp_ns(self.root) + + +class TimestamptzNanoLiteral(_TimestampNanoLiteralBase): + """Literal for timestamptz_ns values, in nanoseconds from 1970-01-01T00:00:00+00:00.""" + + @model_serializer + def ser_model(self) -> str: + return to_human_timestamptz_ns(self.root) + class DecimalLiteral(Literal[Decimal]): def __init__(self, value: Decimal) -> None: @@ -621,6 +719,14 @@ def _(self, _: TimestampType) -> Literal[int]: def _(self, _: TimestamptzType) -> Literal[int]: return TimestampLiteral(timestamptz_to_micros(self.value)) + @to.register(TimestampNanoType) + def _(self, _: TimestampNanoType) -> Literal[int]: + return _nanos_literal(timestamp_to_nanos(self.value), TimestampNanoLiteral) + + @to.register(TimestamptzNanoType) + def _(self, _: TimestamptzNanoType) -> Literal[int]: + return _nanos_literal(timestamptz_to_nanos(self.value), TimestamptzNanoLiteral) + @to.register(UUIDType) def _(self, _: UUIDType) -> Literal[bytes]: return UUIDLiteral(UUID(self.value).bytes) @@ -765,3 +871,57 @@ def _(self, type_var: UUIDType) -> Literal[bytes]: raise TypeError( f"Cannot convert BinaryLiteral into {type_var}, different length: {UUID_BYTES_LENGTH} <> {len(self.value)}" ) + + @to.register(GeometryType) + def _(self, _: GeometryType) -> Literal[bytes]: + return GeometryLiteral(self.value) + + @to.register(GeographyType) + def _(self, _: GeographyType) -> Literal[bytes]: + return GeographyLiteral(self.value) + + +class GeometryLiteral(Literal[bytes]): + """A geometry value, encoded as WKB. Only equality is defined for geometry literals.""" + + def __init__(self, value: bytes) -> None: + super().__init__(value, bytes) + + @model_serializer + def ser_model(self) -> str: + return self.root.hex() + + @singledispatchmethod + def to(self, type_var: IcebergType) -> Literal: # type: ignore + raise TypeError(f"Cannot convert GeometryLiteral into {type_var}") + + @to.register(GeometryType) + def _(self, _: GeometryType) -> Literal[bytes]: + return self + + @to.register(BinaryType) + def _(self, _: BinaryType) -> Literal[bytes]: + return BinaryLiteral(self.value) + + +class GeographyLiteral(Literal[bytes]): + """A geography value, encoded as WKB. Only equality is defined for geography literals.""" + + def __init__(self, value: bytes) -> None: + super().__init__(value, bytes) + + @model_serializer + def ser_model(self) -> str: + return self.root.hex() + + @singledispatchmethod + def to(self, type_var: IcebergType) -> Literal: # type: ignore + raise TypeError(f"Cannot convert GeographyLiteral into {type_var}") + + @to.register(GeographyType) + def _(self, _: GeographyType) -> Literal[bytes]: + return self + + @to.register(BinaryType) + def _(self, _: BinaryType) -> Literal[bytes]: + return BinaryLiteral(self.value) diff --git a/pyiceberg/expressions/visitors.py b/pyiceberg/expressions/visitors.py index 5072d3de11..9676a25160 100644 --- a/pyiceberg/expressions/visitors.py +++ b/pyiceberg/expressions/visitors.py @@ -61,6 +61,8 @@ from pyiceberg.types import ( DoubleType, FloatType, + GeographyType, + GeometryType, IcebergType, NestedField, PrimitiveType, @@ -1179,6 +1181,17 @@ def _is_nan(self, val: Any) -> bool: return False +def _is_geo_predicate(predicate: BoundPredicate) -> bool: + """Return whether a literal or set predicate references a geometry or geography column. + + The lower and upper bounds of geo columns are bounding-box corner points rather than ordered values, so + they cannot be compared with a literal; only the null and value counts can be used for these predicates. + """ + return isinstance(predicate, (BoundLiteralPredicate, BoundSetPredicate)) and isinstance( + predicate.term.ref().field.field_type, (GeometryType, GeographyType) + ) + + class _InclusiveMetricsEvaluator: """Bind an inclusive metrics expression once and evaluate files without mutating prepared state.""" @@ -1208,6 +1221,13 @@ def eval(self, file: DataFile) -> bool: class _InclusiveMetricsEvaluationVisitor(_MetricsEvaluationVisitor): """Evaluate inclusive metrics for one data file.""" + def visit_bound_predicate(self, predicate: BoundPredicate) -> bool: + if _is_geo_predicate(predicate): + if isinstance(predicate, (BoundEqualTo, BoundIn)) and self._contains_nulls_only(predicate.term.ref().field.field_id): + return ROWS_CANNOT_MATCH + return ROWS_MIGHT_MATCH + return super().visit_bound_predicate(predicate) + def _may_contain_null(self, field_id: int) -> bool: # A missing null count means the count is unknown, so the column may contain nulls. null_count = self.null_counts.get(field_id) @@ -1558,6 +1578,15 @@ def __init__(self, struct: StructType, file: DataFile) -> None: super().__init__(file) self.struct = struct + def visit_bound_predicate(self, predicate: BoundPredicate) -> bool: + if _is_geo_predicate(predicate): + if isinstance(predicate, (BoundNotEqualTo, BoundNotIn)) and self._contains_nulls_only( + predicate.term.ref().field.field_id + ): + return ROWS_MUST_MATCH + return ROWS_MIGHT_NOT_MATCH + return super().visit_bound_predicate(predicate) + def visit_is_null(self, term: BoundTerm) -> bool: # no need to check whether the field is required because binding evaluates that case # if the column has any non-null values, the expression does not match diff --git a/pyiceberg/io/pyarrow.py b/pyiceberg/io/pyarrow.py index 1e59107da3..e183422c9e 100644 --- a/pyiceberg/io/pyarrow.py +++ b/pyiceberg/io/pyarrow.py @@ -30,7 +30,9 @@ import functools import importlib import itertools +import json import logging +import math import operator import os import re @@ -147,14 +149,17 @@ visit_with_partner, ) from pyiceberg.table import DOWNCAST_NS_TIMESTAMP_TO_US_ON_WRITE, TableProperties -from pyiceberg.table.deletion_vector import deletion_vectors_from_puffin_file +from pyiceberg.table.delete_file import DeleteFileSet +from pyiceberg.table.deletion_vector import read_deletion_vectors from pyiceberg.table.locations import load_location_provider from pyiceberg.table.metadata import TableMetadata +from pyiceberg.table.metadata_columns import LAST_UPDATED_SEQUENCE_NUMBER, ROW_ID, schema_with_row_lineage from pyiceberg.table.name_mapping import NameMapping, apply_name_mapping -from pyiceberg.table.puffin import PuffinFile from pyiceberg.transforms import IdentityTransform, TruncateTransform from pyiceberg.typedef import EMPTY_DICT, Properties, Record, TableVersion from pyiceberg.types import ( + DEFAULT_GEOMETRY_CRS, + GEOGRAPHY_ALGORITHMS, BinaryType, BooleanType, DateType, @@ -180,12 +185,14 @@ TimeType, UnknownType, UUIDType, + VariantType, strtobool, ) from pyiceberg.utils.concurrent import ExecutorFactory from pyiceberg.utils.config import Config from pyiceberg.utils.datetime import millis_to_datetime from pyiceberg.utils.decimal import unscaled_to_decimal +from pyiceberg.utils.geo import geo_bounds_from_bbox from pyiceberg.utils.properties import get_first_property_value, property_as_bool, property_as_int from pyiceberg.utils.singleton import Singleton from pyiceberg.utils.truncate import truncate_upper_bound_binary_string, truncate_upper_bound_text_string @@ -828,6 +835,10 @@ def visit_unknown(self, _: UnknownType) -> pa.DataType: def visit_binary(self, _: BinaryType) -> pa.DataType: return pa.large_binary() + def visit_variant(self, _: VariantType) -> pa.DataType: + """Convert variant type to the struct that PyArrow reads an unshredded Parquet VARIANT group as.""" + return _VARIANT_ARROW_TYPE + def visit_geometry(self, geometry_type: GeometryType) -> pa.DataType: """Convert geometry type to PyArrow type. @@ -851,15 +862,66 @@ def visit_geography(self, geography_type: GeographyType) -> pa.DataType: import geoarrow.pyarrow as ga wkb_type = ga.wkb().with_crs(geography_type.crs) - # Map Iceberg algorithm to GeoArrow edge type - if geography_type.algorithm == "spherical": - wkb_type = wkb_type.with_edge_type(ga.EdgeType.SPHERICAL) + # Map the Iceberg edge-interpolation algorithm to the GeoArrow edge type of the same name + if geography_type.algorithm in GEOGRAPHY_ALGORITHMS: + wkb_type = wkb_type.with_edge_type(ga.EdgeType[geography_type.algorithm.upper()]) # "planar" is the default edge type in GeoArrow, no need to set explicitly return wkb_type except ImportError: return pa.large_binary() +GEOARROW_WKB_EXTENSION_NAME = "geoarrow.wkb" + + +def _geoarrow_available() -> bool: + """Return whether geoarrow-pyarrow is installed.""" + try: + import geoarrow.pyarrow # noqa: F401 + + return True + except ImportError: + return False + + +def _geoarrow_crs_to_iceberg(crs: Any) -> str: + """Convert the CRS of a GeoArrow extension type to an Iceberg CRS string.""" + if crs is None: + return DEFAULT_GEOMETRY_CRS + if isinstance(crs, str): + return crs + if isinstance(crs, dict) and isinstance(crs_id := crs.get("id"), dict) and "authority" in crs_id and "code" in crs_id: + # PROJJSON, e.g. {"id": {"authority": "EPSG", "code": 3857}} + return f"{crs_id['authority']}:{crs_id['code']}" + raise TypeError(f"Unsupported GeoArrow CRS, expected a string or PROJJSON with an id: {crs}") + + +def _geoarrow_wkb_to_iceberg(extension_type: pa.ExtensionType) -> PrimitiveType: + """Convert a GeoArrow WKB extension type to an Iceberg geometry or geography type. + + The extension metadata is JSON with an optional `crs` and `edges`; planar edges (the default) map to a + geometry, any other edge type to a geography with the algorithm of the same name. + """ + serialized = extension_type.__arrow_ext_serialize__() + metadata = json.loads(serialized) if serialized else {} + crs = _geoarrow_crs_to_iceberg(metadata.get("crs")) + edges = metadata.get("edges") + if edges is None or edges == "planar": + return GeometryType(crs) + if edges in GEOGRAPHY_ALGORITHMS: + return GeographyType(crs, edges) + raise TypeError(f"Unsupported GeoArrow edge type: {edges}") + + +def _cast_to_geo(values: pa.Array, target_type: pa.DataType) -> pa.Array: + """Cast WKB values (binary, large_binary or a WKB extension array) to the Arrow type of a geo field.""" + if isinstance(values.type, pa.ExtensionType): + values = values.storage + if isinstance(target_type, pa.ExtensionType): + return pa.ExtensionArray.from_storage(target_type, values.cast(target_type.storage_type)) + return values.cast(target_type) + + def _convert_scalar(value: Any, iceberg_type: IcebergType) -> pa.scalar: if not isinstance(iceberg_type, PrimitiveType): raise ValueError(f"Expected primitive type, got: {iceberg_type}") @@ -1155,10 +1217,7 @@ def _read_deletes(io: FileIO, data_file: DataFile) -> dict[str, pa.ChunkedArray] for path in table.column("file_path").unique() } elif data_file.file_format == FileFormat.PUFFIN: - with io.new_input(data_file.file_path).open() as fi: - payload = fi.read() - - return {dv.referenced_data_file: dv.to_vector() for dv in deletion_vectors_from_puffin_file(PuffinFile(payload))} + return {dv.referenced_data_file: dv.to_vector() for dv in read_deletion_vectors(io, data_file)} else: raise ValueError(f"Delete file format not supported: {data_file.file_format}") @@ -1240,16 +1299,52 @@ def visit_pyarrow(obj: pa.DataType | pa.Schema, visitor: PyArrowSchemaVisitor[T] @visit_pyarrow.register(pa.Schema) def _(obj: pa.Schema, visitor: PyArrowSchemaVisitor[T]) -> T: - return visitor.schema(obj, visit_pyarrow(pa.struct(obj), visitor)) + # The root is always a struct, even when its columns are named like the fields of a variant + struct = pa.struct(obj) + return visitor.schema(obj, visitor.struct(struct, [visit_pyarrow(field, visitor) for field in struct])) @visit_pyarrow.register(pa.StructType) def _(obj: pa.StructType, visitor: PyArrowSchemaVisitor[T]) -> T: + if _is_variant_group(obj): + return visitor.variant(obj) + results = [visit_pyarrow(field, visitor) for field in obj] return visitor.struct(obj, results) +_VARIANT_METADATA = "metadata" +_VARIANT_VALUE = "value" +_VARIANT_TYPED_VALUE = "typed_value" +_VARIANT_ARROW_TYPE = pa.struct( + [ + pa.field(_VARIANT_METADATA, pa.binary(), nullable=False), + pa.field(_VARIANT_VALUE, pa.binary(), nullable=False), + ] +) + + +def _is_binary_like(arrow_type: pa.DataType) -> bool: + return pa.types.is_binary(arrow_type) or pa.types.is_large_binary(arrow_type) or pa.types.is_binary_view(arrow_type) + + +def _is_variant_group(struct: pa.StructType) -> bool: + """Return whether a struct is a Parquet VARIANT group, unshredded or shredded. + + The spec forbids field ids on the ``metadata``, ``value`` and ``typed_value`` fields of a variant, while + every Iceberg struct field carries one, which tells the two apart. + """ + names = {field.name for field in struct} + if _VARIANT_METADATA not in names or not names <= {_VARIANT_METADATA, _VARIANT_VALUE, _VARIANT_TYPED_VALUE}: + return False + if names == {_VARIANT_METADATA}: + return False + if any(_get_field_id(field) is not None for field in struct): + return False + return all(_is_binary_like(field.type) for field in struct if field.name != _VARIANT_TYPED_VALUE) + + @visit_pyarrow.register(pa.ListType) @visit_pyarrow.register(pa.FixedSizeListType) @visit_pyarrow.register(pa.LargeListType) @@ -1308,6 +1403,10 @@ class PyArrowSchemaVisitor(Generic[T], ABC): def before_field(self, field: pa.Field) -> None: """Override this method to perform an action immediately before visiting a field.""" + def variant(self, variant_group: pa.StructType) -> T: + """Visit a Parquet VARIANT group, which is a plain struct unless a visitor overrides this.""" + return self.struct(variant_group, [visit_pyarrow(field, self) for field in variant_group]) + def after_field(self, field: pa.Field) -> None: """Override this method to perform an action immediately after visiting a field.""" @@ -1372,6 +1471,10 @@ class _HasIds(PyArrowSchemaVisitor[bool]): def schema(self, schema: pa.Schema, struct_result: bool) -> bool: return struct_result + def variant(self, variant_group: pa.StructType) -> bool: + # The fields of a variant never carry field ids + return True + def struct(self, struct: pa.StructType, field_results: builtins.list[bool]) -> bool: return all(field_results) @@ -1418,11 +1521,16 @@ def schema(self, schema: pa.Schema, struct_result: StructType) -> Schema: def struct(self, struct: pa.StructType, field_results: builtins.list[NestedField]) -> StructType: return StructType(*field_results) + def variant(self, variant_group: pa.StructType) -> IcebergType | Schema: + return VariantType() + def field(self, field: pa.Field, field_result: IcebergType) -> NestedField: field_id = self._field_id(field) field_doc = doc_str.decode() if (field.metadata and (doc_str := field.metadata.get(PYARROW_FIELD_DOC_KEY))) else None field_type = field_result - return NestedField(field_id, field.name, field_type, required=not field.nullable, doc=field_doc) + # A column of type unknown is always optional + required = not field.nullable and not isinstance(field_type, UnknownType) + return NestedField(field_id, field.name, field_type, required=required, doc=field_doc) def list(self, list_type: pa.ListType, element_result: IcebergType) -> ListType: element_field = list_type.value_field @@ -1443,7 +1551,10 @@ def map(self, map_type: pa.MapType, key_result: IcebergType, value_result: Icebe return MapType(key_id, key_result, value_id, value_result, value_required=not value_field.nullable) def primitive(self, primitive: pa.DataType) -> PrimitiveType: - if pa.types.is_boolean(primitive): + if isinstance(primitive, pa.ExtensionType) and primitive.extension_name == GEOARROW_WKB_EXTENSION_NAME: + # Only produced when geoarrow-pyarrow has registered the extension type + return _geoarrow_wkb_to_iceberg(primitive) + elif pa.types.is_boolean(primitive): return BooleanType() elif pa.types.is_integer(primitive): width = primitive.bit_width @@ -1605,6 +1716,10 @@ class _ConvertToIcebergWithoutIDs(_ConvertToIceberg): def _field_id(self, field: pa.Field) -> int: return -1 + def variant(self, variant_group: pa.StructType) -> IcebergType | Schema: + # Without field ids a variant cannot be told apart from a struct with metadata and value fields + return PyArrowSchemaVisitor.variant(self, variant_group) + def _get_column_projection_values( file: DataFile, @@ -1632,6 +1747,178 @@ def _get_column_projection_values( return projected_missing_fields +# Temporary column holding each row's position in the data file while filters are applied +_ROW_POSITION_COLUMN = "__pyiceberg_row_position" + + +@dataclass(frozen=True) +class _EqualityDeletes: + """The rows of one equality delete file, projected onto its delete columns. + + Attributes: + field_ids: The field ids of the delete columns (the file's ``equality_ids``). + rows: The delete rows, with one column per delete column named after its field id. + """ + + field_ids: tuple[int, ...] + rows: pa.Table + + +def _field_path(schema: Schema, field_id: int) -> list[str] | None: + """Return the names from the root of the schema to the field, or None when the field is not in the schema.""" + if field_id not in schema._lazy_id_to_field: + return None + parents = schema._lazy_id_to_parent + names = [schema.find_field(field_id).name] + parent_id = parents.get(field_id) + while parent_id is not None: + names.append(schema.find_field(parent_id).name) + parent_id = parents.get(parent_id) + return list(reversed(names)) + + +def _column_by_path(data: pa.RecordBatch | pa.Table, path: list[str]) -> pa.Array | pa.ChunkedArray: + """Return a (possibly nested) column; a null parent struct yields a null value.""" + column = data.column(path[0]) + for name in path[1:]: + column = pc.struct_field(column, name) + return column + + +def _read_equality_deletes( + io: FileIO, + delete_file: DataFile, + name_mapping: NameMapping | None = None, + format_version: TableVersion = TableProperties.DEFAULT_FORMAT_VERSION, +) -> _EqualityDeletes: + """Read the delete columns of an equality delete file, matching them to the table by field id. + + Other columns stored in the delete file are ignored. + """ + if not delete_file.equality_ids: + raise ValueError(f"Equality delete file has no equality_ids: {delete_file.file_path}") + if delete_file.file_format not in (FileFormat.PARQUET, FileFormat.ORC): + raise ValueError(f"Equality delete file format not supported: {delete_file.file_format}") + + with io.new_input(delete_file.file_path).open() as fi: + fragment = _get_file_format(delete_file.file_format).make_fragment(fi) + file_schema = pyarrow_to_schema( + fragment.physical_schema, + name_mapping, + downcast_ns_timestamp_to_us=format_version <= 2, + format_version=format_version, + ) + paths: dict[int, list[str]] = {} + for field_id in delete_file.equality_ids: + path = _field_path(file_schema, field_id) + if path is None: + raise ValueError(f"Equality delete file {delete_file.file_path} is missing the delete column with id {field_id}") + paths[field_id] = path + columns = list(dict.fromkeys(path[0] for path in paths.values())) + table = ds.Scanner.from_fragment(fragment=fragment, columns=columns).to_table() + + rows = pa.table({str(field_id): _column_by_path(table, path) for field_id, path in paths.items()}) + return _EqualityDeletes(field_ids=tuple(paths), rows=rows) + + +def _normalize_key_array(array: pa.Array | pa.ChunkedArray) -> pa.Array: + """Return a plain array that Arrow compute kernels can compare.""" + if isinstance(array, pa.ChunkedArray): + array = array.combine_chunks() + if isinstance(array, pa.DictionaryArray): + array = array.dictionary_decode() + if isinstance(array, pa.ExtensionArray): + array = array.storage + return array + + +def _align_key_arrays(data: pa.Array, deletes: pa.Array) -> tuple[pa.Array, pa.Array]: + """Cast the data and delete values of one delete column to a common type.""" + if data.type == deletes.type: + return data, deletes + try: + return data.cast(deletes.type), deletes + except (pa.ArrowInvalid, pa.ArrowNotImplementedError): + return data, deletes.cast(data.type) + + +def _equality_delete_mask(data_keys: list[pa.Array], delete_keys: list[pa.Array | pa.ChunkedArray]) -> pa.Array: + """Return a mask that is true for each data row equal to any delete row on all delete columns. + + Following the spec, a null value is equal to a null value. + """ + num_rows = len(data_keys[0]) + aligned = [ + _align_key_arrays(_normalize_key_array(data), _normalize_key_array(deletes)) + for data, deletes in zip(data_keys, delete_keys, strict=True) + ] + + if len(aligned) == 1: + data, deletes = aligned[0] + # is_in matches a null input against a null in the value set + return pc.fill_null(pc.is_in(data, value_set=deletes), False) + + # A hash join never matches null keys, so each column is joined on a null flag plus its value with nulls filled + left: dict[str, pa.Array] = {"__pos": pa.array(range(num_rows), type=pa.int64())} + right: dict[str, pa.Array] = {} + keys: list[str] = [] + for index, (data, deletes) in enumerate(aligned): + keys.append(f"null_{index}") + left[keys[-1]] = pc.is_null(data) + right[keys[-1]] = pc.is_null(deletes) + non_null = pc.drop_null(data) if data.null_count < len(data) else pc.drop_null(deletes) + if len(non_null) > 0: + keys.append(f"value_{index}") + left[keys[-1]] = pc.fill_null(data, non_null[0]) + right[keys[-1]] = pc.fill_null(deletes, non_null[0]) + + delete_table = pa.table(right).group_by(keys).aggregate([]) + matched = pa.table(left).join(delete_table, keys=keys, join_type="left semi") + return pc.is_in(left["__pos"], value_set=matched.column("__pos")) + + +def _equality_key_arrays( + batch: pa.RecordBatch, + deletes: _EqualityDeletes, + data_paths: dict[int, list[str] | None], + missing_values: dict[int, Any], +) -> list[pa.Array]: + """Return the values of the delete columns for the rows of a data batch.""" + keys: list[pa.Array] = [] + for field_id in deletes.field_ids: + path = data_paths.get(field_id) + if path is not None: + keys.append(_column_by_path(batch, path)) + else: + # A delete column that the data file does not contain takes its projected value, or null + delete_type = deletes.rows.column(str(field_id)).type + value = missing_values.get(field_id) + keys.append( + pa.nulls(batch.num_rows, delete_type) + if value is None + else pa.repeat(pa.scalar(value, delete_type), batch.num_rows) + ) + return keys + + +def _equality_missing_values( + task: FileScanTask, table_schema: Schema, partition_spec: PartitionSpec | None, field_ids: set[int] +) -> dict[int, Any]: + """Return the projected values of delete columns that a data file does not contain.""" + if not field_ids: + return EMPTY_DICT + values = dict( + _get_column_projection_values( + task.file, prune_columns(table_schema, field_ids, select_full_types=False), table_schema, partition_spec, set() + ) + ) + for field_id in field_ids: + field = table_schema._lazy_id_to_field.get(field_id) + if field_id not in values and field is not None and field.initial_default is not None: + values[field_id] = field.initial_default + return {field_id: value for field_id, value in values.items() if field_id in field_ids} + + def _task_to_record_batches( io: FileIO, task: FileScanTask, @@ -1646,6 +1933,7 @@ def _task_to_record_batches( format_version: TableVersion = TableProperties.DEFAULT_FORMAT_VERSION, downcast_ns_timestamp_to_us: bool | None = None, dictionary_columns: tuple[str, ...] = (), + equality_deletes: list[_EqualityDeletes] | None = None, ) -> Iterator[pa.RecordBatch]: format_kwargs: dict[str, Any] = {"pre_buffer": True, "buffer_size": ONE_MEGABYTE * 8} if dictionary_columns and task.file.file_format == FileFormat.PARQUET: @@ -1678,14 +1966,24 @@ def _task_to_record_batches( bound_file_filter = bind(file_schema, translated_row_filter, case_sensitive=case_sensitive) pyarrow_filter = expression_to_pyarrow(bound_file_filter, file_schema) - file_project_schema = prune_columns(file_schema, projected_field_ids, select_full_types=False) + # Equality deletes compare the delete columns, so they are read even when not projected + equality_field_ids = {field_id for deletes in equality_deletes or () for field_id in deletes.field_ids} + file_project_schema = prune_columns(file_schema, projected_field_ids | equality_field_ids, select_full_types=False) + equality_paths = {field_id: _field_path(file_schema, field_id) for field_id in equality_field_ids} + equality_missing_values = _equality_missing_values( + task, table_schema, partition_spec, {field_id for field_id, path in equality_paths.items() if path is None} + ) + + # Row lineage columns are derived from the position of each row in the file + track_positions = ROW_ID.field_id in projected_schema.field_ids + apply_filter_after_read = bool(positional_deletes) or track_positions fragment_scanner = ds.Scanner.from_fragment( fragment=fragment, schema=physical_schema, # This will push down the query to Arrow. - # But in case there are positional deletes, we have to apply them first - filter=pyarrow_filter if not positional_deletes else None, + # But in case there are positional deletes or row positions are needed, we have to apply them first + filter=pyarrow_filter if not apply_filter_after_read else None, columns=[col.name for col in file_project_schema.columns], ) @@ -1695,27 +1993,50 @@ def _task_to_record_batches( next_index = next_index + len(batch) current_index = next_index - len(batch) current_batch = batch + positions = _row_positions(current_index, next_index) if track_positions else None if positional_deletes: # Create the mask of indices that we're interested in indices = _combine_positional_deletes(positional_deletes, current_index, current_index + len(batch)) current_batch = current_batch.take(indices) - if pyarrow_filter is not None: - # Temporary fix until PyArrow 21 is the minimum supported version - # (https://github.com/apache/arrow/pull/46057): RecordBatch.filter raises - # IndexError on PyArrow <21 when the result is empty; Table.filter does not. - table = pa.Table.from_batches([current_batch]) - table = table.filter(pyarrow_filter) - if table.num_rows == 0: - current_batch = current_batch.slice(0, 0) - else: - current_batch = table.combine_chunks().to_batches()[0] + if positions is not None: + positions = positions.take(indices) + + if equality_deletes and current_batch.num_rows > 0: + # Equality deletes apply after positional deletes, so positions still refer to the data file + deleted = None + for deletes in equality_deletes: + data_keys = _equality_key_arrays(current_batch, deletes, equality_paths, equality_missing_values) + delete_keys = [deletes.rows.column(str(field_id)) for field_id in deletes.field_ids] + mask = _equality_delete_mask(data_keys, delete_keys) + deleted = mask if deleted is None else pc.or_(deleted, mask) + if deleted is not None: + indices = pc.indices_nonzero(pc.invert(deleted)) + current_batch = current_batch.take(indices) + if positions is not None: + positions = positions.take(indices) + + if apply_filter_after_read and pyarrow_filter is not None: + if positions is not None: + current_batch = current_batch.append_column(_ROW_POSITION_COLUMN, positions) + # Temporary fix until PyArrow 21 is the minimum supported version + # (https://github.com/apache/arrow/pull/46057): RecordBatch.filter raises + # IndexError on PyArrow <21 when the result is empty; Table.filter does not. + table = pa.Table.from_batches([current_batch]) + table = table.filter(pyarrow_filter) + if table.num_rows == 0: + current_batch = current_batch.slice(0, 0) + else: + current_batch = table.combine_chunks().to_batches()[0] + if positions is not None: + positions = current_batch.column(_ROW_POSITION_COLUMN) + current_batch = current_batch.drop_columns([_ROW_POSITION_COLUMN]) # skip empty batches if current_batch.num_rows == 0: continue - yield _to_requested_schema( + result_batch = _to_requested_schema( projected_schema, file_project_schema, current_batch, @@ -1723,19 +2044,66 @@ def _task_to_record_batches( projected_missing_fields=projected_missing_fields, allow_timestamp_tz_mismatch=True, ) + yield _inherit_row_lineage(result_batch, projected_schema, task, positions) + + +def _row_positions(start: int, stop: int) -> pa.Array: + """Return the int64 positions ``start`` (inclusive) to ``stop`` (exclusive).""" + return pc.add(pc.cumulative_sum(pa.repeat(pa.scalar(1, pa.int64()), stop - start)), pa.scalar(start - 1, pa.int64())) + + +def _inherit_row_lineage( + batch: pa.RecordBatch, projected_schema: Schema, task: FileScanTask, positions: pa.Array | None +) -> pa.RecordBatch: + """Fill null row lineage metadata columns from the data file's first row ID and data sequence number. + + Rows keep values that are stored in the data file. Files added before row lineage was enabled have no + first_row_id, and both columns stay null for them. + """ + first_row_id = task.file.first_row_id + for index, field in enumerate(projected_schema.fields): + if field.field_id == ROW_ID.field_id: + inherited = ( + None if first_row_id is None or positions is None else pc.add(positions, pa.scalar(first_row_id, pa.int64())) + ) + elif field.field_id == LAST_UPDATED_SEQUENCE_NUMBER.field_id: + inherited = ( + None + if first_row_id is None or task.sequence_number is None + else pa.repeat(pa.scalar(task.sequence_number, pa.int64()), batch.num_rows) + ) + else: + continue + + if inherited is not None: + column = batch.column(index) + values = inherited if column.null_count == len(column) else pc.coalesce(column, inherited) + batch = batch.set_column(index, batch.schema.field(index), values.cast(column.type)) + + return batch def _read_all_delete_files(io: FileIO, tasks: Iterable[FileScanTask]) -> dict[str, list[ChunkedArray]]: deletes_per_file: dict[str, list[ChunkedArray]] = {} - unique_deletes = set(itertools.chain.from_iterable([task.delete_files for task in tasks])) + # A delete file can hold positions for several data files, but it only applies to the data files + # it was planned for (for example, a DV supersedes older position deletes for its data file) + planned_deletes: dict[str, DeleteFileSet] = {} + for task in tasks: + planned_deletes.setdefault(task.file.file_path, DeleteFileSet()).update( + delete_file for delete_file in task.delete_files if delete_file.content != DataFileContent.EQUALITY_DELETES + ) + unique_deletes = DeleteFileSet(itertools.chain.from_iterable(planned_deletes.values())) if len(unique_deletes) > 0: executor = ExecutorFactory.get_or_create() + delete_files = list(unique_deletes) deletes_per_files: Iterator[dict[str, ChunkedArray]] = executor.map( lambda args: _read_deletes(*args), - [(io, delete_file) for delete_file in unique_deletes], + [(io, delete_file) for delete_file in delete_files], ) - for delete in deletes_per_files: + for delete_file, delete in zip(delete_files, deletes_per_files, strict=True): for file, arr in delete.items(): + if delete_file not in planned_deletes.get(file, ()): + continue if file in deletes_per_file: deletes_per_file[file].append(arr) else: @@ -1744,6 +2112,41 @@ def _read_all_delete_files(io: FileIO, tasks: Iterable[FileScanTask]) -> dict[st return deletes_per_file +def _read_all_equality_delete_files( + io: FileIO, + tasks: Iterable[FileScanTask], + name_mapping: NameMapping | None = None, + format_version: TableVersion = TableProperties.DEFAULT_FORMAT_VERSION, +) -> dict[str, list[_EqualityDeletes]]: + """Read every equality delete file planned for the tasks once, and group them by data file path.""" + planned_deletes: dict[str, DeleteFileSet] = {} + for task in tasks: + equality_deletes = [ + delete_file for delete_file in task.delete_files if delete_file.content == DataFileContent.EQUALITY_DELETES + ] + if equality_deletes: + planned_deletes.setdefault(task.file.file_path, DeleteFileSet()).update(equality_deletes) + + if not planned_deletes: + return {} + + unique_deletes = list(DeleteFileSet(itertools.chain.from_iterable(planned_deletes.values()))) + executor = ExecutorFactory.get_or_create() + read_deletes = dict( + zip( + (delete_file.file_path for delete_file in unique_deletes), + executor.map( + lambda delete_file: _read_equality_deletes(io, delete_file, name_mapping, format_version), unique_deletes + ), + strict=True, + ) + ) + return { + data_file_path: [read_deletes[delete_file.file_path] for delete_file in delete_files] + for data_file_path, delete_files in planned_deletes.items() + } + + class ArrowScan: _table_metadata: TableMetadata _io: FileIO @@ -1847,7 +2250,11 @@ def to_record_batches(self, tasks: Iterable[FileScanTask]) -> Iterator[pa.Record ResolveError: When a required field cannot be found in the file ValueError: When a field type in the file cannot be projected to the schema type """ + tasks = list(tasks) deletes_per_file = _read_all_delete_files(self._io, tasks) + equality_deletes_per_file = _read_all_equality_delete_files( + self._io, tasks, self._table_metadata.name_mapping(), self._table_metadata.format_version + ) total_row_count = 0 executor = ExecutorFactory.get_or_create() @@ -1856,7 +2263,7 @@ def batches_for_task(task: FileScanTask) -> list[pa.RecordBatch]: # Materialize the iterator here to ensure execution happens within the executor. # Otherwise, the iterator would be lazily consumed later (in the main thread), # defeating the purpose of using executor.map. - return list(self._record_batches_from_scan_tasks_and_deletes([task], deletes_per_file)) + return list(self._record_batches_from_scan_tasks_and_deletes([task], deletes_per_file, equality_deletes_per_file)) limit_reached = False for batches in executor.map(batches_for_task, tasks): @@ -1876,7 +2283,10 @@ def batches_for_task(task: FileScanTask) -> list[pa.RecordBatch]: break def _record_batches_from_scan_tasks_and_deletes( - self, tasks: Iterable[FileScanTask], deletes_per_file: dict[str, list[ChunkedArray]] + self, + tasks: Iterable[FileScanTask], + deletes_per_file: dict[str, list[ChunkedArray]], + equality_deletes_per_file: dict[str, list[_EqualityDeletes]] | None = None, ) -> Iterator[pa.RecordBatch]: total_row_count = 0 for task in tasks: @@ -1896,6 +2306,7 @@ def _record_batches_from_scan_tasks_and_deletes( self._table_metadata.format_version, self._downcast_ns_timestamp_to_us, self._dictionary_columns, + equality_deletes=(equality_deletes_per_file or {}).get(task.file.file_path), ) for batch in batches: if self._limit is not None: @@ -1916,6 +2327,7 @@ def _to_requested_schema( projected_missing_fields: dict[int, Any] = EMPTY_DICT, allow_timestamp_tz_mismatch: bool = False, format_model: FileFormatModel | None = None, + for_write: bool = False, ) -> pa.RecordBatch: # We could reuse some of these visitors struct_array = visit_with_partner( @@ -1928,12 +2340,24 @@ def _to_requested_schema( projected_missing_fields=projected_missing_fields, allow_timestamp_tz_mismatch=allow_timestamp_tz_mismatch, format_model=format_model, + for_write=for_write, ), ArrowAccessor(file_schema), ) return pa.RecordBatch.from_struct_array(struct_array) +def _to_unshredded_variant(values: pa.Array) -> pa.Array: + """Return a variant column as a struct of its binary metadata and value.""" + if not pa.types.is_struct(values.type) or not _is_binary_like(values.type.field(_VARIANT_METADATA).type): + raise ValueError(f"Unsupported schema projection from {values.type} to variant") + if values.type.get_field_index(_VARIANT_TYPED_VALUE) >= 0: + raise NotImplementedError("Reading shredded variant columns (with a typed_value field) is not yet supported") + if values.type.get_field_index(_VARIANT_VALUE) < 0: + raise ValueError(f"Unsupported schema projection from {values.type} to variant") + return values.cast(_VARIANT_ARROW_TYPE) + + class ArrowProjectionVisitor(SchemaWithPartnerVisitor[pa.Array, pa.Array | None]): _file_schema: Schema _include_field_ids: bool @@ -1941,6 +2365,7 @@ class ArrowProjectionVisitor(SchemaWithPartnerVisitor[pa.Array, pa.Array | None] _projected_missing_fields: dict[int, Any] _allow_timestamp_tz_mismatch: bool _format_model: FileFormatModel | None + _for_write: bool def __init__( self, @@ -1950,6 +2375,7 @@ def __init__( projected_missing_fields: dict[int, Any] = EMPTY_DICT, allow_timestamp_tz_mismatch: bool = False, format_model: FileFormatModel | None = None, + for_write: bool = False, ) -> None: if include_field_ids and format_model is None: raise ValueError("format_model is required when include_field_ids=True") @@ -1961,12 +2387,42 @@ def __init__( # Allowed for reading (aligns with Spark); disallowed for writing to enforce Iceberg spec's strict typing. self._allow_timestamp_tz_mismatch = allow_timestamp_tz_mismatch self._format_model = format_model + # When True, the projection produces a batch for a data file: missing columns are filled from the + # write-default, unknown columns are dropped, and geo columns use a type the Parquet writer supports + self._for_write = for_write + + def _arrow_type(self, field_type: IcebergType) -> pa.DataType: + """Return the Arrow type of a projected column.""" + if self._for_write and isinstance(field_type, GeographyType) and field_type.algorithm != "spherical": + # The Parquet writer only supports spherical edges for the GEOGRAPHY logical type + warnings.warn( + f"Writing geography with the {field_type.algorithm} edge algorithm as plain WKB binary: " + "the Parquet GEOGRAPHY logical type is only written for spherical edges", + stacklevel=2, + ) + return pa.large_binary() + if self._for_write and isinstance(field_type, (GeometryType, GeographyType)) and not _geoarrow_available(): + warnings.warn( + "Writing geometry/geography as plain WKB binary; install pyiceberg[geoarrow] to write the Parquet " + "GEOMETRY/GEOGRAPHY logical types", + stacklevel=2, + ) + return schema_to_pyarrow(field_type, include_field_ids=self._include_field_ids) def _cast_if_needed(self, field: NestedField, values: pa.Array) -> pa.Array: file_field = self._file_schema.find_field(field.field_id) + if isinstance(field.field_type, VariantType): + # The file schema of a variant may be a pruned struct, so the values are checked instead + return _to_unshredded_variant(values) + if field.field_type.is_primitive: - if (target_type := schema_to_pyarrow(field.field_type, include_field_ids=self._include_field_ids)) != values.type: + if (target_type := self._arrow_type(field.field_type)) != values.type: + if isinstance(file_field.field_type, DateType) and isinstance( + field.field_type, (TimestampType, TimestampNanoType) + ): + # Format version 3 date -> timestamp promotion reads a date as midnight + return values.cast(target_type) if field.field_type == TimestampType(): source_tz_compatible = values.type.tz is None or ( self._allow_timestamp_tz_mismatch and values.type.tz in UTC_ALIASES @@ -1996,6 +2452,16 @@ def _cast_if_needed(self, field: NestedField, values: pa.Array) -> pa.Array: elif target_type.unit == "us" and values.type.unit in {"s", "ms", "us"}: return values.cast(target_type) raise ValueError(f"Unsupported schema projection from {values.type} to {target_type}") + elif isinstance(field.field_type, (TimestampNanoType, TimestamptzNanoType)): + # Files written with downcast-ns-timestamp-to-us-on-write store microseconds for nanosecond columns. + # The Arrow type carries the unit, so upcasting is exact. This is deliberately not a promote() rule, + # which would also allow timestamp -> timestamp_ns schema evolution and misread Avro longs. + if ( + pa.types.is_timestamp(values.type) + and values.type.unit in {"s", "ms", "us"} + and (values.type.tz in UTC_ALIASES if target_type.tz else values.type.tz is None) + ): + return values.cast(target_type) elif isinstance(field.field_type, (IntegerType, LongType)): # Cast smaller integer types to target type for cross-platform compatibility # Only allow widening conversions (smaller bit width to larger) @@ -2014,6 +2480,12 @@ def _cast_if_needed(self, field: NestedField, values: pa.Array) -> pa.Array: target_width = target_type.bit_width if source_width < target_width: return values.cast(target_type) + elif isinstance(field.field_type, (GeometryType, GeographyType)): + # Parquet GEOMETRY/GEOGRAPHY columns are read as WKB binary; attach the geo type of the table + # schema (a GeoArrow extension type when geoarrow-pyarrow is installed, else large_binary) + storage_type = values.type.storage_type if isinstance(values.type, pa.ExtensionType) else values.type + if pa.types.is_binary(storage_type) or pa.types.is_large_binary(storage_type): + return _cast_to_geo(values, target_type) if field.field_type != file_field.field_type: target_schema = schema_to_pyarrow( @@ -2048,20 +2520,25 @@ def struct( field_arrays: list[pa.Array] = [] fields: list[pa.Field] = [] for field, field_array in zip(struct.fields, field_results, strict=True): + if self._for_write and isinstance(field.field_type, UnknownType): + # Columns of type unknown are always null and are not stored in data files + continue + # Writes fill missing columns from the write-default only, reads from the initial-default + default = field.write_default if self._for_write else field.initial_default if field_array is not None: array = self._cast_if_needed(field, field_array) field_arrays.append(array) fields.append(self._construct_field(field, array.type)) - elif field.optional or field.initial_default is not None: - # When an optional field is added, or when a required field with a non-null initial default is added - arrow_type = schema_to_pyarrow(field.field_type, include_field_ids=self._include_field_ids) + elif field.optional or default is not None: + # When an optional field is added, or when a required field with a non-null default is added + arrow_type = self._arrow_type(field.field_type) projected_value = self._projected_missing_fields.get(field.field_id) if projected_value is not None: field_arrays.append(pa.repeat(pa.scalar(projected_value, type=arrow_type), len(struct_array))) - elif field.initial_default is None: + elif default is None: field_arrays.append(pa.nulls(len(struct_array), type=arrow_type)) else: - field_arrays.append(pa.repeat(pa.scalar(field.initial_default, type=arrow_type), len(struct_array))) + field_arrays.append(pa.repeat(pa.scalar(default, type=arrow_type), len(struct_array))) fields.append(self._construct_field(field, arrow_type)) else: raise ResolveError(f"Field is required, and could not be found in the file: {field}") @@ -2219,6 +2696,10 @@ def visit_binary(self, binary_type: BinaryType) -> str: def visit_unknown(self, unknown_type: UnknownType) -> str: return "UNKNOWN" + def visit_variant(self, variant_type: VariantType) -> str: + # Unshredded metadata and value are both binary + return "BYTE_ARRAY" + def visit_geometry(self, geometry_type: GeometryType) -> str: return "BYTE_ARRAY" @@ -2304,6 +2785,81 @@ def max_as_bytes(self) -> bytes | None: return self.serialize(self.current_max) +class GeospatialStatsAggregator(StatsAggregator): + """Aggregates the Parquet geospatial statistics of a geometry or geography column into bounding-box bounds. + + Parquet min/max statistics on WKB values are lexicographic and meaningless as geo bounds, so the bounds are + built from the bounding box in `ColumnChunkMetaData.geo_statistics` instead. Any row group without a bounding + box invalidates the bounds of the whole file. + """ + + _BOX_KEYS = ("xmin", "ymin", "zmin", "mmin", "xmax", "ymax", "zmax", "mmax") + + def __init__(self, iceberg_type: PrimitiveType) -> None: + self.current_min = None + self.current_max = None + self.trunc_length = None + self.primitive_type = iceberg_type + self._box: dict[str, float | None] | None = None + self._valid = True + + def update_box(self, geo_statistics: Any | None) -> None: + """Merge the bounding box of a row group; a missing X/Y range invalidates the bounds.""" + box: dict[str, float | None] = {key: getattr(geo_statistics, key, None) for key in self._BOX_KEYS} + if any((value := box[key]) is None or not math.isfinite(value) for key in ("xmin", "ymin", "xmax", "ymax")): + self._valid = False + return + if self._box is None: + self._box = box + return + for key in self._BOX_KEYS: + current, new = self._box[key], box[key] + if current is None or new is None: + # Z and M are only kept when every row group has them + self._box[key] = None + else: + self._box[key] = min(current, new) if key.endswith("min") else max(current, new) + + def _bounds(self) -> tuple[bytes, bytes] | None: + if not self._valid or self._box is None: + return None + xmin, ymin, xmax, ymax = (self._box[key] for key in ("xmin", "ymin", "xmax", "ymax")) + if xmin is None or ymin is None or xmax is None or ymax is None: + return None + if isinstance(self.primitive_type, GeographyType) and xmin > xmax: + # The bounding box wraps the antimeridian; omitting the bounds is always spec-compliant + return None + return geo_bounds_from_bbox( + xmin, + ymin, + xmax, + ymax, + zmin=self._box["zmin"], + zmax=self._box["zmax"], + mmin=self._box["mmin"], + mmax=self._box["mmax"], + ) + + def min_as_bytes(self) -> bytes | None: + bounds = self._bounds() + return bounds[0] if bounds is not None else None + + def max_as_bytes(self) -> bytes | None: + bounds = self._bounds() + return bounds[1] if bounds is not None else None + + +_PARQUET_TIME_UNIT_TO_NANOS = {"milliseconds": 1_000_000, "microseconds": 1_000, "nanoseconds": 1} + + +def _parquet_timestamp_nanos_per_unit(statistics: pq.Statistics) -> int: + """Return the factor that converts raw Parquet timestamp statistics to nanoseconds.""" + time_unit = json.loads(statistics.logical_type.to_json()).get("timeUnit", "nanoseconds") + if time_unit not in _PARQUET_TIME_UNIT_TO_NANOS: + raise ValueError(f"Unsupported Parquet timestamp unit: {time_unit}") + return _PARQUET_TIME_UNIT_TO_NANOS[time_unit] + + DEFAULT_TRUNCATION_LENGTH = 16 TRUNCATION_EXPR = r"^truncate\((\d+)\)$" @@ -2423,6 +2979,10 @@ def primitive(self, primitive: PrimitiveType) -> builtins.list[StatisticsCollect if is_nested and metrics_mode.type in [MetricModeTypes.TRUNCATE, MetricModeTypes.FULL]: metrics_mode = MetricsMode(MetricModeTypes.COUNTS) + if isinstance(primitive, VariantType) and metrics_mode.type in [MetricModeTypes.TRUNCATE, MetricModeTypes.FULL]: + # Variant bounds are Variant objects keyed by field path, which are not collected + metrics_mode = MetricsMode(MetricModeTypes.COUNTS) + return [StatisticsCollector(field_id=self._field_id, iceberg_type=primitive, mode=metrics_mode, column_name=column_name)] @@ -2507,6 +3067,9 @@ def map( return k + v def primitive(self, primitive: PrimitiveType) -> builtins.list[ID2ParquetPath]: + if isinstance(primitive, VariantType): + # Every variant has exactly one metadata value per row, so its column carries the counts of the variant + return [ID2ParquetPath(field_id=self._field_id, parquet_path=".".join([*self._path, _VARIANT_METADATA]))] return [ID2ParquetPath(field_id=self._field_id, parquet_path=".".join(self._path))] @@ -2529,6 +3092,15 @@ def parquet_path_to_id_mapping( return result +def _is_variant_child_path(path: str, parquet_column_mapping: dict[str, int]) -> bool: + """Return whether a Parquet column belongs to a variant group without being its metadata column.""" + parts = path.split(".") + return any( + ".".join([*parts[:index], _VARIANT_METADATA]) in parquet_column_mapping and parts[index] != _VARIANT_METADATA + for index in range(1, len(parts)) + ) + + def data_file_statistics_from_parquet_metadata( parquet_metadata: pq.FileMetaData, stats_columns: dict[int, StatisticsCollector], @@ -2558,9 +3130,10 @@ def data_file_statistics_from_parquet_metadata( null_value_counts: dict[int, int] = {} nan_value_counts: dict[int, int] = {} - col_aggs = {} + col_aggs: dict[int, StatsAggregator] = {} invalidate_col: set[int] = set() + invalidate_null_count: set[int] = set() for r in range(parquet_metadata.num_row_groups): # References: # https://github.com/apache/iceberg/blob/fc381a81a1fdb8f51a0637ca27cd30673bd7aad3/parquet/src/main/java/org/apache/iceberg/parquet/ParquetUtil.java#L232 @@ -2578,6 +3151,10 @@ def data_file_statistics_from_parquet_metadata( for pos in range(parquet_metadata.num_columns): column = row_group.column(pos) + if column.path_in_schema not in parquet_column_mapping and _is_variant_child_path( + column.path_in_schema, parquet_column_mapping + ): + continue field_id = parquet_column_mapping[column.path_in_schema] stats_col = stats_columns[field_id] @@ -2590,6 +3167,19 @@ def data_file_statistics_from_parquet_metadata( value_counts[field_id] = value_counts.get(field_id, 0) + column.num_values + if isinstance(stats_col.iceberg_type, (GeometryType, GeographyType)): + if column.is_stats_set and column.statistics.has_null_count: + null_value_counts[field_id] = null_value_counts.get(field_id, 0) + column.statistics.null_count + else: + invalidate_null_count.add(field_id) + if stats_col.mode != MetricsMode(MetricModeTypes.COUNTS): + geo_agg = col_aggs.setdefault(field_id, GeospatialStatsAggregator(stats_col.iceberg_type)) + if isinstance(geo_agg, GeospatialStatsAggregator): + # Parquet GEOMETRY/GEOGRAPHY logical types carry a bounding box (pyarrow >= 21) + is_geo_stats_set = getattr(column, "is_geo_stats_set", False) + geo_agg.update_box(column.geo_statistics if is_geo_stats_set else None) + continue + if column.is_stats_set: try: statistics = column.statistics @@ -2620,6 +3210,13 @@ def data_file_statistics_from_parquet_metadata( if statistics.max_raw is not None else None ) + elif isinstance(stats_col.iceberg_type, (TimestampNanoType, TimestamptzNanoType)): + # statistics.min/max are datetimes truncated to microseconds, so use the raw int64 values + nanos_per_unit = _parquet_timestamp_nanos_per_unit(statistics) + if statistics.min_raw is not None: + col_aggs[field_id].update_min(statistics.min_raw * nanos_per_unit) + if statistics.max_raw is not None: + col_aggs[field_id].update_max(statistics.max_raw * nanos_per_unit) else: col_aggs[field_id].update_min(statistics.min) col_aggs[field_id].update_max(statistics.max) @@ -2637,6 +3234,9 @@ def data_file_statistics_from_parquet_metadata( col_aggs.pop(field_id, None) null_value_counts.pop(field_id, None) + for field_id in invalidate_null_count: + null_value_counts.pop(field_id, None) + return DataFileStatistics( record_count=parquet_metadata.num_rows, column_sizes=column_sizes, @@ -2721,7 +3321,21 @@ def add_field_metadata(self, field: NestedField, metadata: dict[bytes, bytes], i FileFormatFactory.register(ParquetFormatModel()) -def write_file(io: FileIO, table_metadata: TableMetadata, tasks: Iterator[WriteTask]) -> Iterator[DataFile]: +def _row_lineage_write_schema(schema: Schema) -> Schema: + """Return the schema with the row lineage columns appended, for copying rows with their lineage.""" + if (lineage_schema := schema_with_row_lineage(schema)) is None: + raise ValueError("Cannot preserve row lineage: a table column uses a reserved row lineage column name") + return lineage_schema + + +def write_file( + io: FileIO, table_metadata: TableMetadata, tasks: Iterator[WriteTask], preserve_row_lineage: bool = False +) -> Iterator[DataFile]: + """Write the tasks to data files. + + With ``preserve_row_lineage``, the ``_row_id`` and ``_last_updated_sequence_number`` columns of the tasks are + written as physical columns with their reserved field IDs, so rows copied by a rewrite keep their lineage. + """ from pyiceberg.table import DOWNCAST_NS_TIMESTAMP_TO_US_ON_WRITE, TableProperties file_format = FileFormat( @@ -2739,6 +3353,8 @@ def write_data_file(task: WriteTask) -> DataFile: file_schema = sanitized_schema else: file_schema = table_schema + if preserve_row_lineage: + file_schema = _row_lineage_write_schema(file_schema) downcast_ns_timestamp_to_us = Config().get_bool(DOWNCAST_NS_TIMESTAMP_TO_US_ON_WRITE) or False batches = [ @@ -2749,6 +3365,7 @@ def write_data_file(task: WriteTask) -> DataFile: downcast_ns_timestamp_to_us=downcast_ns_timestamp_to_us, include_field_ids=True, format_model=format_model, + for_write=True, ) for batch in task.record_batches ] @@ -2862,7 +3479,11 @@ def _check_pyarrow_schema_compatible( Raises: ValueError: If the schemas are not compatible. + NotImplementedError: If the schema contains a variant column. """ + if any(isinstance(field.field_type, VariantType) for field in requested_schema._lazy_id_to_field.values()): + raise NotImplementedError("Writing variant is not supported until pyarrow can annotate the Parquet VARIANT logical type") + name_mapping = requested_schema.name_mapping try: @@ -2972,9 +3593,13 @@ def _dataframe_to_data_files( io: FileIO, write_uuid: uuid.UUID | None = None, counter: itertools.count[int] | None = None, + preserve_row_lineage: bool = False, ) -> Iterable[DataFile]: """Convert a PyArrow Table or RecordBatchReader into DataFiles. + With ``preserve_row_lineage``, ``df`` also holds the ``_row_id`` and ``_last_updated_sequence_number`` + columns of copied rows, and they are written into the data files (see :func:`write_file`). + For a ``pa.Table`` the data is materialised in memory and bin-packed into target-sized files (with partition splitting if the table is partitioned). @@ -2996,7 +3621,8 @@ def _dataframe_to_data_files( property_name=TableProperties.WRITE_TARGET_FILE_SIZE_BYTES, default=TableProperties.WRITE_TARGET_FILE_SIZE_BYTES_DEFAULT, ) - name_mapping = table_metadata.schema().name_mapping + table_schema = table_metadata.schema() + name_mapping = (_row_lineage_write_schema(table_schema) if preserve_row_lineage else table_schema).name_mapping downcast_ns_timestamp_to_us = Config().get_bool(DOWNCAST_NS_TIMESTAMP_TO_US_ON_WRITE) or False task_schema = pyarrow_to_schema( df.schema, @@ -3015,6 +3641,7 @@ def _dataframe_to_data_files( yield from write_file( io=io, table_metadata=table_metadata, + preserve_row_lineage=preserve_row_lineage, tasks=( WriteTask(write_uuid=write_uuid, task_id=next(counter), record_batches=batches, schema=task_schema) for batches in bin_pack_record_batches(df, target_file_size) @@ -3026,6 +3653,7 @@ def _dataframe_to_data_files( yield from write_file( io=io, table_metadata=table_metadata, + preserve_row_lineage=preserve_row_lineage, tasks=( WriteTask(write_uuid=write_uuid, task_id=next(counter), record_batches=batches, schema=task_schema) for batches in bin_pack_arrow_table(df, target_file_size) @@ -3036,6 +3664,7 @@ def _dataframe_to_data_files( yield from write_file( io=io, table_metadata=table_metadata, + preserve_row_lineage=preserve_row_lineage, tasks=( WriteTask( write_uuid=write_uuid, diff --git a/pyiceberg/manifest.py b/pyiceberg/manifest.py index 88ca051015..193b2d21c3 100644 --- a/pyiceberg/manifest.py +++ b/pyiceberg/manifest.py @@ -19,13 +19,14 @@ import math import threading from abc import ABC, abstractmethod -from collections.abc import Callable, Iterator +from collections.abc import Callable, Iterator, Mapping from copy import copy from enum import Enum from types import TracebackType from typing import ( Any, Literal, + cast, ) from cachetools import LRUCache @@ -548,6 +549,10 @@ def sort_order_id(self) -> int | None: def first_row_id(self) -> int | None: return self._data[16] + @first_row_id.setter + def first_row_id(self, value: int | None) -> None: + self._data[16] = value + @property def referenced_data_file(self) -> str | None: return self._data[17] @@ -885,6 +890,10 @@ def key_metadata(self) -> bytes | None: def first_row_id(self) -> int | None: return self._data[15] + @first_row_id.setter + def first_row_id(self, value: int | None) -> None: + self._data[15] = value + def has_added_files(self) -> bool: return self.added_files_count is None or self.added_files_count > 0 @@ -916,8 +925,17 @@ def fetch_manifest_entry( read_enums={0: ManifestEntryStatus, 101: FileFormat, 134: DataFileContent}, ) as reader: result = [] + # First row ids are only inherited from data manifests that have one assigned (V3) + next_row_id = self.first_row_id if self.content == ManifestContent.DATA else None for entry in reader: + data_file = entry.data_file + if next_row_id is None: + data_file.first_row_id = None + elif entry.status != ManifestEntryStatus.DELETED and data_file.first_row_id is None: + data_file.first_row_id = next_row_id + next_row_id += data_file.record_count + if discard_deleted and entry.status == ManifestEntryStatus.DELETED: continue _inherit_from_manifest(entry, self) @@ -975,8 +993,10 @@ def get_or_cache(self, manifest_file: ManifestFile) -> ManifestFile: with self._lock: manifest_path = manifest_file.manifest_path - if manifest_path in self._cache: - return self._cache[manifest_path] + cached = self._cache.get(manifest_path) + # A manifest written before a V3 upgrade has no first row id until a later manifest list assigns one + if cached is not None and cached.first_row_id == manifest_file.first_row_id: + return cached self._cache[manifest_path] = manifest_file return manifest_file @@ -1297,6 +1317,8 @@ def prepare_entry(self, entry: ManifestEntry) -> ManifestEntry: class ManifestWriterV2(ManifestWriter): + _content: ManifestContent + def __init__( self, spec: PartitionSpec, @@ -1304,11 +1326,13 @@ def __init__( output_file: OutputFile, snapshot_id: int, avro_compression: AvroCompressionCodec, + content: ManifestContent = ManifestContent.DATA, ): super().__init__(spec, schema, output_file, snapshot_id, avro_compression) + self._content = content def content(self) -> ManifestContent: - return ManifestContent.DATA + return self._content @property def version(self) -> TableVersion: @@ -1318,7 +1342,7 @@ def version(self) -> TableVersion: def _meta(self) -> dict[str, str]: return { **super()._meta, - "content": "data", + "content": "data" if self._content == ManifestContent.DATA else "deletes", } def prepare_entry(self, entry: ManifestEntry) -> ManifestEntry: @@ -1330,6 +1354,58 @@ def prepare_entry(self, entry: ManifestEntry) -> ManifestEntry: return entry +def _layout_version_from_field_count(layouts: Mapping[int, Schema | StructType], field_count: int) -> TableVersion: + """Return the format version whose layout has exactly `field_count` fields. + + A positional record does not carry the format version it was bound to, but each version's + layout has a distinct number of fields, so the field count identifies it. + """ + matches = [version for version, layout in layouts.items() if len(layout.fields) == field_count] + if len(matches) != 1: + raise ValueError(f"Cannot determine layout version for record with {field_count} fields") + return cast(TableVersion, matches[0]) + + +def _rebind_data_file(data_file: DataFile, format_version: TableVersion) -> DataFile: + """Rebind a data file to the layout of the given format version, keeping its values and spec id.""" + target = DATA_FILE_TYPE[format_version] + if len(data_file._data) == len(target.fields): + return data_file + source = DATA_FILE_TYPE[_layout_version_from_field_count(DATA_FILE_TYPE, len(data_file._data))] + target_names = {field.name for field in target.fields} + args = {field.name: value for field, value in zip(source.fields, data_file._data, strict=True) if field.name in target_names} + rebound = DataFile.from_args(_table_format_version=format_version, **args) + if hasattr(data_file, "_spec_id"): + rebound.spec_id = data_file.spec_id + return rebound + + +class ManifestWriterV3(ManifestWriterV2): + """Writes V3 manifest files. + + Sequence number semantics are the same as V2. The V3 data file struct additionally carries + `first_row_id`, `referenced_data_file`, `content_offset` and `content_size_in_bytes`. The + `first_row_id` of added data files is left unset so that readers inherit it from the manifest. + """ + + @property + def version(self) -> TableVersion: + return 3 + + def prepare_entry(self, entry: ManifestEntry) -> ManifestEntry: + entry = super().prepare_entry(entry) + if len(entry.data_file._data) != len(DATA_FILE_TYPE[3].fields): + entry = ManifestEntry.from_args( + _table_format_version=3, + status=entry.status, + snapshot_id=entry.snapshot_id, + sequence_number=entry.sequence_number, + file_sequence_number=entry.file_sequence_number, + data_file=_rebind_data_file(entry.data_file, 3), + ) + return entry + + def write_manifest( format_version: TableVersion, spec: PartitionSpec, @@ -1337,11 +1413,16 @@ def write_manifest( output_file: OutputFile, snapshot_id: int, avro_compression: AvroCompressionCodec, + content: ManifestContent = ManifestContent.DATA, ) -> ManifestWriter: if format_version == 1: + if content != ManifestContent.DATA: + raise ValidationError("Cannot write delete manifests in a v1 table") return ManifestWriterV1(spec, schema, output_file, snapshot_id, avro_compression) elif format_version == 2: - return ManifestWriterV2(spec, schema, output_file, snapshot_id, avro_compression) + return ManifestWriterV2(spec, schema, output_file, snapshot_id, avro_compression, content) + elif format_version == 3: + return ManifestWriterV3(spec, schema, output_file, snapshot_id, avro_compression, content) else: raise ValueError(f"Cannot write manifest for table version: {format_version}") @@ -1465,6 +1546,59 @@ def prepare_manifest(self, manifest_file: ManifestFile) -> ManifestFile: return wrapped_manifest_file +class ManifestListWriterV3(ManifestListWriterV2): + """Writes V3 manifest lists and assigns `first_row_id` to data manifests. + + Follows the spec's first row id assignment: data manifests without a `first_row_id` get a + running value that starts at the snapshot's `first-row-id` and advances by the manifest's + existing and added row counts. Delete manifests are never assigned one, and already assigned + values are preserved. + """ + + _next_row_id: int + + def __init__( + self, + output_file: OutputFile, + snapshot_id: int, + parent_snapshot_id: int | None, + sequence_number: int, + compression: AvroCompressionCodec, + first_row_id: int, + ): + super().__init__(output_file, snapshot_id, parent_snapshot_id, sequence_number, compression) + self._format_version = 3 + self._meta = {**self._meta, "format-version": "3", "first-row-id": str(first_row_id)} + self._next_row_id = first_row_id + + @property + def next_row_id(self) -> int: + """The row id after the last assigned one, which becomes the table's `next-row-id`.""" + return self._next_row_id + + def prepare_manifest(self, manifest_file: ManifestFile) -> ManifestFile: + # Always work on a fresh record: manifests of the parent snapshot are shared through the manifest + # cache, so assigning a first row id must not leak back into them + source_version = _layout_version_from_field_count(MANIFEST_LIST_FILE_SCHEMAS, len(manifest_file._data)) + source = MANIFEST_LIST_FILE_SCHEMAS[source_version] + wrapped_manifest_file = super().prepare_manifest( + ManifestFile.from_args( + _table_format_version=3, + **{field.name: value for field, value in zip(source.fields, manifest_file._data, strict=True)}, + ) + ) + + if wrapped_manifest_file.content == ManifestContent.DATA and wrapped_manifest_file.first_row_id is None: + if wrapped_manifest_file.existing_rows_count is None or wrapped_manifest_file.added_rows_count is None: + raise ValueError( + f"Cannot assign first-row-id to manifest with unknown row counts: {wrapped_manifest_file.manifest_path}" + ) + wrapped_manifest_file.first_row_id = self._next_row_id + self._next_row_id += wrapped_manifest_file.existing_rows_count + wrapped_manifest_file.added_rows_count + + return wrapped_manifest_file + + def write_manifest_list( format_version: TableVersion, output_file: OutputFile, @@ -1472,6 +1606,7 @@ def write_manifest_list( parent_snapshot_id: int | None, sequence_number: int | None, avro_compression: AvroCompressionCodec, + first_row_id: int | None = None, ) -> ManifestListWriter: if format_version == 1: return ManifestListWriterV1(output_file, snapshot_id, parent_snapshot_id, avro_compression) @@ -1479,5 +1614,11 @@ def write_manifest_list( if sequence_number is None: raise ValueError(f"Sequence-number is required for V2 tables: {sequence_number}") return ManifestListWriterV2(output_file, snapshot_id, parent_snapshot_id, sequence_number, avro_compression) + elif format_version == 3: + if sequence_number is None: + raise ValueError(f"Sequence-number is required for V3 tables: {sequence_number}") + if first_row_id is None: + raise ValueError("First-row-id is required for V3 tables") + return ManifestListWriterV3(output_file, snapshot_id, parent_snapshot_id, sequence_number, avro_compression, first_row_id) else: raise ValueError(f"Cannot write manifest list for table version: {format_version}") diff --git a/pyiceberg/partitioning.py b/pyiceberg/partitioning.py index 3d06287b6c..5f064215e5 100644 --- a/pyiceberg/partitioning.py +++ b/pyiceberg/partitioning.py @@ -29,7 +29,6 @@ Field, PlainSerializer, WithJsonSchema, - model_validator, ) from pyiceberg.exceptions import ValidationError @@ -41,6 +40,7 @@ IdentityTransform, MonthTransform, Transform, + TransformSourceMixin, TruncateTransform, UnknownTransform, VoidTransform, @@ -68,17 +68,17 @@ PARTITION_FIELD_ID_START: int = 1000 -class PartitionField(IcebergBaseModel): +class PartitionField(TransformSourceMixin): """PartitionField represents how one partition value is derived from the source column via transformation. Attributes: - source_id(int): The source column id of table's schema. field_id(int): The partition field id across all the table partition specs. transform(Transform): The transform used to produce partition values from source column. name(str): The name of this partition field. + + The source columns are carried by `TransformSourceMixin`. """ - source_id: int = Field(alias="source-id") field_id: int = Field(alias="field-id") transform: Annotated[ # type: ignore Transform, @@ -107,25 +107,10 @@ def __init__( super().__init__(**data) - @model_validator(mode="before") - @classmethod - def map_source_ids_onto_source_id(cls, data: Any) -> Any: - if isinstance(data, dict): - if "source-ids" in data: - if "source-id" in data: - raise ValueError("source-id and source-ids are mutually exclusive") - source_ids = data["source-ids"] - if isinstance(source_ids, list): - if len(source_ids) == 0: - raise ValueError("Empty source-ids is not allowed") - if len(source_ids) > 1: - raise ValueError("Multi argument transforms are not yet supported") - data["source-id"] = source_ids[0] - return data - def __str__(self) -> str: """Return the string representation of the PartitionField class.""" - return f"{self.field_id}: {self.name}: {self.transform}({self.source_id})" + sources = ", ".join(str(source_id) for source_id in self.transform_arguments) + return f"{self.field_id}: {self.name}: {self.transform}({sources})" class PartitionSpec(IcebergBaseModel): @@ -206,7 +191,7 @@ def compatible_with(self, other: PartitionSpec) -> bool: if len(self.fields) != len(other.fields): return False return all( - this_field.source_id == that_field.source_id + this_field.transform_arguments == that_field.transform_arguments and this_field.transform == that_field.transform and this_field.name == that_field.name for this_field, that_field in zip(self.fields, other.fields, strict=True) diff --git a/pyiceberg/schema.py b/pyiceberg/schema.py index 99f983074b..1ee7704cb7 100644 --- a/pyiceberg/schema.py +++ b/pyiceberg/schema.py @@ -61,6 +61,7 @@ TimeType, UnknownType, UUIDType, + VariantType, ) if TYPE_CHECKING: @@ -558,6 +559,8 @@ def primitive(self, primitive: PrimitiveType, primitive_partner: P | None) -> T: return self.visit_geometry(primitive, primitive_partner) elif isinstance(primitive, GeographyType): return self.visit_geography(primitive, primitive_partner) + elif isinstance(primitive, VariantType): + return self.visit_variant(primitive, primitive_partner) else: raise ValueError(f"Type not recognized: {primitive}") @@ -637,6 +640,10 @@ def visit_geometry(self, geometry_type: GeometryType, partner: P | None) -> T: def visit_geography(self, geography_type: GeographyType, partner: P | None) -> T: """Visit a GeographyType.""" + @abstractmethod + def visit_variant(self, variant_type: VariantType, partner: P | None) -> T: + """Visit a VariantType.""" + class PartnerAccessor(Generic[P], ABC): @abstractmethod @@ -764,6 +771,8 @@ def primitive(self, primitive: PrimitiveType) -> T: return self.visit_geometry(primitive) elif isinstance(primitive, GeographyType): return self.visit_geography(primitive) + elif isinstance(primitive, VariantType): + return self.visit_variant(primitive) else: raise ValueError(f"Type not recognized: {primitive}") @@ -843,6 +852,10 @@ def visit_geometry(self, geometry_type: GeometryType) -> T: def visit_geography(self, geography_type: GeographyType) -> T: """Visit a GeographyType.""" + @abstractmethod + def visit_variant(self, variant_type: VariantType) -> T: + """Visit a VariantType.""" + @dataclass(init=True, eq=True, frozen=True) class Accessor: @@ -1347,6 +1360,8 @@ def struct(self, struct: StructType, field_results: builtins.list[Callable[[], I field_type=field_type(), required=field.required, doc=field.doc, + initial_default=field.initial_default, + write_default=field.write_default, ) ) return StructType(*new_fields) @@ -1465,6 +1480,8 @@ def field(self, field: NestedField, field_result: IcebergType | None) -> Iceberg field_type=field_result, doc=field.doc, required=field.required, + initial_default=field.initial_default, + write_default=field.write_default, ) def struct(self, struct: StructType, field_results: builtins.list[IcebergType | None]) -> IcebergType | None: @@ -1687,10 +1704,23 @@ def _(file_type: StringType, read_type: IcebergType) -> IcebergType: def _(file_type: BinaryType, read_type: IcebergType) -> IcebergType: if isinstance(read_type, StringType): return read_type + elif isinstance(read_type, (GeometryType, GeographyType)): + # Like fixed -> uuid, this is a read-time resolution rule rather than a schema evolution: geometry and + # geography values are stored as WKB, and PyArrow reports Parquet GEOMETRY/GEOGRAPHY columns as binary + return read_type else: raise ResolveError(f"Cannot promote an binary to {read_type}") +@promote.register(DateType) +def _(file_type: DateType, read_type: IcebergType) -> IcebergType: + if isinstance(read_type, (TimestampType, TimestampNanoType)): + # Format version 3 allows date -> timestamp and date -> timestamp_ns, reading a date as midnight + return read_type + else: + raise ResolveError(f"Cannot promote a date to {read_type}") + + @promote.register(DecimalType) def _(file_type: DecimalType, read_type: IcebergType) -> IcebergType: if isinstance(read_type, DecimalType): @@ -1711,6 +1741,23 @@ def _(file_type: FixedType, read_type: IcebergType) -> IcebergType: raise ResolveError(f"Cannot promote {file_type} to {read_type}") +@promote.register(StructType) +def _(file_type: StructType, read_type: IcebergType) -> IcebergType: + if isinstance(read_type, VariantType) and _is_variant_struct(file_type): + # A read-time resolution rule rather than a schema evolution: variant values are stored as a group + # with binary metadata and value fields, which surfaces as a struct when a file assigns them field ids + return read_type + else: + raise ResolveError(f"Cannot promote {file_type} to {read_type}") + + +def _is_variant_struct(struct: StructType) -> bool: + """Return whether a struct has exactly the shape of an unshredded variant.""" + return sorted(field.name for field in struct.fields) == ["metadata", "value"] and all( + isinstance(field.field_type, BinaryType) for field in struct.fields + ) + + @promote.register(UnknownType) def _(file_type: UnknownType, read_type: IcebergType) -> IcebergType: # Per V3 Spec, "Unknown" can be promoted to any Primitive type @@ -1758,7 +1805,7 @@ def _is_field_compatible(self, lhs: NestedField) -> bool: try: rhs = self.provided_schema.find_field(lhs.field_id) except ValueError: - if lhs.required: + if lhs.required and lhs.write_default is None: self.rich_table.add_row("❌", str(lhs), "Missing") return False else: diff --git a/pyiceberg/table/__init__.py b/pyiceberg/table/__init__.py index 2c5c26800c..884b70d590 100644 --- a/pyiceberg/table/__init__.py +++ b/pyiceberg/table/__init__.py @@ -40,18 +40,30 @@ _InclusiveMetricsEvaluator, bind, expression_evaluator, + extract_field_ids, inclusive_projection, manifest_evaluator, ) from pyiceberg.io import FileIO, load_file_io -from pyiceberg.manifest import DataFile, DataFileContent, ManifestContent, ManifestEntry, ManifestEntryStatus, ManifestFile +from pyiceberg.manifest import ( + DataFile, + DataFileContent, + FileFormat, + ManifestContent, + ManifestEntry, + ManifestEntryStatus, + ManifestFile, +) from pyiceberg.partitioning import PARTITION_FIELD_ID_START, UNPARTITIONED_PARTITION_SPEC, PartitionKey, PartitionSpec -from pyiceberg.schema import Schema +from pyiceberg.schema import Schema, prune_columns +from pyiceberg.table.delete_file import DeleteFileSet from pyiceberg.table.delete_file_index import DeleteFileIndex +from pyiceberg.table.deletion_vector import DeletionVector from pyiceberg.table.inspect import InspectTable from pyiceberg.table.locations import LocationProvider, load_location_provider from pyiceberg.table.maintenance import MaintenanceTable -from pyiceberg.table.metadata import INITIAL_SEQUENCE_NUMBER, TableMetadata +from pyiceberg.table.metadata import INITIAL_SEQUENCE_NUMBER, SUPPORTED_TABLE_FORMAT_VERSION, TableMetadata +from pyiceberg.table.metadata_columns import is_metadata_column, project_schema, schema_with_row_lineage from pyiceberg.table.name_mapping import NameMapping from pyiceberg.table.refs import MAIN_BRANCH, SnapshotRef from pyiceberg.table.snapshots import ( @@ -111,6 +123,7 @@ import pandas as pd import polars as pl import pyarrow as pa + import pyarrow.compute as pc import ray from duckdb import DuckDBPyConnection from pyiceberg_core.datafusion import IcebergDataFusionTable @@ -332,7 +345,7 @@ def upgrade_table_version(self, format_version: TableVersion) -> Transaction: Returns: The alter table builder. """ - if format_version not in {1, 2}: + if not 1 <= format_version <= SUPPORTED_TABLE_FORMAT_VERSION: raise ValueError(f"Unsupported table format version: {format_version}") if format_version < self.table_metadata.format_version: @@ -565,7 +578,7 @@ def append( table_metadata=self.table_metadata, write_uuid=append_files.commit_uuid, df=df, io=self._table.io ) for data_file in data_files: - append_files.append_data_file(data_file) + append_files._append_written_data_file(data_file) def dynamic_partition_overwrite( self, df: pa.Table, snapshot_properties: dict[str, str] = EMPTY_DICT, branch: str | None = MAIN_BRANCH @@ -638,7 +651,7 @@ def dynamic_partition_overwrite( with self._append_snapshot_producer(snapshot_properties, branch=branch) as append_files: append_files.commit_uuid = append_snapshot_commit_uuid for data_file in data_files: - append_files.append_data_file(data_file) + append_files._append_written_data_file(data_file) def overwrite( self, @@ -736,7 +749,7 @@ def overwrite( table_metadata=self.table_metadata, write_uuid=append_files.commit_uuid, df=df, io=self._table.io ) for data_file in data_files: - append_files.append_data_file(data_file) + append_files._append_written_data_file(data_file) def delete( self, @@ -753,6 +766,8 @@ def delete( - DELETE: In case existing Parquet files can be dropped completely. - OVERWRITE: In case existing Parquet files need to be rewritten to drop rows that match the delete filter. + - DELETE: In case of a format version 3 table with `write.delete.mode=merge-on-read`, rows in + partially matching files are marked as deleted with deletion vectors instead of rewriting the files. Args: delete_filter: A boolean expression to delete rows from a table @@ -762,11 +777,19 @@ def delete( """ from pyiceberg.io.pyarrow import ArrowScan, _dataframe_to_data_files, _expression_to_complementary_pyarrow + use_deletion_vectors = False if ( self.table_metadata.properties.get(TableProperties.DELETE_MODE, TableProperties.DELETE_MODE_DEFAULT) == TableProperties.DELETE_MODE_MERGE_ON_READ ): - warnings.warn("Merge on read is not yet supported, falling back to copy-on-write", stacklevel=2) + if self.table_metadata.format_version < 3: + warnings.warn( + "Merge-on-read deletes require format version 3 (deletion vectors); falling back to copy-on-write", + stacklevel=2, + ) + else: + # The delete phase of an overwrite (which passes an isolation operation) keeps copy-on-write + use_deletion_vectors = _isolation_operation is None if isinstance(delete_filter, str): delete_filter = _parse_row_filter(delete_filter) @@ -777,7 +800,15 @@ def delete( delete_snapshot.delete_by_predicate(delete_filter, case_sensitive) # Check if there are any files that require an actual rewrite of a data file - if delete_snapshot.rewrites_needed is True: + if delete_snapshot.rewrites_needed is True and use_deletion_vectors: + self._delete_with_deletion_vectors( + delete_filter=delete_filter, + case_sensitive=case_sensitive, + snapshot_properties=snapshot_properties, + branch=branch, + starting_snapshot_id=delete_snapshot._starting_snapshot_id, + ) + elif delete_snapshot.rewrites_needed is True: bound_delete_filter = bind(self.table_metadata.schema(), delete_filter, case_sensitive) preserve_row_filter = _expression_to_complementary_pyarrow(bound_delete_filter, self.table_metadata.schema()) @@ -789,6 +820,15 @@ def delete( commit_uuid = uuid.uuid4() counter = itertools.count(0) + # Rows copied into the rewritten files keep their row IDs and last updated sequence numbers, + # which are written as physical columns (see https://iceberg.apache.org/spec/#row-lineage) + projected_schema = self.table_metadata.schema() + preserve_row_lineage = False + if self.table_metadata.format_version >= 3: + if (lineage_schema := schema_with_row_lineage(projected_schema)) is not None: + projected_schema = lineage_schema + preserve_row_lineage = True + replaced_files: list[tuple[DataFile, list[DataFile]]] = [] # This will load the Parquet file into memory, including: # - Filter out the rows based on the delete filter @@ -801,7 +841,7 @@ def delete( df = ArrowScan( table_metadata=self.table_metadata, io=self._table.io, - projected_schema=self.table_metadata.schema(), + projected_schema=projected_schema, row_filter=AlwaysTrue(), ).to_table(tasks=[original_file]) filtered_df = df.filter(preserve_row_filter) @@ -820,6 +860,7 @@ def delete( table_metadata=self.table_metadata, write_uuid=commit_uuid, counter=counter, + preserve_row_lineage=preserve_row_lineage, ) ), ) @@ -837,11 +878,71 @@ def delete( for original_data_file, replaced_data_files in replaced_files: overwrite_snapshot.delete_data_file(original_data_file) for replaced_data_file in replaced_data_files: - overwrite_snapshot.append_data_file(replaced_data_file) + overwrite_snapshot._append_written_data_file(replaced_data_file) if not delete_snapshot.files_affected and not delete_snapshot.rewrites_needed: warnings.warn("Delete operation did not match any records", stacklevel=2) + def _delete_with_deletion_vectors( + self, + delete_filter: BooleanExpression, + case_sensitive: bool, + snapshot_properties: dict[str, str], + branch: str | None, + starting_snapshot_id: int | None, + ) -> None: + """Mark the rows matching the filter as deleted with deletion vectors (merge-on-read). + + Each partially matching data file gets one deletion vector holding the matching positions + and the positions deleted before, replacing the file's previous deletion vector and + file-scoped position delete files. + """ + from pyiceberg.io.pyarrow import _expression_to_complementary_pyarrow + from pyiceberg.table.deletion_vector import write_deletion_vectors + + schema = self.table_metadata.schema() + bound_delete_filter = bind(schema, delete_filter, case_sensitive) + preserve_row_filter = _expression_to_complementary_pyarrow(bound_delete_filter, schema) + filter_field_ids = extract_field_ids(bound_delete_filter) + + file_scan = self._scan(row_filter=delete_filter, case_sensitive=case_sensitive) + if branch is not None: + file_scan = file_scan.use_ref(branch) + + def _plan(task: FileScanTask) -> tuple[DataFile, DeletionVector, list[DataFile]] | None: + positions = _deleted_positions(self.table_metadata, self._table.io, task, filter_field_ids, preserve_row_filter) + if len(positions) == 0: + return None + deletion_vector = DeletionVector.from_positions(task.file.file_path, positions) + previous, replaced = _previous_deletes(self._table.io, task) + if previous is not None: + merged = previous.union(deletion_vector) + if merged.cardinality == previous.cardinality: + # All matching rows were deleted before + return None + deletion_vector = merged + return task.file, deletion_vector, replaced + + executor = ExecutorFactory.get_or_create() + planned = [result for result in executor.map(_plan, file_scan.plan_files()) if result is not None] + if not planned: + return + + commit_uuid = uuid.uuid4() + location = self._table.location_provider().new_data_location(f"00000-0-{commit_uuid}-00001-deletes.puffin") + with self.update_snapshot(snapshot_properties=snapshot_properties, branch=branch).row_delta(commit_uuid) as row_delta: + row_delta._starting_snapshot_id = starting_snapshot_id + row_delta.delete_by_predicate(delete_filter, case_sensitive) + # Track the location before writing, so a partially written file is removed on failure + row_delta._written_delete_file_paths.add(location) + delete_files = write_deletion_vectors(self._table.io, location, [(data_file, dv) for data_file, dv, _ in planned]) + for delete_file in delete_files: + row_delta._add_written_delete_file(delete_file) + for _, _, replaced in planned: + for replaced_delete_file in replaced: + row_delta.remove_delete_file(replaced_delete_file) + row_delta.validate_data_files_exist({data_file for data_file, _, _ in planned}) + def upsert( self, df: pa.Table, @@ -1286,6 +1387,66 @@ def commit_transaction(self) -> Table: return self._table +def _deleted_positions( + table_metadata: TableMetadata, + io: FileIO, + task: FileScanTask, + filter_field_ids: set[int], + preserve_row_filter: pc.Expression, +) -> list[int]: + """Return the positions of the rows in the task's data file that the preserve filter does not keep. + + The rows are read without applying existing deletes, so the positions are relative to the file. + """ + import pyarrow.compute as pc + + from pyiceberg.io.pyarrow import _ROW_POSITION_COLUMN, ArrowScan, _row_positions + + projected_schema = prune_columns(table_metadata.schema(), filter_field_ids, select_full_types=False) + df = ArrowScan(table_metadata=table_metadata, io=io, projected_schema=projected_schema, row_filter=AlwaysTrue()).to_table( + tasks=[FileScanTask(task.file, sequence_number=task.sequence_number)] + ) + if df.num_rows != task.file.record_count: + raise ValueError(f"Expected {task.file.record_count} rows in {task.file.file_path}, but read {df.num_rows}") + + positions = _row_positions(0, df.num_rows) + kept = df.append_column(_ROW_POSITION_COLUMN, positions).filter(preserve_row_filter).column(_ROW_POSITION_COLUMN) + return pc.filter(positions, pc.invert(pc.is_in(positions, value_set=kept))).to_pylist() + + +def _previous_deletes(io: FileIO, task: FileScanTask) -> tuple[DeletionVector | None, list[DataFile]]: + """Return the positions deleted by the task's position deletes, and the delete files a new deletion vector replaces. + + A new deletion vector replaces the data file's previous deletion vector and its file-scoped position delete + files, as in Java. Partition-scoped position delete files are kept; their positions are merged in. + """ + from pyiceberg.io.pyarrow import _read_deletes + from pyiceberg.table.delete_file_index import _referenced_data_file_path + from pyiceberg.table.deletion_vector import read_deletion_vectors + + data_file_path = task.file.file_path + previous: DeletionVector | None = None + replaced: list[DataFile] = [] + for delete_file in task.delete_files: + if delete_file.content != DataFileContent.POSITION_DELETES: + continue + if delete_file.file_format == FileFormat.PUFFIN: + vectors = [dv for dv in read_deletion_vectors(io, delete_file) if dv.referenced_data_file == data_file_path] + replaced.append(delete_file) + else: + positions = _read_deletes(io, delete_file).get(data_file_path) + vectors = ( + [DeletionVector.from_positions(data_file_path, positions.to_pylist())] + if positions is not None and len(positions) > 0 + else [] + ) + if data_file_path in (delete_file.referenced_data_file, _referenced_data_file_path(delete_file)): + replaced.append(delete_file) + for vector in vectors: + previous = vector if previous is None else previous.union(vector) + return previous, replaced + + class Namespace(IcebergRootModel[list[str]]): """Reference to one or more levels of a namespace.""" @@ -2211,10 +2372,7 @@ def projection(self) -> Schema: else: raise ValueError(f"Snapshot not found: {self.snapshot_id}") - if "*" in self.selected_fields: - return current_schema - - return current_schema.select(*self.selected_fields, case_sensitive=self.case_sensitive) + return project_schema(current_schema, self.selected_fields, case_sensitive=self.case_sensitive) def use_ref(self: S, name: str) -> S: if self.snapshot_id is not None: @@ -2237,18 +2395,22 @@ class FileScanTask(ScanTask): """Task representing a data file and its corresponding delete files.""" file: DataFile - delete_files: set[DataFile] + delete_files: DeleteFileSet residual: BooleanExpression + sequence_number: int | None def __init__( self, data_file: DataFile, - delete_files: set[DataFile] | None = None, + delete_files: Iterable[DataFile] | None = None, residual: BooleanExpression = ALWAYS_TRUE, + sequence_number: int | None = None, ) -> None: self.file = data_file - self.delete_files = delete_files or set() + self.delete_files = DeleteFileSet(delete_files if delete_files is not None else []) self.residual = residual + # The data sequence number of the manifest entry, used to derive _last_updated_sequence_number + self.sequence_number = sequence_number @staticmethod def from_rest_response( @@ -2263,21 +2425,13 @@ def from_rest_response( Returns: A FileScanTask with the converted data and delete files. - - Raises: - NotImplementedError: If equality delete files are encountered. """ - from pyiceberg.catalog.rest.scan_planning import RESTEqualityDeleteFile - data_file = _rest_file_to_data_file(rest_task.data_file) - resolved_deletes: set[DataFile] = set() + resolved_deletes = DeleteFileSet() if rest_task.delete_file_references: for idx in rest_task.delete_file_references: - delete_file = delete_files[idx] - if isinstance(delete_file, RESTEqualityDeleteFile): - raise NotImplementedError(f"PyIceberg does not yet support equality deletes: {delete_file.file_path}") - resolved_deletes.add(_rest_file_to_data_file(delete_file)) + resolved_deletes.add(_rest_file_to_data_file(delete_files[idx])) return FileScanTask( data_file=data_file, @@ -2288,7 +2442,7 @@ def from_rest_response( def _rest_file_to_data_file(rest_file: RESTContentFile) -> DataFile: """Convert a REST content file to a manifest DataFile.""" - from pyiceberg.catalog.rest.scan_planning import RESTDataFile + from pyiceberg.catalog.rest.scan_planning import RESTDataFile, RESTEqualityDeleteFile, RESTPositionDeleteFile if isinstance(rest_file, RESTDataFile): column_sizes = rest_file.column_sizes.to_dict() if rest_file.column_sizes else None @@ -2314,6 +2468,11 @@ def _rest_file_to_data_file(rest_file: RESTContentFile) -> DataFile: nan_value_counts=nan_value_counts, split_offsets=rest_file.split_offsets, sort_order_id=rest_file.sort_order_id, + equality_ids=rest_file.equality_ids if isinstance(rest_file, RESTEqualityDeleteFile) else None, + first_row_id=rest_file.first_row_id if isinstance(rest_file, RESTDataFile) else None, + referenced_data_file=rest_file.referenced_data_file if isinstance(rest_file, RESTPositionDeleteFile) else None, + content_offset=rest_file.content_offset if isinstance(rest_file, RESTPositionDeleteFile) else None, + content_size_in_bytes=rest_file.content_size_in_bytes if isinstance(rest_file, RESTPositionDeleteFile) else None, ) data_file.spec_id = rest_file.spec_id return data_file @@ -2444,9 +2603,11 @@ def _plan_files_server_side(self) -> Iterable[FileScanTask]: if self.table_identifier is None: raise ValueError("REST scan planning requires a table identifier") + # Metadata columns are materialized client-side while reading the data files + select = [name for name in self.selected_fields if not is_metadata_column(name, self.case_sensitive)] request = PlanTableScanRequest( snapshot_id=self.snapshot_id, - select=list(self.selected_fields) if self.selected_fields != ("*",) else None, + select=None if "*" in select else select, filter=self.row_filter if self.row_filter != ALWAYS_TRUE else None, case_sensitive=self.case_sensitive, ) @@ -2624,10 +2785,7 @@ def to_snapshot_id_inclusive(self: IAS, to_snapshot_id: int) -> IAS: return self.update(to_snapshot_id=to_snapshot_id) def projection(self) -> Schema: - current_schema = self.table_metadata.schema() - if "*" in self.selected_fields: - return current_schema - return current_schema.select(*self.selected_fields, case_sensitive=self.case_sensitive) + return project_schema(self.table_metadata.schema(), self.selected_fields, case_sensitive=self.case_sensitive) def plan_files(self) -> Iterable[FileScanTask]: """Plans the relevant files added between the specified snapshots.""" @@ -2827,10 +2985,8 @@ def plan_files( data_file = manifest_entry.data_file if data_file.content == DataFileContent.DATA: data_entries.append(manifest_entry) - elif data_file.content == DataFileContent.POSITION_DELETES: + elif data_file.content in (DataFileContent.POSITION_DELETES, DataFileContent.EQUALITY_DELETES): delete_index.add_delete_file(manifest_entry, partition_key=data_file.partition) - elif data_file.content == DataFileContent.EQUALITY_DELETES: - raise ValueError("PyIceberg does not yet support equality deletes: https://github.com/apache/iceberg/issues/6568") else: raise ValueError(f"Unknown DataFileContent ({data_file.content}): {manifest_entry}") @@ -2845,6 +3001,7 @@ def plan_files( residual=residual_evaluators[data_entry.data_file.spec_id](data_entry.data_file).residual_for( data_entry.data_file.partition ), + sequence_number=data_entry.sequence_number, ) for data_entry in data_entries ] diff --git a/pyiceberg/table/delete_file.py b/pyiceberg/table/delete_file.py new file mode 100644 index 0000000000..681b0eef9f --- /dev/null +++ b/pyiceberg/table/delete_file.py @@ -0,0 +1,96 @@ +# Licensed to the Apache Software Foundation (ASF) under one +# or more contributor license agreements. See the NOTICE file +# distributed with this work for additional information +# regarding copyright ownership. The ASF licenses this file +# to you under the Apache License, Version 2.0 (the +# "License"); you may not use this file except in compliance +# with the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, +# software distributed under the License is distributed on an +# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +# KIND, either express or implied. See the License for the +# specific language governing permissions and limitations +# under the License. +from __future__ import annotations + +from collections.abc import Iterable, Iterator, MutableSet +from dataclasses import dataclass +from typing import Any + +from pyiceberg.manifest import DataFile + + +@dataclass(frozen=True, slots=True) +class DeleteFileKey: + """Identity of a delete file, including its referenced content range.""" + + file_path: str + content_offset: int | None + content_size_in_bytes: int | None + + @classmethod + def from_file(cls, delete_file: DataFile) -> DeleteFileKey: + """Create a key from a delete file.""" + return cls( + file_path=delete_file.file_path, + content_offset=delete_file.content_offset, + content_size_in_bytes=delete_file.content_size_in_bytes, + ) + + +class DeleteFileSet(MutableSet[DataFile]): + """Set-like delete-file collection keyed by location and content range.""" + + _files: dict[DeleteFileKey, DataFile] + + def __init__(self, delete_files: Iterable[DataFile] = ()) -> None: + self._files = {} + for delete_file in delete_files: + self.add(delete_file) + + def __contains__(self, delete_file: object) -> bool: + """Return whether the delete file is present.""" + return isinstance(delete_file, DataFile) and DeleteFileKey.from_file(delete_file) in self._files + + def __iter__(self) -> Iterator[DataFile]: + """Return an iterator over delete files.""" + return iter(self._files.values()) + + def __len__(self) -> int: + """Return the number of delete files.""" + return len(self._files) + + def add(self, delete_file: DataFile) -> None: + self._files.setdefault(DeleteFileKey.from_file(delete_file), delete_file) + + def discard(self, delete_file: DataFile) -> None: + self._files.pop(DeleteFileKey.from_file(delete_file), None) + + def update(self, delete_files: Iterable[DataFile]) -> None: + for delete_file in delete_files: + self.add(delete_file) + + def __repr__(self) -> str: + """Return a string representation of the delete file set.""" + return f"{type(self).__name__}({list(self)!r})" + + def __eq__(self, other: Any) -> bool: + """Compare delete file sets by delete file identity.""" + if isinstance(other, DeleteFileSet): + return self._files.keys() == other._files.keys() + + if not isinstance(other, Iterable): + return False + + other_keys: set[DeleteFileKey] = set() + other_count = 0 + for delete_file in other: + if not isinstance(delete_file, DataFile): + return False + other_keys.add(DeleteFileKey.from_file(delete_file)) + other_count += 1 + + return len(other_keys) == other_count and set(self._files) == other_keys diff --git a/pyiceberg/table/delete_file_index.py b/pyiceberg/table/delete_file_index.py index 3f513aabe5..79d5308bcc 100644 --- a/pyiceberg/table/delete_file_index.py +++ b/pyiceberg/table/delete_file_index.py @@ -16,11 +16,19 @@ # under the License. from __future__ import annotations -from bisect import bisect_left +from bisect import bisect_left, bisect_right from pyiceberg.expressions import EqualTo from pyiceberg.expressions.visitors import _InclusiveMetricsEvaluator -from pyiceberg.manifest import INITIAL_SEQUENCE_NUMBER, POSITIONAL_DELETE_SCHEMA, DataFile, ManifestEntry +from pyiceberg.manifest import ( + INITIAL_SEQUENCE_NUMBER, + POSITIONAL_DELETE_SCHEMA, + DataFile, + DataFileContent, + FileFormat, + ManifestEntry, +) +from pyiceberg.table.delete_file import DeleteFileSet from pyiceberg.typedef import Record PATH_FIELD_ID = 2147483546 @@ -48,17 +56,38 @@ def _ensure_indexed(self) -> None: self._buffer = None def filter_by_seq(self, seq: int) -> list[DataFile]: + return [delete_file for delete_file, _ in self.filter_entries_by_seq(seq)] + + def filter_entries_by_seq(self, seq: int) -> list[tuple[DataFile, int]]: + """Return the delete files, with their sequence numbers, that apply to data with the given sequence number.""" self._ensure_indexed() if not self._files: return [] start_idx = bisect_left(self._seqs, seq) - return [delete_file for delete_file, _ in self._files[start_idx:]] + return self._files[start_idx:] def referenced_delete_files(self) -> list[DataFile]: self._ensure_indexed() return [data_file for data_file, _ in self._files] +class EqualityDeletes(PositionDeletes): + """Collects equality delete files and indexes them by sequence number.""" + + __slots__ = () + + def filter_entries_by_seq(self, seq: int) -> list[tuple[DataFile, int]]: + """Return the delete files that apply to data with the given sequence number. + + Equality deletes only apply to data files with a strictly lower data sequence number. + """ + self._ensure_indexed() + if not self._files: + return [] + start_idx = bisect_right(self._seqs, seq) + return self._files[start_idx:] + + def _has_path_bounds(delete_file: DataFile) -> bool: lower = delete_file.lower_bounds upper = delete_file.upper_bounds @@ -96,6 +125,10 @@ def _referenced_data_file_path(delete_file: DataFile) -> str | None: return None +def _is_deletion_vector(delete_file: DataFile) -> bool: + return delete_file.file_format == FileFormat.PUFFIN + + def _partition_key(spec_id: int, partition: Record | None) -> tuple[int, Record]: if partition: return spec_id, partition @@ -103,18 +136,48 @@ def _partition_key(spec_id: int, partition: Record | None) -> tuple[int, Record] class DeleteFileIndex: - """Indexes position delete files by partition and by exact data file path.""" + """Indexes delete files for scan planning. + + Deletion vectors are indexed by referenced data file, position delete files by partition and by exact + data file path, and equality delete files by partition or globally when written with an unpartitioned spec. + """ def __init__(self) -> None: self._by_partition: dict[tuple[int, Record], PositionDeletes] = {} self._by_path: dict[str, PositionDeletes] = {} + self._dvs_by_path: dict[str, PositionDeletes] = {} + self._eq_by_partition: dict[tuple[int, Record], EqualityDeletes] = {} + self._eq_global: EqualityDeletes = EqualityDeletes() + self._has_global_eq_deletes = False def is_empty(self) -> bool: - return not self._by_partition and not self._by_path + return ( + not self._by_partition + and not self._by_path + and not self._dvs_by_path + and not self._eq_by_partition + and not self._has_global_eq_deletes + ) def add_delete_file(self, manifest_entry: ManifestEntry, partition_key: Record | None = None) -> None: delete_file = manifest_entry.data_file seq = manifest_entry.sequence_number or INITIAL_SEQUENCE_NUMBER + + if delete_file.content == DataFileContent.EQUALITY_DELETES: + # Equality deletes stored with an unpartitioned spec apply to every partition + if not partition_key: + self._eq_global.add(delete_file, seq) + self._has_global_eq_deletes = True + else: + key = _partition_key(delete_file.spec_id or 0, partition_key) + self._eq_by_partition.setdefault(key, EqualityDeletes()).add(delete_file, seq) + return + + # A deletion vector applies only to the data file it references + if _is_deletion_vector(delete_file) and delete_file.referenced_data_file is not None: + self._dvs_by_path.setdefault(delete_file.referenced_data_file, PositionDeletes()).add(delete_file, seq) + return + target_path = _referenced_data_file_path(delete_file) if target_path: @@ -125,25 +188,43 @@ def add_delete_file(self, manifest_entry: ManifestEntry, partition_key: Record | deletes = self._by_partition.setdefault(key, PositionDeletes()) deletes.add(delete_file, seq) - def for_data_file(self, seq_num: int, data_file: DataFile, partition_key: Record | None = None) -> set[DataFile]: + def for_data_file(self, seq_num: int, data_file: DataFile, partition_key: Record | None = None) -> DeleteFileSet: if self.is_empty(): - return set() + return DeleteFileSet() - deletes: set[DataFile] = set() + candidates: list[tuple[DataFile, int]] = [] spec_id = data_file.spec_id or 0 key = _partition_key(spec_id, partition_key) partition_deletes = self._by_partition.get(key) if partition_deletes: - for delete_file in partition_deletes.filter_by_seq(seq_num): - if _applies_to_data_file(delete_file, data_file): - deletes.add(delete_file) + candidates.extend( + (delete_file, seq) + for delete_file, seq in partition_deletes.filter_entries_by_seq(seq_num) + if _applies_to_data_file(delete_file, data_file) + ) path_deletes = self._by_path.get(data_file.file_path) if path_deletes: - deletes.update(path_deletes.filter_by_seq(seq_num)) + candidates.extend(path_deletes.filter_entries_by_seq(seq_num)) + + dvs = self._dvs_by_path.get(data_file.file_path) + dv_entries = dvs.filter_entries_by_seq(seq_num) if dvs else [] + if dv_entries: + # At most one DV applies to a data file; it replaces older position deletes for that file. + # The entries are sorted by sequence number, so the last one is the most recent DV. + dv, dv_seq = dv_entries[-1] + candidates = [(delete_file, seq) for delete_file, seq in candidates if seq > dv_seq] + candidates.append((dv, dv_seq)) + + if self._has_global_eq_deletes: + candidates.extend(self._eq_global.filter_entries_by_seq(seq_num)) - return deletes + eq_partition_deletes = self._eq_by_partition.get(key) + if eq_partition_deletes: + candidates.extend(eq_partition_deletes.filter_entries_by_seq(seq_num)) + + return DeleteFileSet(delete_file for delete_file, _ in candidates) def referenced_delete_files(self) -> list[DataFile]: data_files: list[DataFile] = [] @@ -154,4 +235,12 @@ def referenced_delete_files(self) -> list[DataFile]: for deletes in self._by_path.values(): data_files.extend(deletes.referenced_delete_files()) + for deletes in self._dvs_by_path.values(): + data_files.extend(deletes.referenced_delete_files()) + + for eq_deletes in self._eq_by_partition.values(): + data_files.extend(eq_deletes.referenced_delete_files()) + + data_files.extend(self._eq_global.referenced_delete_files()) + return data_files diff --git a/pyiceberg/table/deletion_vector.py b/pyiceberg/table/deletion_vector.py index 88fb3daf73..e3d8404668 100644 --- a/pyiceberg/table/deletion_vector.py +++ b/pyiceberg/table/deletion_vector.py @@ -14,19 +14,45 @@ # KIND, either express or implied. See the License for the # specific language governing permissions and limitations # under the License. +import io import math -from typing import TYPE_CHECKING +import struct +import zlib +from collections.abc import Iterable +from typing import TYPE_CHECKING, cast from pyroaring import BitMap, FrozenBitMap -from pyiceberg.table.puffin import PuffinFile +from pyiceberg.table.puffin import PuffinBlob, PuffinBlobMetadata, PuffinFile, PuffinWriter if TYPE_CHECKING: import pyarrow as pa + from pyiceberg.io import FileIO + from pyiceberg.manifest import DataFile + EMPTY_BITMAP = FrozenBitMap() MAX_JAVA_SIGNED = int(math.pow(2, 31)) - 1 +# Largest addressable position, mirroring Java's RoaringPositionBitmap.MAX_POSITION +# (toPosition(Integer.MAX_VALUE - 1, Integer.MIN_VALUE)): the high 32 bits hold the bitmap +# key and the low 32 bits the position within that bitmap. +MAX_POSITION = ((MAX_JAVA_SIGNED - 1) << 32) | 0x80000000 PROPERTY_REFERENCED_DATA_FILE = "referenced-data-file" +PROPERTY_CARDINALITY = "cardinality" +DELETION_VECTOR_V1_BLOB_TYPE = "deletion-vector-v1" +DELETION_VECTOR_MAGIC = b"\xd1\xd3\x39\x64" +# Reserved field id of the row position (_pos) metadata column, referenced by +# deletion-vector-v1 blob metadata (Java: MetadataColumns.ROW_POSITION) +ROW_POSITION_FIELD_ID = 2147483645 +# Snapshot id and sequence number of a DV blob are inherited from the manifest entry at commit +_INHERITED = -1 +_MAX_DELETION_VECTOR_CONTENT_SIZE = 2**31 - 1 +_DV_BLOB_LENGTH = struct.Struct(">I") +_DV_BLOB_MAGIC = struct.Struct("I") +_DV_BLOB_MAGIC_NUMBER = 1681511377 +_ROARING_BITMAP_COUNT_SIZE_BYTES = 8 +_DV_BLOB_MIN_SIZE_BYTES = _DV_BLOB_LENGTH.size + _DV_BLOB_MAGIC.size + _ROARING_BITMAP_COUNT_SIZE_BYTES + _DV_BLOB_CRC.size class DeletionVector: @@ -37,11 +63,71 @@ def __init__(self, referenced_data_file: str, bitmaps: list[BitMap]) -> None: self.referenced_data_file = referenced_data_file self._bitmaps = bitmaps + @classmethod + def from_positions(cls, referenced_data_file: str, positions: Iterable[int]) -> "DeletionVector": + """Create a deletion vector marking the given row positions of a data file as deleted.""" + bitmaps_by_key: dict[int, BitMap] = {} + for position in positions: + if position < 0 or position > MAX_POSITION: + raise ValueError(f"Invalid position: {position}, must be between 0 and {MAX_POSITION}") + key = position >> 32 + if (bitmap := bitmaps_by_key.get(key)) is None: + bitmap = bitmaps_by_key[key] = BitMap() + bitmap.add(position & 0xFFFFFFFF) + + if not bitmaps_by_key: + raise ValueError("Deletion vector must contain at least one position") + + # A list indexed by key, padding gaps with the empty bitmap (mirrors _deserialize_bitmap) + bitmaps: list[BitMap] = [bitmaps_by_key.get(key, EMPTY_BITMAP) for key in range(max(bitmaps_by_key) + 1)] + return cls(referenced_data_file, bitmaps) + + @property + def cardinality(self) -> int: + """Return the number of deleted positions.""" + return sum(len(bitmap) for bitmap in self._bitmaps) + + def union(self, other: "DeletionVector") -> "DeletionVector": + """Return a deletion vector with the positions deleted by either vector.""" + if other.referenced_data_file != self.referenced_data_file: + raise ValueError(f"Cannot union deletion vectors of {self.referenced_data_file} and {other.referenced_data_file}") + bitmaps: list[BitMap] = [] + for key in range(max(len(self._bitmaps), len(other._bitmaps))): + left = self._bitmaps[key] if key < len(self._bitmaps) else EMPTY_BITMAP + right = other._bitmaps[key] if key < len(other._bitmaps) else EMPTY_BITMAP + bitmaps.append(BitMap(left) | right) + return DeletionVector(self.referenced_data_file, bitmaps) + + def serialize(self) -> bytes: + """Serialize the positions as a portable 64-bit Roaring bitmap (the vector inside a deletion-vector-v1 blob).""" + return self._serialize_bitmap(self._bitmaps) + + def to_blob(self) -> PuffinBlob: + """Return the deletion-vector-v1 Puffin blob of this vector; offset and length are assigned by the writer.""" + metadata = PuffinBlobMetadata( + type=DELETION_VECTOR_V1_BLOB_TYPE, + fields=[ROW_POSITION_FIELD_ID], + snapshot_id=_INHERITED, + sequence_number=_INHERITED, + offset=0, + length=0, + properties={ + PROPERTY_REFERENCED_DATA_FILE: self.referenced_data_file, + PROPERTY_CARDINALITY: str(self.cardinality), + }, + ) + return PuffinBlob(metadata=metadata, payload=_serialize_dv_blob(self.serialize())) + @staticmethod def _deserialize_bitmap(pl: bytes) -> list[BitMap]: number_of_bitmaps = int.from_bytes(pl[0:8], byteorder="little") pl = pl[8:] + # Every bitmap contributes at least a 4-byte key, so a count that cannot fit + # in the remaining payload is invalid and must not be used as a loop bound. + if number_of_bitmaps * 4 > len(pl): + raise ValueError(f"Payload declares {number_of_bitmaps} bitmaps, but only holds {len(pl)} bytes") + bitmaps = [] last_key = -1 for _ in range(number_of_bitmaps): @@ -67,6 +153,25 @@ def _deserialize_bitmap(pl: bytes) -> list[BitMap]: return bitmaps + @staticmethod + def _serialize_bitmap(bitmaps: list[BitMap]) -> bytes: + # Counterpart of _deserialize_bitmap: number of bitmaps (8 bytes, little-endian), then for each + # non-empty bitmap in ascending key order its key (4 bytes, little-endian) and portable payload. + non_empty = [(key, bitmap) for key, bitmap in enumerate(bitmaps) if len(bitmap) > 0] + + with io.BytesIO() as out: + out.write(len(non_empty).to_bytes(8, byteorder="little")) + for key, bitmap in non_empty: + # Java's RoaringPositionBitmap rejects keys above Integer.MAX_VALUE - 1 + if key > MAX_JAVA_SIGNED - 1: + raise ValueError(f"Key {key} is too large, max {MAX_JAVA_SIGNED - 1} for compatibility with the Java impl") + out.write(key.to_bytes(4, byteorder="little")) + # Run-length encode so contiguous deletes stay compact, matching Java's BitmapPositionDeleteIndex + optimized = BitMap(bitmap) + optimized.run_optimize() + out.write(optimized.serialize()) + return out.getvalue() + @staticmethod def _bitmaps_to_chunked_array(bitmaps: list[BitMap]) -> "pa.ChunkedArray": import pyarrow as pa @@ -77,17 +182,156 @@ def to_vector(self) -> "pa.ChunkedArray": return self._bitmaps_to_chunked_array(self._bitmaps) -def _extract_vector_payload(blob_payload: bytes) -> bytes: - """Strip deletion-vector-v1 blob framing: length(4 big-endian) + DV magic(4) ... CRC(4 big-endian).""" - length_prefix = int.from_bytes(blob_payload[0:4], "big") - return blob_payload[8 : 4 + length_prefix] +def _serialize_dv_blob(bitmap_payload: bytes) -> bytes: + # Counterpart of _deserialize_dv_blob: 4-byte big-endian length of magic + vector, the magic, + # the vector, and a 4-byte big-endian CRC-32 of magic + vector. + bitmap_data = _DV_BLOB_MAGIC.pack(_DV_BLOB_MAGIC_NUMBER) + bitmap_payload + if len(bitmap_data) > _MAX_DELETION_VECTOR_CONTENT_SIZE: + raise ValueError(f"Cannot write deletion vector larger than 2GB: {len(bitmap_data)}") + return _DV_BLOB_LENGTH.pack(len(bitmap_data)) + bitmap_data + _DV_BLOB_CRC.pack(zlib.crc32(bitmap_data) & 0xFFFFFFFF) + + +def _deserialize_dv_blob(blob: bytes, record_count: int | None = None) -> list[BitMap]: + # The DV blob encoding matches Iceberg Java's BitmapPositionDeleteIndex: + # 4-byte big-endian bitmap-data length, 4-byte little-endian magic number, + # portable Roaring bitmap data, and 4-byte big-endian CRC-32. + if len(blob) < _DV_BLOB_MIN_SIZE_BYTES: + raise ValueError(f"Invalid deletion vector blob length: {len(blob)}") + + bitmap_data_length = _DV_BLOB_LENGTH.unpack_from(blob)[0] + expected_bitmap_data_length = len(blob) - _DV_BLOB_LENGTH.size - _DV_BLOB_CRC.size + if bitmap_data_length != expected_bitmap_data_length: + raise ValueError(f"Invalid bitmap data length: {bitmap_data_length}, expected {expected_bitmap_data_length}") + + bitmap_data_offset = _DV_BLOB_LENGTH.size + crc_offset = bitmap_data_offset + bitmap_data_length + bitmap_data = blob[bitmap_data_offset:crc_offset] + + magic_number = _DV_BLOB_MAGIC.unpack_from(bitmap_data)[0] + if magic_number != _DV_BLOB_MAGIC_NUMBER: + raise ValueError(f"Invalid magic number: {magic_number}, expected {_DV_BLOB_MAGIC_NUMBER}") + + checksum = zlib.crc32(bitmap_data) & 0xFFFFFFFF + expected_checksum = _DV_BLOB_CRC.unpack_from(blob, crc_offset)[0] + if checksum != expected_checksum: + raise ValueError("Invalid CRC") + + bitmaps = DeletionVector._deserialize_bitmap(bitmap_data[_DV_BLOB_MAGIC.size :]) + if record_count is not None: + cardinality = sum(len(bitmap) for bitmap in bitmaps) + if cardinality != record_count: + raise ValueError(f"Invalid cardinality: {cardinality}, expected {record_count}") + + return bitmaps + + +def _validate_deletion_vector_content(dv: "DataFile") -> None: + content_offset = dv.content_offset + content_size_in_bytes = dv.content_size_in_bytes + referenced_data_file = dv.referenced_data_file + + if content_offset is None: + raise ValueError(f"Invalid deletion vector, content offset is missing: {dv.file_path}") + if content_size_in_bytes is None: + raise ValueError(f"Invalid deletion vector, content size is missing: {dv.file_path}") + if content_offset < 0: + raise ValueError(f"Invalid deletion vector, content offset cannot be negative: {content_offset}") + if content_size_in_bytes < 0: + raise ValueError(f"Invalid deletion vector, content size cannot be negative: {content_size_in_bytes}") + if content_size_in_bytes > _MAX_DELETION_VECTOR_CONTENT_SIZE: + raise ValueError(f"Cannot read deletion vector larger than 2GB: {content_size_in_bytes}") + if referenced_data_file is None: + raise ValueError(f"Invalid deletion vector, referenced data file is missing: {dv.file_path}") + + +def has_deletion_vector_content_reference(dv: "DataFile") -> bool: + """Return whether a deletion vector is described by manifest content-range metadata.""" + return dv.content_offset is not None or dv.content_size_in_bytes is not None or dv.referenced_data_file is not None + + +def _read_deletion_vector(io: "FileIO", dv: "DataFile") -> DeletionVector: + _validate_deletion_vector_content(dv) + + content_offset = cast(int, dv.content_offset) + content_size_in_bytes = cast(int, dv.content_size_in_bytes) + referenced_data_file = cast(str, dv.referenced_data_file) + + with io.new_input(dv.file_path).open() as fi: + fi.seek(content_offset) + payload = fi.read(content_size_in_bytes) + + if len(payload) != content_size_in_bytes: + raise ValueError(f"Could not read deletion vector, expected {content_size_in_bytes} bytes, got {len(payload)}") + + return DeletionVector( + referenced_data_file=referenced_data_file, + bitmaps=_deserialize_dv_blob(payload, dv.record_count), + ) + + +def read_deletion_vectors(io: "FileIO", dv: "DataFile") -> list[DeletionVector]: + """Read deletion vectors from a delete file or its manifest content range.""" + if has_deletion_vector_content_reference(dv): + return [_read_deletion_vector(io, dv)] + + with io.new_input(dv.file_path).open() as fi: + return deletion_vectors_from_puffin_file(PuffinFile(fi.read())) def deletion_vectors_from_puffin_file(puffin_file: PuffinFile) -> list[DeletionVector]: + """Read all deletion vectors stored in a Puffin file, skipping blobs of other types.""" + deletion_vectors = [] + for blob in puffin_file.footer.blobs: + if blob.type != DELETION_VECTOR_V1_BLOB_TYPE: + continue + if PROPERTY_REFERENCED_DATA_FILE not in blob.properties: + raise ValueError(f"Invalid deletion vector blob at offset {blob.offset}, {PROPERTY_REFERENCED_DATA_FILE} is missing") + deletion_vectors.append( + DeletionVector( + referenced_data_file=blob.properties[PROPERTY_REFERENCED_DATA_FILE], + bitmaps=_deserialize_dv_blob(puffin_file.get_blob_payload(blob)), + ) + ) + return deletion_vectors + + +def write_deletion_vectors( + io: "FileIO", location: str, deletion_vectors: Iterable[tuple["DataFile", DeletionVector]] +) -> list["DataFile"]: + """Write deletion vectors into one Puffin file and return a delete file entry per vector. + + Args: + io: The FileIO to write the Puffin file with. + location: The location of the Puffin file. + deletion_vectors: Pairs of the referenced data file and its deletion vector. The data file + provides the partition and partition spec of the returned delete file. + + Returns: + The position delete files, one per deletion vector, carrying the blob's content range. + """ + from pyiceberg.manifest import DataFile, DataFileContent, FileFormat + + written: list[tuple[DataFile, DeletionVector, PuffinBlobMetadata]] = [] + with PuffinWriter(io.new_output(location)) as writer: + for data_file, dv in deletion_vectors: + if dv.referenced_data_file != data_file.file_path: + raise ValueError(f"Deletion vector for {dv.referenced_data_file} does not reference {data_file.file_path}") + written.append((data_file, dv, writer.add_blob(dv.to_blob()))) + + file_size = cast(int, writer.file_size) return [ - DeletionVector( - referenced_data_file=blob.properties[PROPERTY_REFERENCED_DATA_FILE], - bitmaps=DeletionVector._deserialize_bitmap(_extract_vector_payload(puffin_file.get_blob_payload(blob))), + DataFile.from_args( + _table_format_version=3, + spec_id=data_file.spec_id, + content=DataFileContent.POSITION_DELETES, + file_path=location, + file_format=FileFormat.PUFFIN, + partition=data_file.partition, + record_count=dv.cardinality, + file_size_in_bytes=file_size, + referenced_data_file=dv.referenced_data_file, + content_offset=blob.offset, + content_size_in_bytes=blob.length, ) - for blob in puffin_file.footer.blobs + for data_file, dv, blob in written ] diff --git a/pyiceberg/table/inspect.py b/pyiceberg/table/inspect.py index e24e251fa4..aa64b48326 100644 --- a/pyiceberg/table/inspect.py +++ b/pyiceberg/table/inspect.py @@ -26,7 +26,7 @@ from pyiceberg.manifest import DataFile, DataFileContent, ManifestContent, ManifestFile, PartitionFieldSummary from pyiceberg.partitioning import PartitionSpec from pyiceberg.table.snapshots import Snapshot, ancestors_of -from pyiceberg.types import PrimitiveType +from pyiceberg.types import GeographyType, GeometryType, IcebergType, PrimitiveType, UnknownType, VariantType from pyiceberg.utils.concurrent import ExecutorFactory from pyiceberg.utils.singleton import _convert_to_hashable_type @@ -39,7 +39,26 @@ def _readable_bound(field_type: PrimitiveType, bound: bytes | None) -> Any | None: - return from_bytes(field_type, bound) if bound is not None else None + if bound is None or isinstance(field_type, UnknownType): + return None + if isinstance(field_type, VariantType): + # Variant bounds are serialized Variant objects keyed by field path, surfaced as raw bytes + return bound + return from_bytes(field_type, bound) + + +def _readable_bound_arrow_type(field_type: IcebergType) -> pa.DataType: + """Return the Arrow type of a readable lower/upper bound for a column.""" + import pyarrow as pa + + from pyiceberg.io.pyarrow import schema_to_pyarrow + + if isinstance(field_type, (GeometryType, GeographyType, VariantType)): + # Geo bounds are WKB bounding-box points, surfaced as raw bytes rather than as a GeoArrow extension type + return pa.binary() + if isinstance(field_type, UnknownType): + return pa.null() + return schema_to_pyarrow(field_type) class InspectTable: @@ -176,7 +195,7 @@ def entries(self, snapshot_id: int | None = None) -> pa.Table: readable_metrics_struct = [] def _readable_metrics_struct(bound_type: PrimitiveType) -> pa.StructType: - pa_bound_type = schema_to_pyarrow(bound_type) + pa_bound_type = _readable_bound_arrow_type(bound_type) return pa.struct( [ pa.field("column_size", pa.int64(), nullable=True), @@ -222,6 +241,10 @@ def _readable_metrics_struct(bound_type: PrimitiveType) -> pa.StructType: pa.field("split_offsets", pa.list_(pa.int64()), nullable=True), pa.field("equality_ids", pa.list_(pa.int32()), nullable=True), pa.field("sort_order_id", pa.int32(), nullable=True), + pa.field("first_row_id", pa.int64(), nullable=True), + pa.field("referenced_data_file", pa.string(), nullable=True), + pa.field("content_offset", pa.int64(), nullable=True), + pa.field("content_size_in_bytes", pa.int64(), nullable=True), ] ), nullable=False, @@ -272,7 +295,7 @@ def _readable_metrics_struct(bound_type: PrimitiveType) -> pa.StructType: "partition": partition_record_dict, "record_count": entry.data_file.record_count, "file_size_in_bytes": entry.data_file.file_size_in_bytes, - "column_sizes": dict(entry.data_file.column_sizes), + "column_sizes": dict(entry.data_file.column_sizes or {}), "value_counts": dict(entry.data_file.value_counts or {}), "null_value_counts": dict(entry.data_file.null_value_counts or {}), "nan_value_counts": dict(entry.data_file.nan_value_counts or {}), @@ -282,6 +305,10 @@ def _readable_metrics_struct(bound_type: PrimitiveType) -> pa.StructType: "split_offsets": entry.data_file.split_offsets, "equality_ids": entry.data_file.equality_ids, "sort_order_id": entry.data_file.sort_order_id, + "first_row_id": entry.data_file.first_row_id, + "referenced_data_file": entry.data_file.referenced_data_file, + "content_offset": entry.data_file.content_offset, + "content_size_in_bytes": entry.data_file.content_size_in_bytes, "spec_id": entry.data_file.spec_id, }, "readable_metrics": readable_metrics, @@ -783,6 +810,10 @@ def _get_files_from_manifest( "split_offsets": data_file.split_offsets, "equality_ids": data_file.equality_ids, "sort_order_id": data_file.sort_order_id, + "first_row_id": data_file.first_row_id, + "referenced_data_file": data_file.referenced_data_file, + "content_offset": data_file.content_offset, + "content_size_in_bytes": data_file.content_size_in_bytes, "readable_metrics": readable_metrics, } ) @@ -803,7 +834,9 @@ def _get_files_schema(self) -> pa.Schema: ``file_size_in_bytes``, ``column_sizes``, ``value_counts``, ``null_value_counts``, ``nan_value_counts``, ``lower_bounds``, ``upper_bounds``, ``key_metadata``, ``split_offsets``, - ``equality_ids``, ``sort_order_id``, and ``readable_metrics``. + ``equality_ids``, ``sort_order_id``, ``first_row_id``, + ``referenced_data_file``, ``content_offset``, + ``content_size_in_bytes``, and ``readable_metrics``. """ import pyarrow as pa @@ -813,7 +846,7 @@ def _get_files_schema(self) -> pa.Schema: readable_metrics_struct = [] def _readable_metrics_struct(bound_type: PrimitiveType) -> pa.StructType: - pa_bound_type = schema_to_pyarrow(bound_type) + pa_bound_type = _readable_bound_arrow_type(bound_type) return pa.struct( [ pa.field("column_size", pa.int64(), nullable=True), @@ -852,6 +885,10 @@ def _readable_metrics_struct(bound_type: PrimitiveType) -> pa.StructType: pa.field("split_offsets", pa.list_(pa.int64()), nullable=True), pa.field("equality_ids", pa.list_(pa.int32()), nullable=True), pa.field("sort_order_id", pa.int32(), nullable=True), + pa.field("first_row_id", pa.int64(), nullable=True), + pa.field("referenced_data_file", pa.string(), nullable=True), + pa.field("content_offset", pa.int64(), nullable=True), + pa.field("content_size_in_bytes", pa.int64(), nullable=True), pa.field("readable_metrics", pa.struct(readable_metrics_struct), nullable=True), ] ) diff --git a/pyiceberg/table/metadata.py b/pyiceberg/table/metadata.py index 8236f12229..e7566c56fc 100644 --- a/pyiceberg/table/metadata.py +++ b/pyiceberg/table/metadata.py @@ -68,9 +68,10 @@ INITIAL_SEQUENCE_NUMBER = 0 INITIAL_SPEC_ID = 0 +INITIAL_ROW_ID = 0 DEFAULT_SCHEMA_ID = 0 -SUPPORTED_TABLE_FORMAT_VERSION = 2 +SUPPORTED_TABLE_FORMAT_VERSION = 3 def cleanup_snapshot_id(data: dict[str, Any]) -> dict[str, Any]: @@ -615,8 +616,12 @@ def construct_refs(self) -> TableMetadata: """The table’s highest assigned sequence number, a monotonically increasing long that tracks the order of snapshots in a table.""" - next_row_id: int | None = Field(alias="next-row-id", default=None) - """A long higher than all assigned row IDs; the next snapshot's `first-row-id`.""" + next_row_id: int = Field(alias="next-row-id") + """A long higher than all assigned row IDs; the next snapshot's `first-row-id`. + + Required in V3: starting at 0 for a missing value could hand out row IDs that are already assigned. + Set to 0 when a table is created or upgraded to V3. + """ encryption_keys: list[EncryptedKey] = Field(alias="encryption-keys", default_factory=list) """An optional list of encryption keys used for table encryption.""" @@ -629,9 +634,6 @@ def serialize_model(self, handler: ModelWrapSerializerWithoutInfo) -> dict[str, serialized.pop("encryption-keys", None) return serialized - def model_dump_json(self, exclude_none: bool = True, exclude: Any | None = None, by_alias: bool = True, **kwargs: Any) -> str: - raise NotImplementedError("Writing V3 is not yet supported, see: https://github.com/apache/iceberg-python/issues/1551") - TableMetadata = Annotated[TableMetadataV1 | TableMetadataV2 | TableMetadataV3, Field(discriminator="format_version")] @@ -700,6 +702,7 @@ def new_table_metadata( properties=properties, last_partition_id=fresh_partition_spec.last_assigned_field_id, table_uuid=table_uuid, + next_row_id=INITIAL_ROW_ID, ) else: raise ValidationError(f"Unknown format version: {format_version}") diff --git a/pyiceberg/table/metadata_columns.py b/pyiceberg/table/metadata_columns.py new file mode 100644 index 0000000000..8eb2e487e3 --- /dev/null +++ b/pyiceberg/table/metadata_columns.py @@ -0,0 +1,123 @@ +# Licensed to the Apache Software Foundation (ASF) under one +# or more contributor license agreements. See the NOTICE file +# distributed with this work for additional information +# regarding copyright ownership. The ASF licenses this file +# to you under the Apache License, Version 2.0 (the +# "License"); you may not use this file except in compliance +# with the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, +# software distributed under the License is distributed on an +# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +# KIND, either express or implied. See the License for the +# specific language governing permissions and limitations +# under the License. +"""Reserved metadata columns that can be requested in a scan projection. + +See https://iceberg.apache.org/spec/#reserved-field-ids and https://iceberg.apache.org/spec/#row-lineage. +""" + +from __future__ import annotations + +from pyiceberg.schema import Schema +from pyiceberg.types import LongType, NestedField + +ROW_ID = NestedField( + field_id=2147483540, + name="_row_id", + field_type=LongType(), + required=False, + doc="Implicit row ID that is automatically assigned", +) +LAST_UPDATED_SEQUENCE_NUMBER = NestedField( + field_id=2147483539, + name="_last_updated_sequence_number", + field_type=LongType(), + required=False, + doc="Sequence number when the row was last updated", +) + +ROW_LINEAGE_COLUMNS: tuple[NestedField, ...] = (ROW_ID, LAST_UPDATED_SEQUENCE_NUMBER) +_METADATA_COLUMNS_BY_NAME: dict[str, NestedField] = {field.name: field for field in ROW_LINEAGE_COLUMNS} +_METADATA_COLUMNS_BY_ID: dict[int, NestedField] = {field.field_id: field for field in ROW_LINEAGE_COLUMNS} + + +def _lookup(name: str, case_sensitive: bool = True) -> NestedField | None: + if case_sensitive: + return _METADATA_COLUMNS_BY_NAME.get(name) + return next((field for field in ROW_LINEAGE_COLUMNS if field.name.lower() == name.lower()), None) + + +def is_metadata_column(name: str, case_sensitive: bool = True) -> bool: + """Return whether the name refers to a supported reserved metadata column.""" + return _lookup(name, case_sensitive) is not None + + +def is_metadata_field_id(field_id: int) -> bool: + """Return whether the field ID belongs to a supported reserved metadata column.""" + return field_id in _METADATA_COLUMNS_BY_ID + + +def schema_with_row_lineage(schema: Schema) -> Schema | None: + """Return the schema with the row lineage columns appended, or None when a table column uses one of their names. + + Copy-on-write rewrites write this schema, so the row lineage of copied rows is kept in the new data files. + """ + if any(_has_field(schema, field.name, case_sensitive=True) for field in ROW_LINEAGE_COLUMNS): + return None + return Schema( + *schema.fields, + *ROW_LINEAGE_COLUMNS, + schema_id=schema.schema_id, + identifier_field_ids=schema.identifier_field_ids, + ) + + +def project_schema(schema: Schema, selected_fields: tuple[str, ...], case_sensitive: bool = True) -> Schema: + """Select columns from a table schema, appending any requested metadata columns in the requested order. + + Args: + schema: The table schema to select from. + selected_fields: Column names to select; ``*`` selects all table columns. + case_sensitive: Whether name lookups are case-sensitive. + + Returns: + The projected schema, with metadata columns after the table columns. + """ + table_fields: list[str] = [] + metadata_fields: list[NestedField] = [] + for name in selected_fields: + metadata_field = _lookup(name, case_sensitive) + # A table column with the same name takes precedence over the metadata column + if metadata_field is not None and not _has_field(schema, name, case_sensitive): + if metadata_field not in metadata_fields: + metadata_fields.append(metadata_field) + else: + table_fields.append(name) + + if "*" in table_fields: + projected = schema + elif table_fields: + projected = schema.select(*table_fields, case_sensitive=case_sensitive) + else: + projected = Schema(schema_id=schema.schema_id) + + if not metadata_fields: + return projected + + return Schema( + *projected.fields, + *metadata_fields, + schema_id=projected.schema_id, + identifier_field_ids=projected.identifier_field_ids, + ) + + +def _has_field(schema: Schema, name: str, case_sensitive: bool) -> bool: + try: + schema.find_field(name, case_sensitive) + return True + except ValueError: + return False diff --git a/pyiceberg/table/puffin.py b/pyiceberg/table/puffin.py index 13315bf802..f60b7c755f 100644 --- a/pyiceberg/table/puffin.py +++ b/pyiceberg/table/puffin.py @@ -14,11 +14,15 @@ # KIND, either express or implied. See the License for the # specific language governing permissions and limitations # under the License. +from dataclasses import dataclass +from types import TracebackType from typing import TYPE_CHECKING import zstandard from pydantic import Field +from pyiceberg import __version__ +from pyiceberg.io import OutputFile from pyiceberg.typedef import IcebergBaseModel from pyiceberg.utils.deprecated import deprecated @@ -27,6 +31,8 @@ # Short for: Puffin Fratercula arctica, version 1 MAGIC_BYTES = b"PFA1" +# Footer trailer: footer payload size (4 bytes), flags (4 bytes) and the closing magic +_FOOTER_TRAILER_SIZE = 12 class PuffinBlobMetadata(IcebergBaseModel): @@ -82,3 +88,90 @@ def to_vector(self) -> dict[str, "pa.ChunkedArray"]: from pyiceberg.table.deletion_vector import deletion_vectors_from_puffin_file # local import avoids the cycle return {dv.referenced_data_file: dv.to_vector() for dv in deletion_vectors_from_puffin_file(self)} + + +@dataclass(frozen=True) +class PuffinBlob: + """A blob to write into a Puffin file: its metadata and serialized payload.""" + + metadata: PuffinBlobMetadata + payload: bytes + + +class PuffinWriter: + """Assembles a Puffin file from blobs and writes it to an output file. + + The writer is blob-agnostic: callers supply already-serialized blobs (for example via + ``DeletionVector.to_blob()``). The offset and length of each blob are assigned when it is + added, and the file is written on ``close()`` (or on exiting the context manager without + an exception), after which ``file_size`` holds the size of the written file. + """ + + closed: bool + file_size: int | None + _output_file: OutputFile + _blobs: list[PuffinBlob] + _next_offset: int + _created_by: str + + def __init__(self, output_file: OutputFile, created_by: str | None = None) -> None: + self.closed = False + self.file_size = None + self._output_file = output_file + self._blobs = [] + self._next_offset = len(MAGIC_BYTES) + self._created_by = created_by if created_by is not None else f"PyIceberg version {__version__}" + + def __enter__(self) -> "PuffinWriter": + """Open the writer.""" + return self + + def __exit__( + self, + exc_type: type[BaseException] | None, + exc_value: BaseException | None, + traceback: TracebackType | None, + ) -> None: + """Write the Puffin file, unless the body raised.""" + if exc_type is not None: + # Skip writing a half-populated file + self.closed = True + return + self.close() + + @property + def location(self) -> str: + """Return the location of the output file.""" + return self._output_file.location + + def add_blob(self, blob: PuffinBlob) -> PuffinBlobMetadata: + """Add a blob and return its metadata with the offset and length it has in the file.""" + if self.closed: + raise RuntimeError("Cannot add blob to closed Puffin writer") + metadata = blob.metadata.model_copy(update={"offset": self._next_offset, "length": len(blob.payload)}) + self._blobs.append(PuffinBlob(metadata=metadata, payload=blob.payload)) + self._next_offset += len(blob.payload) + return metadata + + def close(self) -> int: + """Write the blobs and the footer to the output file and return the size of the file.""" + if self.closed: + raise RuntimeError("Puffin writer is already closed") + self.closed = True + + footer = Footer(blobs=[blob.metadata for blob in self._blobs], properties={"created-by": self._created_by}) + footer_payload = footer.model_dump_json(by_alias=True, exclude_none=True).encode("utf-8") + + with self._output_file.create(overwrite=True) as output_stream: + output_stream.write(MAGIC_BYTES) + for blob in self._blobs: + output_stream.write(blob.payload) + output_stream.write(MAGIC_BYTES) + output_stream.write(footer_payload) + output_stream.write(len(footer_payload).to_bytes(4, byteorder="little")) + # Flags: the footer payload is not compressed + output_stream.write((0).to_bytes(4, byteorder="little")) + output_stream.write(MAGIC_BYTES) + + self.file_size = self._next_offset + len(MAGIC_BYTES) + len(footer_payload) + _FOOTER_TRAILER_SIZE + return self.file_size diff --git a/pyiceberg/table/snapshots.py b/pyiceberg/table/snapshots.py index 0450df2861..db72a137eb 100644 --- a/pyiceberg/table/snapshots.py +++ b/pyiceberg/table/snapshots.py @@ -27,7 +27,7 @@ from pyiceberg.environment_context import EnvironmentContext from pyiceberg.io import FileIO -from pyiceberg.manifest import DataFile, DataFileContent, ManifestFile, _manifests +from pyiceberg.manifest import DataFile, DataFileContent, FileFormat, ManifestFile, _manifests from pyiceberg.partitioning import UNPARTITIONED_PARTITION_SPEC, PartitionSpec from pyiceberg.schema import Schema @@ -37,6 +37,7 @@ ADDED_DATA_FILES = "added-data-files" ADDED_DELETE_FILES = "added-delete-files" +ADDED_DVS = "added-dvs" ADDED_EQUALITY_DELETES = "added-equality-deletes" ADDED_FILE_SIZE = "added-files-size" ADDED_POSITION_DELETES = "added-position-deletes" @@ -46,6 +47,7 @@ DELETED_RECORDS = "deleted-records" ADDED_EQUALITY_DELETE_FILES = "added-equality-delete-files" REMOVED_DELETE_FILES = "removed-delete-files" +REMOVED_DVS = "removed-dvs" REMOVED_EQUALITY_DELETES = "removed-equality-deletes" REMOVED_EQUALITY_DELETE_FILES = "removed-equality-delete-files" REMOVED_FILE_SIZE = "removed-files-size" @@ -94,6 +96,17 @@ class IsolationLevel(str, Enum): SNAPSHOT = "snapshot" +def _is_deletion_vector(data_file: DataFile) -> bool: + return data_file.content == DataFileContent.POSITION_DELETES and data_file.file_format == FileFormat.PUFFIN + + +def _content_size_in_bytes(data_file: DataFile) -> int: + """Return the size of a file's content; a deletion vector only accounts for its blob, as in Java.""" + if _is_deletion_vector(data_file) and data_file.content_size_in_bytes is not None: + return data_file.content_size_in_bytes + return data_file.file_size_in_bytes + + class UpdateMetrics: added_file_size: int removed_file_size: int @@ -103,6 +116,8 @@ class UpdateMetrics: removed_eq_delete_files: int added_pos_delete_files: int removed_pos_delete_files: int + added_dvs: int + removed_dvs: int added_delete_files: int removed_delete_files: int added_records: int @@ -121,6 +136,8 @@ def __init__(self) -> None: self.removed_eq_delete_files = 0 self.added_pos_delete_files = 0 self.removed_pos_delete_files = 0 + self.added_dvs = 0 + self.removed_dvs = 0 self.added_delete_files = 0 self.removed_delete_files = 0 self.added_records = 0 @@ -131,14 +148,17 @@ def __init__(self) -> None: self.removed_eq_deletes = 0 def add_file(self, data_file: DataFile) -> None: - self.added_file_size += data_file.file_size_in_bytes + self.added_file_size += _content_size_in_bytes(data_file) if data_file.content == DataFileContent.DATA: self.added_data_files += 1 self.added_records += data_file.record_count elif data_file.content == DataFileContent.POSITION_DELETES: self.added_delete_files += 1 - self.added_pos_delete_files += 1 + if _is_deletion_vector(data_file): + self.added_dvs += 1 + else: + self.added_pos_delete_files += 1 self.added_pos_deletes += data_file.record_count elif data_file.content == DataFileContent.EQUALITY_DELETES: self.added_delete_files += 1 @@ -148,14 +168,17 @@ def add_file(self, data_file: DataFile) -> None: raise ValueError(f"Unknown data file content: {data_file.content}") def remove_file(self, data_file: DataFile) -> None: - self.removed_file_size += data_file.file_size_in_bytes + self.removed_file_size += _content_size_in_bytes(data_file) if data_file.content == DataFileContent.DATA: self.removed_data_files += 1 self.deleted_records += data_file.record_count elif data_file.content == DataFileContent.POSITION_DELETES: self.removed_delete_files += 1 - self.removed_pos_delete_files += 1 + if _is_deletion_vector(data_file): + self.removed_dvs += 1 + else: + self.removed_pos_delete_files += 1 self.removed_pos_deletes += data_file.record_count elif data_file.content == DataFileContent.EQUALITY_DELETES: self.removed_delete_files += 1 @@ -174,6 +197,8 @@ def to_dict(self) -> dict[str, str]: set_when_positive(properties, self.removed_eq_delete_files, REMOVED_EQUALITY_DELETE_FILES) set_when_positive(properties, self.added_pos_delete_files, ADDED_POSITION_DELETE_FILES) set_when_positive(properties, self.removed_pos_delete_files, REMOVED_POSITION_DELETE_FILES) + set_when_positive(properties, self.added_dvs, ADDED_DVS) + set_when_positive(properties, self.removed_dvs, REMOVED_DVS) set_when_positive(properties, self.added_delete_files, ADDED_DELETE_FILES) set_when_positive(properties, self.removed_delete_files, REMOVED_DELETE_FILES) set_when_positive(properties, self.added_records, ADDED_RECORDS) diff --git a/pyiceberg/table/sorting.py b/pyiceberg/table/sorting.py index 76bc109935..f7b59649c6 100644 --- a/pyiceberg/table/sorting.py +++ b/pyiceberg/table/sorting.py @@ -29,7 +29,7 @@ from pyiceberg.exceptions import ValidationError from pyiceberg.schema import Schema -from pyiceberg.transforms import IdentityTransform, Transform, parse_transform +from pyiceberg.transforms import IdentityTransform, Transform, TransformSourceMixin, parse_transform from pyiceberg.typedef import IcebergBaseModel from pyiceberg.types import IcebergType @@ -60,11 +60,12 @@ def __repr__(self) -> str: return f"NullOrder.{self.name}" -class SortField(IcebergBaseModel): +class SortField(TransformSourceMixin): """Sort order field. + The source columns are carried by `TransformSourceMixin`. + Args: - source_id (int): Source column id from the table’s schema. transform (str): Transform that is used to produce values to be sorted on from the source column. This is the same transform as described in partition transforms. direction (SortDirection): Sort direction, that can only be either asc or desc. @@ -97,23 +98,6 @@ def set_null_order(cls, values: dict[str, Any]) -> dict[str, Any]: values["null-order"] = NullOrder.NULLS_FIRST if values["direction"] == SortDirection.ASC else NullOrder.NULLS_LAST return values - @model_validator(mode="before") - @classmethod - def map_source_ids_onto_source_id(cls, data: Any) -> Any: - if isinstance(data, dict): - if "source-ids" in data: - if "source-id" in data: - raise ValueError("source-id and source-ids are mutually exclusive") - source_ids = data["source-ids"] - if isinstance(source_ids, list): - if len(source_ids) == 0: - raise ValueError("Empty source-ids is not allowed") - if len(source_ids) > 1: - raise ValueError("Multi argument transforms are not yet supported") - data["source-id"] = source_ids[0] - return data - - source_id: int = Field(alias="source-id") transform: Annotated[ # type: ignore Transform, BeforeValidator(parse_transform), @@ -128,8 +112,8 @@ def __str__(self) -> str: if isinstance(self.transform, IdentityTransform): # In the case of an identity transform, we can omit the transform return f"{self.source_id} {self.direction} {self.null_order}" - else: - return f"{self.transform}({self.source_id}) {self.direction} {self.null_order}" + sources = ", ".join(str(source_id) for source_id in self.transform_arguments) + return f"{self.transform}({sources}) {self.direction} {self.null_order}" INITIAL_SORT_ORDER_ID = 1 diff --git a/pyiceberg/table/update/__init__.py b/pyiceberg/table/update/__init__.py index 64838b0bd6..e98b7093c9 100644 --- a/pyiceberg/table/update/__init__.py +++ b/pyiceberg/table/update/__init__.py @@ -29,7 +29,7 @@ from pyiceberg.exceptions import CommitFailedException from pyiceberg.partitioning import PARTITION_FIELD_ID_START, PartitionSpec from pyiceberg.schema import Schema -from pyiceberg.table.metadata import SUPPORTED_TABLE_FORMAT_VERSION, TableMetadata, TableMetadataUtil +from pyiceberg.table.metadata import INITIAL_ROW_ID, SUPPORTED_TABLE_FORMAT_VERSION, TableMetadata, TableMetadataUtil from pyiceberg.table.refs import MAIN_BRANCH, SnapshotRef, SnapshotRefType from pyiceberg.table.snapshots import ( MetadataLogEntry, @@ -321,10 +321,15 @@ def _( elif update.format_version == base_metadata.format_version: return base_metadata - updated_metadata = base_metadata.model_copy(update={"format_version": update.format_version}) + updated_metadata = TableMetadataUtil._construct_without_validation( + base_metadata.model_copy(update={"format_version": update.format_version}) + ) + if update.format_version >= 3 and base_metadata.format_version < 3: + # Rows in snapshots committed before the upgrade have no row ids, so row id assignment starts at 0 + updated_metadata = updated_metadata.model_copy(update={"next_row_id": INITIAL_ROW_ID}) context.add_update(update) - return TableMetadataUtil._construct_without_validation(updated_metadata) + return updated_metadata @_apply_table_update.register(SetPropertiesUpdate) @@ -354,6 +359,8 @@ def _(update: RemovePropertiesUpdate, base_metadata: TableMetadata, context: _Ta @_apply_table_update.register(AddSchemaUpdate) def _(update: AddSchemaUpdate, base_metadata: TableMetadata, context: _TableMetadataUpdateContext) -> TableMetadata: + update.schema_.check_format_version_compatibility(base_metadata.format_version) + metadata_updates: dict[str, Any] = { "last_column_id": max(base_metadata.last_column_id, update.schema_.highest_field_id), "schemas": base_metadata.schemas + [update.schema_], @@ -435,7 +442,7 @@ def _(update: AddSnapshotUpdate, base_metadata: TableMetadata, context: _TableMe elif base_metadata.snapshot_by_id(update.snapshot.snapshot_id) is not None: raise ValueError(f"Snapshot with id {update.snapshot.snapshot_id} already exists") elif ( - base_metadata.format_version == 2 + base_metadata.format_version >= 2 and update.snapshot.sequence_number is not None and update.snapshot.sequence_number <= base_metadata.last_sequence_number and update.snapshot.parent_snapshot_id is not None @@ -444,32 +451,27 @@ def _(update: AddSnapshotUpdate, base_metadata: TableMetadata, context: _TableMe f"Cannot add snapshot with sequence number {update.snapshot.sequence_number} " f"older than last sequence number {base_metadata.last_sequence_number}" ) - elif base_metadata.format_version >= 3 and update.snapshot.first_row_id is None: - raise ValueError("Cannot add snapshot without first row id") - elif ( - base_metadata.format_version >= 3 - and update.snapshot.first_row_id is not None - and base_metadata.next_row_id is not None - and update.snapshot.first_row_id < base_metadata.next_row_id - ): - raise ValueError( - f"Cannot add a snapshot with first row id smaller than the table's next-row-id " - f"{update.snapshot.first_row_id} < {base_metadata.next_row_id}" - ) + + metadata_updates: dict[str, Any] = { + "last_updated_ms": update.snapshot.timestamp_ms, + "last_sequence_number": update.snapshot.sequence_number, + "snapshots": base_metadata.snapshots + [update.snapshot], + } + + if base_metadata.format_version >= 3: + if update.snapshot.first_row_id is None: + raise ValueError("Cannot add snapshot without first row id") + if update.snapshot.first_row_id < base_metadata.next_row_id: + raise ValueError( + f"Cannot add a snapshot with first row id smaller than the table's next-row-id " + f"{update.snapshot.first_row_id} < {base_metadata.next_row_id}" + ) + if update.snapshot.added_rows is None: + raise ValueError("Cannot add snapshot without added rows") + metadata_updates["next_row_id"] = base_metadata.next_row_id + update.snapshot.added_rows context.add_update(update) - return base_metadata.model_copy( - update={ - "last_updated_ms": update.snapshot.timestamp_ms, - "last_sequence_number": update.snapshot.sequence_number, - "snapshots": base_metadata.snapshots + [update.snapshot], - "next_row_id": base_metadata.next_row_id + update.snapshot.added_rows - if base_metadata.format_version >= 3 - and base_metadata.next_row_id is not None - and update.snapshot.added_rows is not None - else None, - } - ) + return base_metadata.model_copy(update=metadata_updates) @_apply_table_update.register(SetSnapshotRefUpdate) diff --git a/pyiceberg/table/update/schema.py b/pyiceberg/table/update/schema.py index 828f1e877a..2d918fe9ca 100644 --- a/pyiceberg/table/update/schema.py +++ b/pyiceberg/table/update/schema.py @@ -49,8 +49,22 @@ UpdatesAndRequirements, UpdateTableMetadata, ) +from pyiceberg.transforms import DayTransform, MonthTransform, VoidTransform, YearTransform from pyiceberg.typedef import L, TableVersion -from pyiceberg.types import IcebergType, ListType, MapType, NestedField, PrimitiveType, StructType +from pyiceberg.types import ( + DateType, + GeographyType, + GeometryType, + IcebergType, + ListType, + MapType, + NestedField, + PrimitiveType, + StructType, + TimestampNanoType, + TimestampType, + UnknownType, +) if TYPE_CHECKING: import pyarrow as pa @@ -144,13 +158,24 @@ def case_sensitive(self, case_sensitive: bool) -> UpdateSchema: return self def union_by_name( - # TODO: Move TableProperties.DEFAULT_FORMAT_VERSION to separate file and set that as format_version default. self, new_schema: Schema | pa.Schema, - format_version: TableVersion = 2, + format_version: TableVersion | None = None, ) -> UpdateSchema: + """Merge a new schema into the current schema by matching fields by name. + + Args: + new_schema: The schema to merge in. + format_version: Format version used to convert a PyArrow schema; defaults to the table's format version. + + Returns: + This for method chaining. + """ from pyiceberg.catalog import Catalog + if format_version is None: + format_version = self._transaction.table_metadata.format_version + visit_with_partner( Catalog._convert_schema_if_needed(new_schema, format_version=format_version), -1, @@ -181,7 +206,8 @@ def add_column( field_type: Type for the new column. doc: Documentation string for the new column. required: Whether the new column is required. - default_value: Default value for the new column. + default_value: Default value for the new column. Requires format version 3, unless incompatible + changes are allowed. Returns: This for method chaining. @@ -224,6 +250,9 @@ def add_column( new_id = self.assign_new_column_id() new_type = assign_fresh_schema_ids(field_type, self.assign_new_column_id) + if default_value is not None and isinstance(new_type, (GeometryType, GeographyType, UnknownType)): + raise ValueError(f"Invalid default value: columns of type {new_type} must default to null: {full_name}") + if default_value is not None: try: # To make sure that the value is valid for the type @@ -233,6 +262,13 @@ def add_column( else: initial_default = default_value # type: ignore + format_version = self._transaction.table_metadata.format_version + if initial_default is not None and format_version < 3 and not self._allow_incompatible_changes: + raise ValueError( + f"Incompatible change: default values require format version 3 or higher, " + f"current format version is {format_version}: {full_name}" + ) + if (required and initial_default is None) and not self._allow_incompatible_changes: # Table format version 1 and 2 cannot add required column because there is no initial value raise ValueError(f"Incompatible change: cannot add required column: {'.'.join(path)}") @@ -398,6 +434,16 @@ def _set_column_default_value(self, path: str | tuple[str, ...], default_value: field = self._schema.find_field(name, self._case_sensitive) + if default_value is not None and isinstance(field.field_type, (GeometryType, GeographyType, UnknownType)): + raise ValueError(f"Invalid default value: columns of type {field.field_type} must default to null: {name}") + + format_version = self._transaction.table_metadata.format_version + if default_value is not None and format_version < 3 and not self._allow_incompatible_changes: + raise ValueError( + f"Incompatible change: default values require format version 3 or higher, " + f"current format version is {format_version}: {name}" + ) + if default_value is not None: try: # To make sure that the value is valid for the type @@ -470,10 +516,14 @@ def update_column( raise ValidationError(f"Cannot change column type: {field.field_type} is not a primitive") if not self._allow_incompatible_changes and field.field_type != field_type: + if isinstance(field_type, (GeometryType, GeographyType)): + # promote() accepts binary -> geometry/geography only to resolve file schemas on read + raise ValidationError(f"Cannot change column type: {full_name}: {field.field_type} -> {field_type}") try: promote(field.field_type, field_type) except ResolveError as e: raise ValidationError(f"Cannot change column type: {full_name}: {field.field_type} -> {field_type}") from e + self._validate_v3_promotion(full_name, field, field_type) # if other updates for the same field exist in one transaction: if updated := self._updates.get(field.field_id): @@ -502,6 +552,32 @@ def update_column( return self + def _validate_v3_promotion(self, full_name: str, field: NestedField, field_type: IcebergType) -> None: + """Validate the type promotions that the spec only allows from format version 3.""" + is_unknown_promotion = isinstance(field.field_type, UnknownType) + is_date_promotion = isinstance(field.field_type, DateType) and isinstance(field_type, (TimestampType, TimestampNanoType)) + if not (is_unknown_promotion or is_date_promotion): + return + + format_version = self._transaction.table_metadata.format_version + if format_version < 3: + raise ValidationError( + f"Cannot change column type: {full_name}: {field.field_type} -> {field_type} " + f"requires format version 3 or higher, current format version is {format_version}" + ) + + if is_date_promotion: + # Only allowed when every partition transform on the column produces the same value for the new type + for spec in self._transaction.table_metadata.partition_specs: + for partition_field in spec.fields: + if partition_field.source_id == field.field_id and not isinstance( + partition_field.transform, (YearTransform, MonthTransform, DayTransform, VoidTransform) + ): + raise ValidationError( + f"Cannot change column type: {full_name}: {field.field_type} -> {field_type}, " + f"the column is partitioned by {partition_field.transform}" + ) + def _find_for_move(self, name: str) -> int | None: try: return self._schema.find_field(name, self._case_sensitive).field_id diff --git a/pyiceberg/table/update/snapshot.py b/pyiceberg/table/update/snapshot.py index 8024e808b2..c1a06ea750 100644 --- a/pyiceberg/table/update/snapshot.py +++ b/pyiceberg/table/update/snapshot.py @@ -25,6 +25,7 @@ from dataclasses import dataclass from datetime import datetime from functools import cached_property +from types import TracebackType from typing import TYPE_CHECKING, Generic from pyiceberg.avro.codecs import AvroCompressionCodec @@ -47,16 +48,19 @@ from pyiceberg.manifest import ( DataFile, DataFileContent, + FileFormat, ManifestContent, ManifestEntry, ManifestEntryStatus, ManifestFile, + ManifestListWriterV3, ManifestWriter, write_manifest, write_manifest_list, ) from pyiceberg.partitioning import PartitionSpec from pyiceberg.schema import Schema +from pyiceberg.table.delete_file import DeleteFileSet from pyiceberg.table.refs import MAIN_BRANCH, SnapshotRef, SnapshotRefType from pyiceberg.table.snapshots import ( Operation, @@ -141,6 +145,7 @@ class _SnapshotProducer(UpdateTableMetadata[U], Generic[U]): _parent_snapshot_id: int | None _starting_snapshot_id: int | None _added_data_files: list[DataFile] + _written_data_files: list[DataFile] _manifest_num_counter: itertools.count[int] _deleted_data_files: set[DataFile] _compression: AvroCompressionCodec @@ -168,6 +173,7 @@ def __init__( self._operation = operation self._snapshot_id = self._transaction.table_metadata.new_snapshot_id() self._added_data_files = [] + self._written_data_files = [] self._deleted_data_files = set() self.snapshot_properties = snapshot_properties self._manifest_num_counter = itertools.count(0) @@ -209,18 +215,19 @@ def delete_data_file(self, data_file: DataFile) -> _SnapshotProducer[U]: self._deleted_data_files.add(data_file) return self - def _calculate_added_rows(self, manifests: list[ManifestFile]) -> int: - """Calculate the number of added rows from a list of manifest files.""" - added_rows = 0 - for manifest in manifests: - if manifest.added_snapshot_id is None or manifest.added_snapshot_id == self._snapshot_id: - if manifest.added_rows_count is None: - raise ValueError( - "Cannot determine number of added rows in snapshot because " - f"the entry for manifest {manifest.manifest_path} is missing the field `added-rows-count`" - ) - added_rows += manifest.added_rows_count - return added_rows + def _append_written_data_file(self, data_file: DataFile) -> _SnapshotProducer[U]: + """Append a data file that was written for this operation, so it is removed again if the operation fails.""" + self._written_data_files.append(data_file) + return self.append_data_file(data_file) + + def _delete_written_data_files(self) -> None: + """Best-effort removal of the data files written for this operation.""" + for data_file in self._written_data_files: + try: + self._io.delete(data_file.file_path) + except Exception: + logger.warning("Failed to delete uncommitted data file: %s", data_file.file_path, exc_info=True) + self._written_data_files.clear() @abstractmethod def _deleted_entries(self) -> list[ManifestEntry]: ... @@ -311,6 +318,8 @@ def _summary(self, snapshot_properties: dict[str, str] = EMPTY_DICT) -> Summary: schema=schema, ) + self._update_summary(ssc) + previous_snapshot = ( table_metadata.snapshot_by_id(self._parent_snapshot_id) if self._parent_snapshot_id is not None else None ) @@ -320,6 +329,9 @@ def _summary(self, snapshot_properties: dict[str, str] = EMPTY_DICT) -> Summary: previous_summary=previous_snapshot.summary if previous_snapshot is not None else None, ) + def _update_summary(self, collector: SnapshotSummaryCollector) -> None: + """Record additional file changes of this operation in the snapshot summary.""" + def _commit(self) -> UpdatesAndRequirements: new_manifests = self._manifests() next_sequence_number = self._transaction.table_metadata.next_sequence_number() @@ -334,20 +346,26 @@ def _commit(self) -> UpdatesAndRequirements: manifest_list_file_path = location_provider.new_metadata_location(file_name) self._written_manifest_lists.append(manifest_list_file_path) + table_metadata = self._transaction.table_metadata + first_row_id: int | None = None + added_rows: int | None = None + if table_metadata.format_version >= 3: + first_row_id = table_metadata.next_row_id + with write_manifest_list( - format_version=self._transaction.table_metadata.format_version, + format_version=table_metadata.format_version, output_file=self._io.new_output(manifest_list_file_path), snapshot_id=self._snapshot_id, parent_snapshot_id=self._parent_snapshot_id, sequence_number=next_sequence_number, avro_compression=self._compression, + first_row_id=first_row_id, ) as writer: writer.add_manifests(new_manifests) - first_row_id: int | None = None - - if self._transaction.table_metadata.format_version >= 3: - first_row_id = self._transaction.table_metadata.next_row_id + if isinstance(writer, ManifestListWriterV3) and first_row_id is not None: + # Row ids assigned to the manifests in this snapshot, including pre-upgrade manifests without one + added_rows = writer.next_row_id - first_row_id snapshot = Snapshot( snapshot_id=self._snapshot_id, @@ -355,8 +373,9 @@ def _commit(self) -> UpdatesAndRequirements: manifest_list=manifest_list_file_path, sequence_number=next_sequence_number, summary=summary, - schema_id=self._transaction.table_metadata.current_schema_id, + schema_id=table_metadata.current_schema_id, first_row_id=first_row_id, + added_rows=added_rows, ) add_snapshot_update = AddSnapshotUpdate(snapshot=snapshot) @@ -397,7 +416,7 @@ def schema(self) -> Schema: def spec(self, spec_id: int) -> PartitionSpec: return self._transaction.table_metadata.specs()[spec_id] - def new_manifest_writer(self, spec: PartitionSpec) -> ManifestWriter: + def new_manifest_writer(self, spec: PartitionSpec, content: ManifestContent = ManifestContent.DATA) -> ManifestWriter: return write_manifest( format_version=self._transaction.table_metadata.format_version, spec=spec, @@ -405,6 +424,7 @@ def new_manifest_writer(self, spec: PartitionSpec) -> ManifestWriter: output_file=self.new_manifest_output(), snapshot_id=self._snapshot_id, avro_compression=self._compression, + content=content, ) def new_manifest_output(self) -> OutputFile: @@ -418,8 +438,20 @@ def fetch_manifest_entry(self, manifest: ManifestFile, discard_deleted: bool = T return manifest.fetch_manifest_entry(io=self._io, discard_deleted=discard_deleted) def commit(self) -> None: + try: + updates, requirements = self._commit() + except Exception: + # Nothing references the files written so far, so remove them instead of leaving orphans + self._clean_all_uncommitted() + raise self._transaction._register_snapshot_producer(self) - super().commit() + self._transaction._apply(updates, requirements) + + def __exit__(self, exctype: type[BaseException] | None, excinst: BaseException | None, exctb: TracebackType | None) -> None: + """Commit the snapshot, or remove the files written for it if an exception has been raised.""" + if excinst is not None: + self._clean_all_uncommitted() + super().__exit__(exctype, excinst, exctb) def _cleanup_uncommitted(self) -> None: """Delete manifest files and manifest lists from failed retry attempts.""" @@ -439,7 +471,8 @@ def _cleanup_uncommitted(self) -> None: self._written_manifest_lists = self._written_manifest_lists[-1:] def _clean_all_uncommitted(self) -> None: - """Clean up all manifests and manifest lists on abort.""" + """Clean up all written data files, manifests and manifest lists on abort.""" + self._delete_written_data_files() for path in itertools.chain(self._uncommitted_manifests, self._written_manifests): try: self._io.delete(path) @@ -882,6 +915,171 @@ def _validate_required_deletes(self, deleted_entries: list[ManifestEntry]) -> No raise ValidationException(f"Missing required files to delete: {', '.join(sorted(missing))}") +class _RowDelta(_SnapshotProducer["_RowDelta"]): + """Adds delete files (deletion vectors) and removes the delete files they replace. + + Produces a DELETE snapshot when only delete files are added, and an OVERWRITE snapshot when + data files change as well. Existing data and delete manifests are carried forward; delete + manifests holding a replaced delete file are rewritten with that entry marked as DELETED. + """ + + _added_delete_files: list[DataFile] + _removed_delete_files: DeleteFileSet + _referenced_data_files: set[DataFile] + _written_delete_file_paths: set[str] + + def __init__( + self, + operation: Operation, + transaction: Transaction, + io: FileIO, + commit_uuid: uuid.UUID | None = None, + snapshot_properties: dict[str, str] = EMPTY_DICT, + branch: str | None = MAIN_BRANCH, + ) -> None: + super().__init__(operation, transaction, io, commit_uuid, snapshot_properties, branch) + self._added_delete_files = [] + self._removed_delete_files = DeleteFileSet() + self._referenced_data_files = set() + self._written_delete_file_paths = set() + + def add_delete_file(self, delete_file: DataFile) -> _RowDelta: + """Add a delete file to the table.""" + if delete_file.content == DataFileContent.DATA: + raise ValueError(f"Expected a delete file, got a data file: {delete_file.file_path}") + if self._transaction.table_metadata.format_version >= 3 and delete_file.file_format != FileFormat.PUFFIN: + raise ValueError(f"Only deletion vectors can be added to format version 3 tables: {delete_file.file_path}") + self._added_delete_files.append(delete_file) + return self + + def _add_written_delete_file(self, delete_file: DataFile) -> _RowDelta: + """Add a delete file written for this operation, so it is removed again if the operation fails.""" + self._written_delete_file_paths.add(delete_file.file_path) + return self.add_delete_file(delete_file) + + def remove_delete_file(self, delete_file: DataFile) -> _RowDelta: + """Remove a delete file that is replaced by this operation, such as a superseded deletion vector.""" + self._removed_delete_files.add(delete_file) + return self + + def validate_data_files_exist(self, data_files: set[DataFile]) -> _RowDelta: + """Require that the given data files are neither removed nor receive new deletes concurrently.""" + self._referenced_data_files |= data_files + return self + + def _delete_written_data_files(self) -> None: + super()._delete_written_data_files() + for path in self._written_delete_file_paths: + try: + self._io.delete(path) + except Exception: + logger.warning("Failed to delete uncommitted delete file: %s", path, exc_info=True) + self._written_delete_file_paths.clear() + + @property + def _has_changes(self) -> bool: + return bool(self._added_data_files or self._deleted_data_files or self._added_delete_files or self._removed_delete_files) + + def _commit(self) -> UpdatesAndRequirements: + if not self._has_changes: + return (), () + # Mirrors Java's BaseRowDelta: only adding deletes is a DELETE, anything else an OVERWRITE + if self._added_delete_files and not self._added_data_files and not self._deleted_data_files: + self._operation = Operation.DELETE + else: + self._operation = Operation.OVERWRITE + return super()._commit() + + def _update_summary(self, collector: SnapshotSummaryCollector) -> None: + schema = self.schema() + specs = self._transaction.table_metadata.specs() + for delete_file in self._added_delete_files: + collector.add_file(delete_file, schema=schema, partition_spec=specs[delete_file.spec_id]) + for delete_file in self._removed_delete_files: + collector.remove_file(delete_file, schema=schema, partition_spec=specs[delete_file.spec_id]) + + def _manifests(self) -> list[ManifestFile]: + return self._write_added_delete_manifests() + super()._manifests() + + def _write_added_delete_manifests(self) -> list[ManifestFile]: + by_spec: dict[int, list[DataFile]] = defaultdict(list) + for delete_file in self._added_delete_files: + by_spec[delete_file.spec_id].append(delete_file) + + manifests = [] + for spec_id, delete_files in by_spec.items(): + with self.new_manifest_writer(self.spec(spec_id), content=ManifestContent.DELETES) as writer: + for delete_file in delete_files: + writer.add( + ManifestEntry.from_args( + status=ManifestEntryStatus.ADDED, + snapshot_id=self._snapshot_id, + sequence_number=None, + file_sequence_number=None, + data_file=delete_file, + ) + ) + manifests.append(writer.to_manifest_file()) + return manifests + + def _existing_manifests(self) -> list[ManifestFile]: + """Carry the parent's manifests forward, rewriting delete manifests that hold a removed delete file.""" + if self._parent_snapshot_id is None: + if self._removed_delete_files: + raise ValidationException("Cannot remove delete files from a table without snapshots") + return [] + + parent = self._transaction.table_metadata.snapshot_by_id(self._parent_snapshot_id) + if parent is None: + raise ValueError(f"Snapshot could not be found: {self._parent_snapshot_id}") + + existing_manifests = [] + found = DeleteFileSet() + for manifest in parent.manifests(io=self._io): + if manifest.content != ManifestContent.DELETES or not self._removed_delete_files: + existing_manifests.append(manifest) + continue + + entries = manifest.fetch_manifest_entry(io=self._io, discard_deleted=True) + if not any(entry.data_file in self._removed_delete_files for entry in entries): + existing_manifests.append(manifest) + continue + + with self.new_manifest_writer(self.spec(manifest.partition_spec_id), content=ManifestContent.DELETES) as writer: + for entry in entries: + if entry.data_file in self._removed_delete_files: + found.add(entry.data_file) + writer.delete(entry) + else: + writer.existing(entry) + existing_manifests.append(writer.to_manifest_file()) + + if missing := [delete_file.file_path for delete_file in self._removed_delete_files if delete_file not in found]: + raise ValidationException(f"Missing required delete files to remove: {', '.join(sorted(missing))}") + + return existing_manifests + + def _deleted_entries(self) -> list[ManifestEntry]: + """Return no separate entries; removed delete files are recorded in their rewritten delete manifests.""" + return [] + + def _validate_concurrency(self) -> None: + """Run the base validation, and check the referenced data files still exist and have no new deletes.""" + from pyiceberg.table.update.validate import _validate_data_files_exist, _validate_no_new_deletes_for_data_files + + super()._validate_concurrency() + + if self._commit_window is None or self._commit_window.is_empty() or self._commit_window.head is None: + return + + if self._referenced_data_files: + table = self._transaction._table + head = self._commit_window.head + base = self._commit_window.base + _validate_data_files_exist(table, head, self._referenced_data_files, base) + _validate_no_new_deletes_for_data_files(table, head, None, self._referenced_data_files, base) + + class UpdateSnapshot: _transaction: Transaction _io: FileIO @@ -939,6 +1137,16 @@ def delete(self) -> _DeleteFiles: snapshot_properties=self._snapshot_properties, ) + def row_delta(self, commit_uuid: uuid.UUID | None = None) -> _RowDelta: + return _RowDelta( + commit_uuid=commit_uuid, + operation=Operation.DELETE, + transaction=self._transaction, + io=self._io, + branch=self._branch, + snapshot_properties=self._snapshot_properties, + ) + class _ManifestMergeManager(Generic[U]): _target_size_bytes: int diff --git a/pyiceberg/transforms.py b/pyiceberg/transforms.py index 54c01d9bed..89bf9b0384 100644 --- a/pyiceberg/transforms.py +++ b/pyiceberg/transforms.py @@ -17,6 +17,7 @@ import base64 import datetime as py_datetime +import re import struct from abc import ABC, abstractmethod from collections.abc import Callable @@ -27,7 +28,7 @@ from uuid import UUID import mmh3 -from pydantic import Field, PositiveInt, PrivateAttr +from pydantic import Field, PositiveInt, PrivateAttr, model_serializer, model_validator from pyiceberg.expressions import ( BoundEqualTo, @@ -62,9 +63,11 @@ Literal, LongLiteral, TimestampLiteral, + TimestampNanoLiteral, + TimestamptzNanoLiteral, literal, ) -from pyiceberg.typedef import IcebergRootModel, L +from pyiceberg.typedef import IcebergBaseModel, IcebergRootModel, L from pyiceberg.types import ( BinaryType, DateType, @@ -82,6 +85,7 @@ TimestamptzType, TimeType, UUIDType, + VariantType, ) from pyiceberg.utils import datetime from pyiceberg.utils.decimal import decimal_to_bytes, truncate_decimal @@ -107,6 +111,8 @@ HOUR = "hour" BUCKET_PARSER = ParseNumberFromBrackets(BUCKET) +# name[argument] with nothing after the closing bracket +_BRACKET_FORM = re.compile(r"^[a-z_]+\[[^\]]*\]$") TRUNCATE_PARSER = ParseNumberFromBrackets(TRUNCATE) @@ -212,9 +218,11 @@ def parse_transform(v: Any) -> Transform[Any, Any]: return IdentityTransform() elif v == VOID: return VoidTransform() - elif v.startswith(BUCKET): + elif _BRACKET_FORM.match(v) and v.startswith(f"{BUCKET}["): + # A known transform with a malformed argument is an error; an unknown name that merely shares the + # prefix (bucketv2[4], bucket[4]v2) is a different transform and is preserved as-is below return BucketTransform(num_buckets=BUCKET_PARSER.match(v)) - elif v.startswith(TRUNCATE): + elif _BRACKET_FORM.match(v) and v.startswith(f"{TRUNCATE}["): return TruncateTransform(width=TRUNCATE_PARSER.match(v)) elif v == YEAR: return YearTransform() @@ -472,8 +480,10 @@ def year_func(v: Any) -> int: elif isinstance(source, (TimestampNanoType, TimestamptzNanoType)): def year_func(v: Any) -> int: - # python datetime has no nanoseconds support. - # nanosecond datetimes will be expressed as int as a workaround + # python datetime has no nanoseconds support, so datetimes are converted with microsecond precision + if isinstance(v, py_datetime.datetime): + v = datetime.datetime_to_nanos(v) + return datetime.nanos_to_years(v) else: @@ -532,8 +542,10 @@ def month_func(v: Any) -> int: elif isinstance(source, (TimestampNanoType, TimestamptzNanoType)): def month_func(v: Any) -> int: - # python datetime has no nanoseconds support. - # nanosecond datetimes will be expressed as int as a workaround + # python datetime has no nanoseconds support, so datetimes are converted with microsecond precision + if isinstance(v, py_datetime.datetime): + v = datetime.datetime_to_nanos(v) + return datetime.nanos_to_months(v) else: @@ -593,8 +605,10 @@ def day_func(v: Any) -> int: elif isinstance(source, (TimestampNanoType, TimestamptzNanoType)): def day_func(v: Any) -> int: - # python datetime has no nanoseconds support. - # nanosecond datetimes will be expressed as int as a workaround + # python datetime has no nanoseconds support, so datetimes are converted with microsecond precision + if isinstance(v, py_datetime.datetime): + v = datetime.datetime_to_nanos(v) + return datetime.nanos_to_days(v) else: @@ -654,8 +668,10 @@ def hour_func(v: Any) -> int: elif isinstance(source, (TimestampNanoType, TimestamptzNanoType)): def hour_func(v: Any) -> int: - # python datetime has no nanoseconds support. - # nanosecond datetimes will be expressed as int as a workaround + # python datetime has no nanoseconds support, so datetimes are converted with microsecond precision + if isinstance(v, py_datetime.datetime): + v = datetime.datetime_to_nanos(v) + return datetime.nanos_to_hours(v) else: @@ -706,7 +722,7 @@ def transform(self, source: IcebergType) -> Callable[[S | None], S | None]: return lambda v: v def can_transform(self, source: IcebergType) -> bool: - return source.is_primitive and not isinstance(source, (GeographyType, GeometryType)) + return source.is_primitive and not isinstance(source, (GeographyType, GeometryType, VariantType)) def result_type(self, source: IcebergType) -> IcebergType: return source @@ -955,6 +971,16 @@ def _(_type: IcebergType, value: int) -> str: return datetime.to_human_timestamptz(value) +@_int_to_human_string.register(TimestampNanoType) +def _(_type: IcebergType, value: int) -> str: + return datetime.to_human_timestamp_ns(value) + + +@_int_to_human_string.register(TimestamptzNanoType) +def _(_type: IcebergType, value: int) -> str: + return datetime.to_human_timestamptz_ns(value) + + class UnknownTransform(Transform[S, T]): """A transform that represents when an unknown transform is provided. @@ -987,10 +1013,24 @@ def project(self, name: str, pred: BoundPredicate) -> UnboundPredicate | None: def strict_project(self, name: str, pred: BoundPredicate) -> UnboundPredicate | None: return None + def __str__(self) -> str: + """Return the original transform name so it round-trips through serialization.""" + return self._transform + def __repr__(self) -> str: """Return the string representation of the UnknownTransform class.""" return f"UnknownTransform(transform={repr(self._transform)})" + def __eq__(self, other: Any) -> bool: + """Compare the preserved transform name, since every unknown transform shares one root value.""" + if isinstance(other, UnknownTransform): + return self._transform == other._transform + return False + + def __hash__(self) -> int: + """Hash the preserved transform name so distinct unknown transforms do not collide.""" + return hash((self.root, self._transform)) + def pyarrow_transform(self, source: IcebergType) -> "Callable[[pa.Array], pa.Array]": raise NotImplementedError() @@ -1003,8 +1043,9 @@ class VoidTransform(Transform[S, None], Singleton): def transform(self, source: IcebergType) -> Callable[[S | None], T | None]: return lambda v: None - def can_transform(self, _: IcebergType) -> bool: - return True + def can_transform(self, source: IcebergType) -> bool: + # Variant cannot be the source of any partition or sort field + return not isinstance(source, VariantType) def result_type(self, source: IcebergType) -> IcebergType: return source @@ -1031,12 +1072,82 @@ def pyarrow_transform(self, source: IcebergType) -> "Callable[[pa.Array], pa.Arr return lambda arr: pa.nulls(len(arr), type=arr.type) +class TransformSourceMixin(IcebergBaseModel): + """Shared `source-id` and `source-ids` handling for fields that apply a transform to source columns. + + Both partition fields and sort fields carry this pair, and the spec writes only one of the + two: `source-id` for a transform with a single argument, `source-ids` for a multi-argument + transform. This mixin owns reading, serializing and reporting them, so that neither field + type has to reach into the raw keys. + + Attributes: + source_id(int): The source column id of the table's schema. + source_ids(list[int] | None): The source column ids of a multi-argument transform. + """ + + source_id: int = Field(alias="source-id") + source_ids: list[int] | None = Field(alias="source-ids", default=None, repr=False) + + @property + def transform_arguments(self) -> list[int]: + """Return the source column ids that the transform is applied to.""" + source_ids = self.source_ids + if source_ids is not None and len(source_ids) > 1: + return list(source_ids) + return [self.source_id] + + @property + def is_multi_argument(self) -> bool: + """Return True if the transform takes more than one source column.""" + return len(self.transform_arguments) > 1 + + @model_validator(mode="before") + @classmethod + def map_source_ids_onto_source_id(cls, data: Any) -> Any: + if not isinstance(data, dict) or "source-ids" not in data: + return data + + if "source-id" in data: + raise ValueError("source-id and source-ids are mutually exclusive") + + source_ids = data["source-ids"] + if not isinstance(source_ids, list): + return data + if len(source_ids) == 0: + raise ValueError("Empty source-ids is not allowed") + + data["source-id"] = source_ids[0] + if len(source_ids) == 1: + data.pop("source-ids", None) + return data + + if data.get("transform") is None: + raise ValueError("Transform is required for a multi-argument field") + # Multi-argument transforms cannot be evaluated; per the spec, v3 readers + # must read tables with such transforms, ignoring them + data["transform"] = UnknownTransform(transform=str(data["transform"])) + return data + + @model_serializer(mode="wrap") + def _serialize_source_ids(self, handler: Any) -> Any: + serialized = handler(self) + # Per the spec, single-argument transforms write only source-id and + # multi-argument transforms write only source-ids + if self.is_multi_argument: + serialized.pop("source-id", None) + else: + serialized.pop("source-ids", None) + return serialized + + def _truncate_number( name: str, pred: BoundLiteralPredicate, transform: Callable[[Any | None], Any | None] ) -> UnboundPredicate | None: boundary = pred.literal - if not isinstance(boundary, (LongLiteral, DecimalLiteral, DateLiteral, TimestampLiteral)): + if not isinstance( + boundary, (LongLiteral, DecimalLiteral, DateLiteral, TimestampLiteral, TimestampNanoLiteral, TimestamptzNanoLiteral) + ): raise ValueError(f"Expected a numeric literal, got: {type(boundary)}") if isinstance(pred, BoundLessThan): @@ -1058,7 +1169,9 @@ def _truncate_number_strict( ) -> UnboundPredicate | None: boundary = pred.literal - if not isinstance(boundary, (LongLiteral, DecimalLiteral, DateLiteral, TimestampLiteral)): + if not isinstance( + boundary, (LongLiteral, DecimalLiteral, DateLiteral, TimestampLiteral, TimestampNanoLiteral, TimestamptzNanoLiteral) + ): raise ValueError(f"Expected a numeric literal, got: {type(boundary)}") if isinstance(pred, BoundLessThan): diff --git a/pyiceberg/types.py b/pyiceberg/types.py index 3295beb666..bec35747b0 100644 --- a/pyiceberg/types.py +++ b/pyiceberg/types.py @@ -66,11 +66,16 @@ DEFAULT_GEOGRAPHY_CRS = "OGC:CRS84" DEFAULT_GEOGRAPHY_ALGORITHM = "spherical" -# Regex patterns for parsing geometry and geography type strings -# Matches: geometry, geometry('CRS'), geometry('crs'), geometry("CRS") -GEOMETRY_REGEX = re.compile(r"geometry(?:\(\s*['\"]([^'\"]+)['\"]\s*\))?$") -# Matches: geography, geography('CRS'), geography('crs', 'algo') -GEOGRAPHY_REGEX = re.compile(r"geography(?:\(\s*['\"]([^'\"]+)['\"](?:\s*,\s*['\"]([^'\"]+)['\"])?\s*\))?$") +GEOGRAPHY_ALGORITHMS = frozenset({"spherical", "vincenty", "thomas", "andoyer", "karney"}) + +# Regex patterns for parsing geometry and geography type strings, compatible with the Java implementation. +# Values are unquoted in the spec form (geometry(srid:3857), geography(srid:4326, vincenty)); single- or +# double-quoted values are also accepted because older PyIceberg versions wrote them that way. +GEOMETRY_REGEX = re.compile(r"""^geometry\s*(?:\(\s*(?:'([^']*)'|"([^"]*)"|([^)]*?))\s*\))?$""", re.IGNORECASE) +GEOGRAPHY_REGEX = re.compile( + r"""^geography\s*(?:\(\s*(?:'([^']*)'|"([^"]*)"|([^,)]*?))\s*(?:,\s*(?:'(\w*)'|"(\w*)"|(\w*))\s*)?\))?$""", + re.IGNORECASE, +) def transform_dict_value_to_str(d: dict[str, Any]) -> dict[str, str]: @@ -105,6 +110,14 @@ def _parse_fixed_type(fixed: Any) -> int: return fixed +def _first_group(match: re.Match[str], *groups: int) -> str | None: + """Return the first non-empty capture group among the given group indices.""" + for group in groups: + if value := match.group(group): + return value + return None + + def _parse_geometry_type(geometry: Any) -> str: """Parse geometry type string and return CRS. @@ -115,10 +128,9 @@ def _parse_geometry_type(geometry: Any) -> str: The CRS string (defaults to DEFAULT_GEOMETRY_CRS if not specified). """ if isinstance(geometry, str): - match = GEOMETRY_REGEX.match(geometry) + match = GEOMETRY_REGEX.match(geometry.strip()) if match: - crs = match.group(1) - return crs if crs else DEFAULT_GEOMETRY_CRS + return _first_group(match, 1, 2, 3) or DEFAULT_GEOMETRY_CRS else: raise ValidationError(f"Could not parse {geometry} into a GeometryType") elif isinstance(geometry, dict): @@ -127,6 +139,16 @@ def _parse_geometry_type(geometry: Any) -> str: return geometry +def _validate_geography_algorithm(algorithm: str) -> str: + """Lower-case the edge-interpolation algorithm and check that the spec allows it.""" + normalized = algorithm.lower() + if normalized not in GEOGRAPHY_ALGORITHMS: + raise ValidationError( + f"Invalid geography edge algorithm: {algorithm}, must be one of: {', '.join(sorted(GEOGRAPHY_ALGORITHMS))}" + ) + return normalized + + def _parse_geography_type(geography: Any) -> tuple[str, str]: """Parse geography type string and return (CRS, algorithm). @@ -137,17 +159,17 @@ def _parse_geography_type(geography: Any) -> tuple[str, str]: Tuple of (CRS, algorithm) with defaults applied where not specified. """ if isinstance(geography, str): - match = GEOGRAPHY_REGEX.match(geography) + match = GEOGRAPHY_REGEX.match(geography.strip()) if match: - crs = match.group(1) if match.group(1) else DEFAULT_GEOGRAPHY_CRS - algorithm = match.group(2) if match.group(2) else DEFAULT_GEOGRAPHY_ALGORITHM - return crs, algorithm + crs = _first_group(match, 1, 2, 3) or DEFAULT_GEOGRAPHY_CRS + algorithm = _first_group(match, 4, 5, 6) or DEFAULT_GEOGRAPHY_ALGORITHM + return crs, _validate_geography_algorithm(algorithm) else: raise ValidationError(f"Could not parse {geography} into a GeographyType") elif isinstance(geography, dict): crs = geography.get("crs", DEFAULT_GEOGRAPHY_CRS) algorithm = geography.get("algorithm", DEFAULT_GEOGRAPHY_ALGORITHM) - return crs, algorithm + return crs, _validate_geography_algorithm(algorithm) else: return geography @@ -188,9 +210,9 @@ def handle_primitive_type(cls, v: Any, handler: ValidatorFunctionWrapHandler) -> # a CRS string (or CRS, algorithm tuple). If we try to parse those as type # strings here, we'd re-enter this validator (or raise) instead of letting # pydantic validate the raw root values. - if cls.__name__ == "GeometryType" and not v.startswith("geometry"): + if cls.__name__ == "GeometryType" and not v.lower().startswith("geometry"): return handler(v) - if cls.__name__ == "GeographyType" and not v.startswith("geography"): + if cls.__name__ == "GeographyType" and not v.lower().startswith("geography"): return handler(v) if v == "boolean": @@ -223,15 +245,17 @@ def handle_primitive_type(cls, v: Any, handler: ValidatorFunctionWrapHandler) -> return BinaryType() if v == "unknown": return UnknownType() + if v == "variant": + return VariantType() if v.startswith("fixed"): return FixedType(_parse_fixed_type(v)) if v.startswith("decimal"): precision, scale = _parse_decimal_type(v) return DecimalType(precision, scale) - if v.startswith("geometry"): + if v.lower().startswith("geometry"): crs = _parse_geometry_type(v) return GeometryType(crs) - if v.startswith("geography"): + if v.lower().startswith("geography"): crs, algorithm = _parse_geography_type(v) return GeographyType(crs, algorithm) else: @@ -505,9 +529,19 @@ def __repr__(self) -> str: return f"NestedField({', '.join(parts)})" - def __getnewargs__(self) -> tuple[int, str, IcebergType, bool, str | None]: + @model_validator(mode="after") + def check_null_only_types(self) -> NestedField: + """Validate that columns of the types that must default to null are optional and have no defaults.""" + if isinstance(self.field_type, (UnknownType, GeometryType, GeographyType)): + if self.initial_default is not None or self.write_default is not None: + raise ValueError(f"Columns of type {self.field_type} must default to null: {self.name}") + if isinstance(self.field_type, UnknownType) and self.required: + raise ValueError(f"Columns of type {self.field_type} must be optional: {self.name}") + return self + + def __getnewargs__(self) -> tuple[int, str, IcebergType, bool, str | None, Any, Any]: """Pickle the NestedField class.""" - return (self.field_id, self.name, self.field_type, self.required, self.doc) + return (self.field_id, self.name, self.field_type, self.required, self.doc, self.initial_default, self.write_default) @property def optional(self) -> bool: @@ -973,6 +1007,29 @@ def minimum_format_version(self) -> TableVersion: return 3 +class VariantType(PrimitiveType): + """A variant data type in Iceberg (v3+) for storing semi-structured values. + + Values are stored in the Variant binary encoding as a ``metadata`` and a ``value`` binary. The spec + defines variant as a semi-structured type; PyIceberg models it as a primitive because it cannot be + projected into and has no Iceberg field ids below it. + + Example: + >>> column_foo = VariantType() + >>> isinstance(column_foo, VariantType) + True + >>> column_foo + VariantType() + >>> str(column_foo) + 'variant' + """ + + root: Literal["variant"] = Field(default="variant") + + def minimum_format_version(self) -> TableVersion: + return 3 + + class GeometryType(PrimitiveType): """A geometry data type in Iceberg (v3+) for storing spatial geometries. @@ -990,7 +1047,7 @@ class GeometryType(PrimitiveType): >>> GeometryType("EPSG:4326") GeometryType(crs='EPSG:4326') >>> str(GeometryType("EPSG:4326")) - "geometry('EPSG:4326')" + 'geometry(EPSG:4326)' """ root: str = Field(default=DEFAULT_GEOMETRY_CRS) @@ -1000,10 +1057,10 @@ def __init__(self, crs: str = DEFAULT_GEOMETRY_CRS) -> None: @model_serializer def ser_model(self) -> str: - """Serialize the model to a string.""" + """Serialize the model to the spec (and Java) form, e.g. geometry(srid:3857).""" if self.crs == DEFAULT_GEOMETRY_CRS: return "geometry" - return f"geometry('{self.crs}')" + return f"geometry({self.crs})" @property def crs(self) -> str: @@ -1018,9 +1075,7 @@ def __repr__(self) -> str: def __str__(self) -> str: """Return the string representation.""" - if self.crs == DEFAULT_GEOMETRY_CRS: - return "geometry" - return f"geometry('{self.crs}')" + return self.ser_model() def __hash__(self) -> int: """Return the hash of the CRS.""" @@ -1056,10 +1111,10 @@ class GeographyType(PrimitiveType): 'OGC:CRS84' >>> column_foo.algorithm 'spherical' - >>> GeographyType("EPSG:4326", "planar") - GeographyType(crs='EPSG:4326', algorithm='planar') - >>> str(GeographyType("EPSG:4326", "planar")) - "geography('EPSG:4326', 'planar')" + >>> GeographyType("EPSG:4326", "vincenty") + GeographyType(crs='EPSG:4326', algorithm='vincenty') + >>> str(GeographyType("EPSG:4326", "vincenty")) + 'geography(EPSG:4326, vincenty)' """ root: tuple[str, str] = Field(default=(DEFAULT_GEOGRAPHY_CRS, DEFAULT_GEOGRAPHY_ALGORITHM)) @@ -1069,12 +1124,12 @@ def __init__(self, crs: str = DEFAULT_GEOGRAPHY_CRS, algorithm: str = DEFAULT_GE @model_serializer def ser_model(self) -> str: - """Serialize the model to a string.""" + """Serialize the model to the spec (and Java) form, e.g. geography(srid:4326, vincenty).""" if self.crs == DEFAULT_GEOGRAPHY_CRS and self.algorithm == DEFAULT_GEOGRAPHY_ALGORITHM: return "geography" if self.algorithm == DEFAULT_GEOGRAPHY_ALGORITHM: - return f"geography('{self.crs}')" - return f"geography('{self.crs}', '{self.algorithm}')" + return f"geography({self.crs})" + return f"geography({self.crs}, {self.algorithm})" @property def crs(self) -> str: @@ -1096,11 +1151,7 @@ def __repr__(self) -> str: def __str__(self) -> str: """Return the string representation.""" - if self.crs == DEFAULT_GEOGRAPHY_CRS and self.algorithm == DEFAULT_GEOGRAPHY_ALGORITHM: - return "geography" - if self.algorithm == DEFAULT_GEOGRAPHY_ALGORITHM: - return f"geography('{self.crs}')" - return f"geography('{self.crs}', '{self.algorithm}')" + return self.ser_model() def __hash__(self) -> int: """Return the hash of the tuple.""" diff --git a/pyiceberg/utils/datetime.py b/pyiceberg/utils/datetime.py index ea7329ea20..17554aaa94 100644 --- a/pyiceberg/utils/datetime.py +++ b/pyiceberg/utils/datetime.py @@ -222,6 +222,26 @@ def to_human_timestamp(timestamp_micros: int) -> str: return (EPOCH_TIMESTAMP + timedelta(microseconds=timestamp_micros)).isoformat() +def to_human_timestamp_ns(timestamp_nanos: int) -> str: + """Convert a TimestampNanoType value to human string with nanosecond precision.""" + # Python datetime only supports microsecond precision, so render the + # microsecond timestamp with a fixed 6 fractional digits and append the + # remaining 3 sub-microsecond nanosecond digits for full 9-digit precision. + micros = nanos_to_micros(timestamp_nanos) + sub_micros = timestamp_nanos - micros * 1_000 + return (EPOCH_TIMESTAMP + timedelta(microseconds=micros)).isoformat(timespec="microseconds") + f"{sub_micros:03d}" + + +def to_human_timestamptz_ns(timestamp_nanos: int) -> str: + """Convert a TimestamptzNanoType value to human string with nanosecond precision.""" + micros = nanos_to_micros(timestamp_nanos) + sub_micros = timestamp_nanos - micros * 1_000 + iso = (EPOCH_TIMESTAMPTZ + timedelta(microseconds=micros)).isoformat(timespec="microseconds") + # Insert the sub-microsecond nanosecond digits before the "+00:00" zone offset. + timestamp_part, _, offset = iso.rpartition("+") + return f"{timestamp_part}{sub_micros:03d}+{offset}" + + def micros_to_hours(micros: int) -> int: """Convert a timestamp in microseconds to hours from 1970-01-01T00:00.""" return micros // 3_600_000_000 diff --git a/pyiceberg/utils/geo.py b/pyiceberg/utils/geo.py new file mode 100644 index 0000000000..1137d169b0 --- /dev/null +++ b/pyiceberg/utils/geo.py @@ -0,0 +1,82 @@ +# Licensed to the Apache Software Foundation (ASF) under one +# or more contributor license agreements. See the NOTICE file +# distributed with this work for additional information +# regarding copyright ownership. The ASF licenses this file +# to you under the Apache License, Version 2.0 (the +# "License"); you may not use this file except in compliance +# with the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, +# software distributed under the License is distributed on an +# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +# KIND, either express or implied. See the License for the +# specific language governing permissions and limitations +# under the License. +"""Helpers for geometry and geography column bounds. + +Iceberg stores the lower and upper bounds of a geometry or geography column as the corner points of the +bounding box. Each point is serialized as the concatenation of little-endian IEEE 754 doubles +``x:y`` (16 bytes), ``x:y:z`` (24 bytes) or ``x:y:z:m`` (32 bytes); when only ``m`` is present, ``z`` is NaN. +""" + +from __future__ import annotations + +import math +import struct +from dataclasses import dataclass + + +@dataclass(frozen=True) +class GeospatialBound: + """A corner point of a geospatial bounding box.""" + + x: float + y: float + z: float | None = None + m: float | None = None + + def to_bytes(self) -> bytes: + """Serialize the point to the Iceberg geospatial bound encoding.""" + if self.m is not None: + z = self.z if self.z is not None else math.nan + return struct.pack("<4d", self.x, self.y, z, self.m) + if self.z is not None: + return struct.pack("<3d", self.x, self.y, self.z) + return struct.pack("<2d", self.x, self.y) + + @staticmethod + def from_bytes(buf: bytes) -> GeospatialBound: + """Deserialize a point from the Iceberg geospatial bound encoding.""" + if len(buf) == 16: + x, y = struct.unpack("<2d", buf) + return GeospatialBound(x, y) + if len(buf) == 24: + x, y, z = struct.unpack("<3d", buf) + return GeospatialBound(x, y, z) + if len(buf) == 32: + x, y, z, m = struct.unpack("<4d", buf) + return GeospatialBound(x, y, None if math.isnan(z) else z, m) + raise ValueError(f"Invalid geospatial bound of {len(buf)} bytes, expected 16, 24 or 32") + + +def geo_bounds_from_bbox( + xmin: float, + ymin: float, + xmax: float, + ymax: float, + zmin: float | None = None, + zmax: float | None = None, + mmin: float | None = None, + mmax: float | None = None, +) -> tuple[bytes, bytes]: + """Encode a bounding box as the serialized (lower, upper) bound points. + + Z and M are only included when both their minimum and maximum are known. + """ + has_z = zmin is not None and zmax is not None + has_m = mmin is not None and mmax is not None + lower = GeospatialBound(xmin, ymin, zmin if has_z else None, mmin if has_m else None) + upper = GeospatialBound(xmax, ymax, zmax if has_z else None, mmax if has_m else None) + return lower.to_bytes(), upper.to_bytes() diff --git a/pyiceberg/utils/parsing.py b/pyiceberg/utils/parsing.py index 200904fd97..dd3552abb8 100644 --- a/pyiceberg/utils/parsing.py +++ b/pyiceberg/utils/parsing.py @@ -28,7 +28,8 @@ class ParseNumberFromBrackets: def __init__(self, prefix: str): self.prefix = prefix - self.regex = re.compile(rf"{prefix}\[(\d+)\]") + # anchored: a name such as truncate[8]v2 is a different transform, not truncate[8] + self.regex = re.compile(rf"^{prefix}\[(\d+)\]$") def match(self, str_repr: str) -> int: matches = self.regex.search(str_repr) diff --git a/pyiceberg/utils/schema_conversion.py b/pyiceberg/utils/schema_conversion.py index 7e6343141b..259c52f4c6 100644 --- a/pyiceberg/utils/schema_conversion.py +++ b/pyiceberg/utils/schema_conversion.py @@ -17,6 +17,9 @@ """Utility class for converting between Avro and Iceberg schemas.""" import logging +import uuid +from datetime import date, datetime, time +from decimal import Decimal from typing import ( Any, ) @@ -48,13 +51,17 @@ PrimitiveType, StringType, StructType, + TimestampNanoType, TimestampType, + TimestamptzNanoType, TimestamptzType, TimeType, UnknownType, UUIDType, + VariantType, ) -from pyiceberg.utils.decimal import decimal_required_bytes +from pyiceberg.utils.datetime import date_to_days, datetime_to_micros, datetime_to_nanos, time_to_micros +from pyiceberg.utils.decimal import decimal_required_bytes, decimal_to_bytes logger = logging.getLogger(__name__) @@ -74,6 +81,7 @@ ("date", "int"): DateType(), ("time-micros", "long"): TimeType(), ("timestamp-micros", "long"): TimestampType(), + ("timestamp-nanos", "long"): TimestampNanoType(), ("uuid", "fixed"): UUIDType(), ("uuid", "string"): UUIDType(), } @@ -81,6 +89,34 @@ AvroType = str | Any +def _to_avro_default(field_type: IcebergType, value: Any) -> Any: + """Encode a default value as the JSON default of the Avro type that the field is written as. + + Avro encodes the default of a logical type as its underlying type: dates as int days, times and timestamps as + long micros (or nanos), and bytes, fixed, decimal and uuid values as strings of the code points 0-255. + """ + if isinstance(field_type, DateType): + return date_to_days(value) if isinstance(value, date) else value + if isinstance(field_type, TimeType): + return time_to_micros(value) if isinstance(value, time) else value + if isinstance(field_type, (TimestampType, TimestamptzType)): + return datetime_to_micros(value) if isinstance(value, datetime) else value + if isinstance(field_type, (TimestampNanoType, TimestamptzNanoType)): + return datetime_to_nanos(value) if isinstance(value, datetime) else value + if isinstance(field_type, UUIDType): + uuid_bytes = value.bytes if isinstance(value, uuid.UUID) else value + return uuid_bytes.decode("latin-1") + if isinstance(field_type, DecimalType): + scaled = Decimal(value).quantize(Decimal(1).scaleb(-field_type.scale)) + unscaled = decimal_to_bytes(scaled, byte_length=decimal_required_bytes(field_type.precision)) + return unscaled.decode("latin-1") + if isinstance(field_type, (BinaryType, FixedType)): + return bytes(value).decode("latin-1") + if isinstance(field_type, PrimitiveType): + return value + raise ValueError(f"Cannot encode an Avro default for {field_type}: {value}") + + class AvroSchemaConversion: def avro_to_iceberg(self, avro_schema: dict[str, Any]) -> Schema: """Convert an Apache Avro into an Apache Iceberg schema equivalent. @@ -233,12 +269,14 @@ def _convert_field(self, field: dict[str, Any]) -> NestedField: raise ValueError(f"Cannot convert field, missing {FIELD_ID_PROP}: {field}") plain_type, required = self._resolve_union(field["type"]) + field_type = self._convert_schema(plain_type) return NestedField( field_id=field[FIELD_ID_PROP], name=field["name"], - field_type=self._convert_schema(plain_type), - required=required, + field_type=field_type, + # A column of type unknown is always optional, even though Avro encodes it as a plain null + required=required and not isinstance(field_type, UnknownType), doc=field.get("doc"), ) @@ -381,6 +419,11 @@ def _convert_logical_type(self, avro_logical_type: dict[str, Any]) -> IcebergTyp return TimestamptzType() else: return TimestampType() + elif logical_type == "timestamp-nanos": + if avro_logical_type.get("adjust-to-utc", False) is True: + return TimestamptzNanoType() + else: + return TimestampNanoType() elif (logical_type, physical_type) in LOGICAL_FIELD_TYPE_MAPPING: return LOGICAL_FIELD_TYPE_MAPPING[(logical_type, physical_type)] else: @@ -531,19 +574,22 @@ def field(self, field: NestedField, field_result: AvroType) -> AvroType: original_name = field.name sanitized_name = make_compatible_name(original_name) + is_null = isinstance(field.field_type, UnknownType) result = { "name": sanitized_name, FIELD_ID_PROP: field.field_id, - "type": field_result if field.required else ["null", field_result], + # A union of null with null is not valid Avro + "type": field_result if field.required or is_null else ["null", field_result], } if original_name != sanitized_name: result[ICEBERG_FIELD_NAME_PROP] = original_name - if field.write_default is not None: - result["default"] = field.write_default - elif field.optional: + if field.optional or is_null: + # The Avro default of a union must match its first branch, which is null result["default"] = None + elif field.write_default is not None: + result["default"] = _to_avro_default(field.field_type, field.write_default) if field.doc is not None: result["doc"] = field.doc @@ -639,6 +685,14 @@ def visit_binary(self, binary_type: BinaryType) -> AvroType: def visit_unknown(self, unknown_type: UnknownType) -> AvroType: return "null" + def visit_variant(self, variant_type: VariantType) -> AvroType: + """Convert variant type to an Avro record whose metadata and value fields have no field ids.""" + return { + "type": "record", + "logicalType": "variant", + "fields": [{"name": "metadata", "type": "bytes"}, {"name": "value", "type": "bytes"}], + } + def visit_geometry(self, geometry_type: GeometryType) -> AvroType: """Convert geometry type to Avro bytes (WKB format per Iceberg spec).""" return "bytes" diff --git a/pyiceberg/variant.py b/pyiceberg/variant.py new file mode 100644 index 0000000000..5c6b4e3280 --- /dev/null +++ b/pyiceberg/variant.py @@ -0,0 +1,304 @@ +# Licensed to the Apache Software Foundation (ASF) under one +# or more contributor license agreements. See the NOTICE file +# distributed with this work for additional information +# regarding copyright ownership. The ASF licenses this file +# to you under the Apache License, Version 2.0 (the +# "License"); you may not use this file except in compliance +# with the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, +# software distributed under the License is distributed on an +# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +# KIND, either express or implied. See the License for the +# specific language governing permissions and limitations +# under the License. +r"""Decoder for the Variant binary encoding. + +Iceberg stores variant values in the Variant binary encoding defined by the Parquet project +(https://github.com/apache/parquet-format/blob/master/VariantEncoding.md): a ``metadata`` binary holding +the dictionary of object keys, and a ``value`` binary holding the encoded value. Scans return variant +columns as a struct of these two binaries; the functions in this module decode one value. + +Example: + >>> from pyiceberg.variant import to_json, to_python + >>> to_python(b"\\x01\\x00\\x00", b"\\x35just a string") + 'just a string' + >>> to_json(b"\\x01\\x01\\x00\\x01a", b"\\x02\\x01\\x00\\x00\\x02\\x0c\\x01") + '{"a":1}' +""" + +from __future__ import annotations + +import base64 +import json +import math +import struct +import uuid +from collections.abc import Callable +from datetime import date, datetime, time, timedelta +from decimal import Decimal +from typing import Any + +from pyiceberg.utils.datetime import EPOCH_DATE, EPOCH_TIMESTAMP, EPOCH_TIMESTAMPTZ +from pyiceberg.utils.decimal import unscaled_to_decimal + +VARIANT_VERSION = 1 + +# Basic types, stored in the lower two bits of the first byte of a value +_PRIMITIVE = 0 +_SHORT_STRING = 1 +_OBJECT = 2 + +# Primitive type ids, stored in the upper six bits of the first byte of a primitive value +_NULL = 0 +_TRUE = 1 +_FALSE = 2 +_INT8 = 3 +_INT16 = 4 +_INT32 = 5 +_INT64 = 6 +_DOUBLE = 7 +_DECIMAL4 = 8 +_DECIMAL8 = 9 +_DECIMAL16 = 10 +_DATE = 11 +_TIMESTAMPTZ = 12 +_TIMESTAMP = 13 +_FLOAT = 14 +_BINARY = 15 +_STRING = 16 +_TIME = 17 +_TIMESTAMPTZ_NANOS = 18 +_TIMESTAMP_NANOS = 19 +_UUID = 20 + +_FIXED_WIDTH_FORMATS = { + _INT8: " None: + self.nanos = nanos + self.utc = utc + + def to_datetime(self) -> datetime: + """Return the timestamp truncated to microseconds.""" + return (EPOCH_TIMESTAMPTZ if self.utc else EPOCH_TIMESTAMP) + timedelta(microseconds=self.nanos // 1000) + + def isoformat(self) -> str: + """Return the ISO-8601 representation with nanosecond precision.""" + seconds = self.to_datetime().replace(microsecond=0).isoformat() + fraction = f".{self.nanos % 1_000_000_000:09d}" + if self.utc: + # Place the fraction before the "+00:00" offset + return f"{seconds[:-6]}{fraction}{seconds[-6:]}" + return f"{seconds}{fraction}" + + +def _read_unsigned(buf: bytes, pos: int, size: int) -> int: + if pos + size > len(buf): + raise ValueError("Invalid variant: unexpected end of buffer") + return int.from_bytes(buf[pos : pos + size], "little", signed=False) + + +def _parse_metadata(metadata: bytes) -> list[str]: + """Return the dictionary of object keys stored in variant metadata.""" + if len(metadata) == 0: + raise ValueError("Invalid variant metadata: empty buffer") + header = metadata[0] + version = header & 0x0F + if version != VARIANT_VERSION: + raise ValueError(f"Unsupported variant metadata version: {version}") + offset_size = ((header >> 6) & 0x03) + 1 + dictionary_size = _read_unsigned(metadata, 1, offset_size) + offsets_start = 1 + offset_size + offsets = [_read_unsigned(metadata, offsets_start + index * offset_size, offset_size) for index in range(dictionary_size + 1)] + strings_start = offsets_start + (dictionary_size + 1) * offset_size + if strings_start + offsets[-1] > len(metadata): + raise ValueError("Invalid variant metadata: unexpected end of buffer") + return [ + metadata[strings_start + offsets[index] : strings_start + offsets[index + 1]].decode("utf-8") + for index in range(dictionary_size) + ] + + +def _decode_primitive(type_id: int, value: bytes, pos: int) -> Any: + if type_id == _NULL: + return None + if type_id == _TRUE: + return True + if type_id == _FALSE: + return False + if (fmt := _FIXED_WIDTH_FORMATS.get(type_id)) is not None: + width = struct.calcsize(fmt) + if pos + width > len(value): + raise ValueError("Invalid variant: unexpected end of buffer") + return struct.unpack_from(fmt, value, pos)[0] + if (decimal_width := _DECIMAL_WIDTHS.get(type_id)) is not None: + scale = _read_unsigned(value, pos, 1) + unscaled = int.from_bytes(value[pos + 1 : pos + 1 + decimal_width], "little", signed=True) + return unscaled_to_decimal(unscaled, scale) + if type_id == _DATE: + return EPOCH_DATE + timedelta(days=struct.unpack_from(" Any: + if pos >= len(value): + raise ValueError("Invalid variant: unexpected end of buffer") + value_metadata = value[pos] + basic_type = value_metadata & 0x03 + header = value_metadata >> 2 + + if basic_type == _PRIMITIVE: + return _decode_primitive(header, value, pos + 1) + + if basic_type == _SHORT_STRING: + data = value[pos + 1 : pos + 1 + header] + if len(data) != header: + raise ValueError("Invalid variant: unexpected end of buffer") + return data.decode("utf-8") + + if basic_type == _OBJECT: + offset_size = (header & 0x03) + 1 + field_id_size = ((header >> 2) & 0x03) + 1 + is_large = (header >> 4) & 0x01 + num_elements_size = 4 if is_large else 1 + num_elements = _read_unsigned(value, pos + 1, num_elements_size) + field_ids_start = pos + 1 + num_elements_size + offsets_start = field_ids_start + num_elements * field_id_size + values_start = offsets_start + (num_elements + 1) * offset_size + result: dict[str, Any] = {} + for index in range(num_elements): + field_id = _read_unsigned(value, field_ids_start + index * field_id_size, field_id_size) + if field_id >= len(keys): + raise ValueError(f"Invalid variant: field id {field_id} is not in the metadata dictionary") + offset = _read_unsigned(value, offsets_start + index * offset_size, offset_size) + result[keys[field_id]] = _decode(keys, value, values_start + offset) + return result + + # The remaining basic type is an array + offset_size = (header & 0x03) + 1 + is_large = (header >> 2) & 0x01 + num_elements_size = 4 if is_large else 1 + num_elements = _read_unsigned(value, pos + 1, num_elements_size) + offsets_start = pos + 1 + num_elements_size + values_start = offsets_start + (num_elements + 1) * offset_size + return [ + _decode(keys, value, values_start + _read_unsigned(value, offsets_start + index * offset_size, offset_size)) + for index in range(num_elements) + ] + + +def _convert(decoded: Any, nano_timestamp: Callable[[_NanoTimestamp], Any]) -> Any: + if isinstance(decoded, _NanoTimestamp): + return nano_timestamp(decoded) + if isinstance(decoded, dict): + return {key: _convert(item, nano_timestamp) for key, item in decoded.items()} + if isinstance(decoded, list): + return [_convert(item, nano_timestamp) for item in decoded] + return decoded + + +def to_python(metadata: bytes, value: bytes) -> Any: + """Decode a variant value into Python objects. + + Objects decode to ``dict``, arrays to ``list``, and primitives to ``None``, ``bool``, ``int``, ``float``, + ``Decimal``, ``str``, ``bytes``, ``date``, ``time``, ``datetime`` (timezone-aware in UTC for timestamps + with a time zone) or ``UUID``. Nanosecond timestamps are truncated to microseconds, the precision + of ``datetime``; use ``to_json`` to keep all digits. + + Args: + metadata: The variant metadata, holding the dictionary of object keys. + value: The encoded variant value. + + Returns: + The decoded value. + + Raises: + ValueError: If the buffers are not a valid variant. + """ + decoded = _decode(_parse_metadata(bytes(metadata)), bytes(value), 0) + return _convert(decoded, _NanoTimestamp.to_datetime) + + +def _json_text(decoded: Any) -> str: + if decoded is None: + return "null" + if isinstance(decoded, bool): + return "true" if decoded else "false" + if isinstance(decoded, int): + return str(decoded) + if isinstance(decoded, float): + # JSON has no representation for NaN and infinity + return json.dumps(str(decoded)) if math.isnan(decoded) or math.isinf(decoded) else repr(decoded) + if isinstance(decoded, Decimal): + return format(decoded, "f") + if isinstance(decoded, str): + return json.dumps(decoded, ensure_ascii=False) + if isinstance(decoded, bytes): + return json.dumps(base64.b64encode(decoded).decode("ascii")) + if isinstance(decoded, (date, time, datetime, _NanoTimestamp)): + return json.dumps(decoded.isoformat()) + if isinstance(decoded, uuid.UUID): + return json.dumps(str(decoded)) + if isinstance(decoded, dict): + return "{" + ",".join(f"{json.dumps(key, ensure_ascii=False)}:{_json_text(item)}" for key, item in decoded.items()) + "}" + if isinstance(decoded, list): + return "[" + ",".join(_json_text(item) for item in decoded) + "]" + raise ValueError(f"Cannot convert {decoded!r} to JSON") + + +def to_json(metadata: bytes, value: bytes) -> str: + """Decode a variant value into a JSON string. + + Decimals are written as JSON numbers with all their digits; binary values as base64 strings; dates, + times and timestamps as ISO-8601 strings; UUIDs as strings; NaN and infinite floats as strings. + + Args: + metadata: The variant metadata, holding the dictionary of object keys. + value: The encoded variant value. + + Returns: + The JSON text of the value. + + Raises: + ValueError: If the buffers are not a valid variant. + """ + return _json_text(_decode(_parse_metadata(bytes(metadata)), bytes(value), 0)) + + +__all__ = ["to_json", "to_python", "VARIANT_VERSION"] diff --git a/tests/avro/test_file.py b/tests/avro/test_file.py index 2d3ddeefab..c6d0095013 100644 --- a/tests/avro/test_file.py +++ b/tests/avro/test_file.py @@ -15,15 +15,17 @@ # specific language governing permissions and limitations # under the License. import inspect +import json from _decimal import Decimal from datetime import datetime from enum import Enum +from io import StringIO from tempfile import TemporaryDirectory from typing import Any from uuid import UUID import pytest -from fastavro import reader, writer +from fastavro import json_writer, reader, writer import pyiceberg.avro.file as avro from pyiceberg.avro.codecs.deflate import DeflateCodec @@ -41,6 +43,7 @@ from pyiceberg.schema import Schema from pyiceberg.typedef import Record, TableVersion from pyiceberg.types import ( + BinaryType, BooleanType, DateType, DecimalType, @@ -50,7 +53,9 @@ IntegerType, LongType, NestedField, + PrimitiveType, StringType, + TimestampNanoType, TimestampType, TimestamptzType, TimeType, @@ -453,3 +458,54 @@ def field_uuid(self) -> UUID: for idx, field in enumerate(all_primitives_schema.as_struct()): assert record[idx] == avro_entry[idx], f"Invalid {field}" assert record[idx] == avro_entry_read_with_fastavro[idx], f"Invalid {field} read with fastavro" + + +@pytest.mark.parametrize( + "field_type, value", + [ + (BooleanType(), True), + (IntegerType(), 34), + (LongType(), 429496729622), + (FloatType(), 1.5), + (DoubleType(), 2.25), + (StringType(), "blue"), + (BinaryType(), b"\x00\x7f\xff"), + (FixedType(3), b"\x00\x7f\xff"), + (DecimalType(6, 2), Decimal("-123.45")), + (DateType(), 19052), + (TimeType(), 69922000000), + (TimestampType(), 1677629965000000), + (TimestamptzType(), 1677629965000000), + (TimestampNanoType(), 1677629965000000001), + (UUIDType(), UUID("12345678-1234-5678-1234-567812345678")), + ], +) +def test_write_default_is_avro_default(field_type: PrimitiveType, value: Any) -> None: + """The Avro default of a field decodes to the same value as the written field.""" + id_field = NestedField(field_id=1, name="id", field_type=IntegerType(), required=True) + value_field = NestedField(field_id=2, name="value", field_type=field_type, required=True) + without_value = Schema(id_field) + with_value = Schema(id_field, value_field) + with_default = Schema( + id_field, NestedField(field_id=2, name="value", field_type=field_type, required=True, write_default=value) + ) + + with TemporaryDirectory() as tmpdir: + without_value_file, with_value_file = f"{tmpdir}/without_value.avro", f"{tmpdir}/with_value.avro" + with avro.AvroOutputFile[Record](PyArrowFileIO().new_output(without_value_file), without_value, "test") as out: + out.write_block([Record(1)]) + with avro.AvroOutputFile[Record](PyArrowFileIO().new_output(with_value_file), with_value, "test") as out: + out.write_block([Record(1, value)]) + + with open(with_value_file, "rb") as fo: + (expected,) = list(reader(fo)) + + # Schema resolution fills the field that is missing from the file with the raw Avro default + reader_schema = AvroSchemaConversion().iceberg_to_avro(with_default, schema_name="test") + with open(without_value_file, "rb") as fo: + (actual,) = list(reader(fo, reader_schema=reader_schema)) + + # The Avro JSON encoding of the written value is the expected JSON default + encoded = StringIO() + json_writer(encoded, reader_schema, [expected]) + assert actual == json.loads(encoded.getvalue()) diff --git a/tests/catalog/test_catalog_behaviors.py b/tests/catalog/test_catalog_behaviors.py index b859e2d541..97fea5f248 100644 --- a/tests/catalog/test_catalog_behaviors.py +++ b/tests/catalog/test_catalog_behaviors.py @@ -25,6 +25,7 @@ from typing import Any import pyarrow as pa +import pyarrow.parquet as pq import pytest from pydantic_core import ValidationError from pytest_lazy_fixtures import lf @@ -133,7 +134,7 @@ def test_create_table_without_namespace(catalog: Catalog, table_schema_nested: S catalog.create_table(table_name, table_schema_nested) -@pytest.mark.parametrize("format_version", [1, 2]) +@pytest.mark.parametrize("format_version", [1, 2, 3]) def test_create_table_transaction(catalog: Catalog, format_version: int) -> None: identifier = f"default.arrow_create_table_transaction_{catalog.name}_{format_version}" try: @@ -720,6 +721,200 @@ def test_append_table(catalog: Catalog, table_schema_simple: Schema, test_table_ assert df == table.scan().to_arrow() +def test_v3_row_lineage_on_commit(catalog: Catalog, table_schema_simple: Schema, test_table_identifier: Identifier) -> None: + catalog.create_namespace(Catalog.namespace_from(test_table_identifier)) + table = catalog.create_table(test_table_identifier, table_schema_simple, properties={"format-version": "3"}) + assert table.format_version == 3 + assert table.metadata.next_row_id == 0 + + df = pa.Table.from_pydict( + {"foo": ["a", "b", "c"], "bar": [1, 2, 3], "baz": [True, False, True]}, + schema=schema_to_pyarrow(table_schema_simple), + ) + table.append(df) + table.append(df.slice(0, 2)) + + first, second = table.snapshots() + assert (first.first_row_id, first.added_rows) == (0, 3) + assert (second.first_row_id, second.added_rows) == (3, 2) + assert table.metadata.next_row_id == 5 + + # Fast appends carry over the manifests of earlier snapshots with their first row ids + manifests = second.manifests(table.io) + assert {m.first_row_id for m in manifests} == {0, 3} + assert all(entry.data_file.first_row_id is None for m in manifests for entry in _raw_entries(table, m)) + + # A copy-on-write delete keeps the row ids of the rows it copies into the rewritten data file + table.delete("bar = 1") + reloaded = catalog.load_table(test_table_identifier) + assert reloaded.metadata.next_row_id == reloaded.snapshots()[-1].first_row_id + reloaded.snapshots()[-1].added_rows # type: ignore[operator] + assert reloaded.metadata.next_row_id == sum(s.added_rows or 0 for s in reloaded.snapshots()) + assert sorted(reloaded.scan().to_arrow()["bar"].to_pylist()) == [2, 2, 3] + + +def _lineage(table: Any) -> dict[int, tuple[int, int]]: + """Return the row id and last updated sequence number of each row, keyed by ``bar``.""" + rows = table.scan(selected_fields=("bar", "_row_id", "_last_updated_sequence_number")).to_arrow().to_pylist() + return {row["bar"]: (row["_row_id"], row["_last_updated_sequence_number"]) for row in rows} + + +def _v3_table_with_two_files(catalog: Catalog, schema: Schema, identifier: Identifier) -> Any: + catalog.create_namespace(Catalog.namespace_from(identifier)) + table = catalog.create_table(identifier, schema, properties={"format-version": "3"}) + arrow_schema = schema_to_pyarrow(schema) + table.append( + pa.Table.from_pydict({"foo": ["a", "b", "c"], "bar": [1, 2, 3], "baz": [True, False, True]}, schema=arrow_schema) + ) + table.append(pa.Table.from_pydict({"foo": ["d", "e"], "bar": [4, 5], "baz": [False, True]}, schema=arrow_schema)) + assert _lineage(table) == {1: (0, 1), 2: (1, 1), 3: (2, 1), 4: (3, 2), 5: (4, 2)} + return table + + +def test_v3_copy_on_write_delete_preserves_row_lineage( + catalog: Catalog, table_schema_simple: Schema, test_table_identifier: Identifier +) -> None: + table = _v3_table_with_two_files(catalog, table_schema_simple, test_table_identifier) + + table.delete("bar = 5") + + assert table.snapshots()[-1].summary.operation == Operation.OVERWRITE + assert _lineage(table) == {1: (0, 1), 2: (1, 1), 3: (2, 1), 4: (3, 2)} + # Copied rows still count as added rows, so the next row id stays above every assigned id + assert table.metadata.next_row_id == 6 + + # The lineage columns are stored in the rewritten file but do not surface in normal reads + rewritten = [task.file for task in table.scan(row_filter="bar = 4").plan_files()] + assert len(rewritten) == 1 + physical_schema = pq.read_schema(table.io.new_input(rewritten[0].file_path).open()) + assert physical_schema.field("_row_id").metadata[b"PARQUET:field_id"] == b"2147483540" + assert physical_schema.field("_last_updated_sequence_number").metadata[b"PARQUET:field_id"] == b"2147483539" + assert table.scan().to_arrow().column_names == ["foo", "bar", "baz"] + assert table.scan(selected_fields=("bar", "_row_id")).to_arrow().column_names == ["bar", "_row_id"] + + +def test_v3_copy_on_write_delete_preserves_row_lineage_through_deletion_vector( + catalog: Catalog, table_schema_simple: Schema, test_table_identifier: Identifier +) -> None: + table = _v3_table_with_two_files(catalog, table_schema_simple, test_table_identifier) + table.delete("bar = 1") + assert _lineage(table) == {2: (1, 1), 3: (2, 1), 4: (3, 2), 5: (4, 2)} + + # Positions of a later deletion vector refer to the rewritten file, which keeps the copied lineage + table.transaction().set_properties( + {TableProperties.DELETE_MODE: TableProperties.DELETE_MODE_MERGE_ON_READ} + ).commit_transaction() + table.delete("bar = 2") + assert table.snapshots()[-1].summary.operation == Operation.DELETE + assert _lineage(table) == {3: (2, 1), 4: (3, 2), 5: (4, 2)} + + +def test_v3_overwrite_preserves_row_lineage( + catalog: Catalog, table_schema_simple: Schema, test_table_identifier: Identifier +) -> None: + table = _v3_table_with_two_files(catalog, table_schema_simple, test_table_identifier) + + table.overwrite( + pa.Table.from_pydict({"foo": ["f"], "bar": [6], "baz": [False]}, schema=schema_to_pyarrow(table_schema_simple)), + overwrite_filter="bar = 5", + ) + + # Row 4 is copied with its lineage, row 6 is new and gets an id from its own snapshot + lineage = _lineage(table) + assert {bar: lineage[bar] for bar in (1, 2, 3, 4)} == {1: (0, 1), 2: (1, 1), 3: (2, 1), 4: (3, 2)} + assert lineage[6] == (6, 4) + assert table.metadata.next_row_id == 7 + + +def test_v3_upsert_preserves_row_lineage_of_unchanged_rows( + catalog: Catalog, table_schema_simple: Schema, test_table_identifier: Identifier +) -> None: + table = _v3_table_with_two_files(catalog, table_schema_simple, test_table_identifier) + + result = table.upsert( + pa.Table.from_pydict( + {"foo": ["a", "B"], "bar": [1, 2], "baz": [True, False]}, schema=schema_to_pyarrow(table_schema_simple) + ) + ) + assert (result.rows_updated, result.rows_inserted) == (1, 0) + + lineage = _lineage(table) + # Untouched rows keep their lineage, including the ones copied out of the rewritten file + assert {bar: lineage[bar] for bar in (1, 3, 4, 5)} == {1: (0, 1), 3: (2, 1), 4: (3, 2), 5: (4, 2)} + # The updated row is written as a new row with a fresh id and the sequence number of its snapshot + updated_snapshot = table.snapshots()[-1] + assert lineage[2][0] >= 5 + assert lineage[2][1] == updated_snapshot.sequence_number + + +def _raw_entries(table: Any, manifest: Any) -> list[Any]: + """Read manifest entries without first row id inheritance, to check what the writer stored.""" + from pyiceberg.avro.file import AvroFile + from pyiceberg.manifest import MANIFEST_ENTRY_SCHEMAS, DataFile, ManifestEntry + + with AvroFile[ManifestEntry]( + table.io.new_input(manifest.manifest_path), MANIFEST_ENTRY_SCHEMAS[3], read_types={-1: ManifestEntry, 2: DataFile} + ) as reader: + return list(reader) + + +def test_v3_upgrade_then_append(catalog: Catalog, table_schema_simple: Schema, test_table_identifier: Identifier) -> None: + catalog.create_namespace(Catalog.namespace_from(test_table_identifier)) + table = catalog.create_table(test_table_identifier, table_schema_simple, properties={"format-version": "2"}) + df = pa.Table.from_pydict( + {"foo": ["a", "b"], "bar": [1, 2], "baz": [True, False]}, + schema=schema_to_pyarrow(table_schema_simple), + ) + table.append(df) + + with table.transaction() as txn: + txn.upgrade_table_version(format_version=3) + table = catalog.load_table(test_table_identifier) + assert table.format_version == 3 + assert table.metadata.next_row_id == 0 + assert table.snapshots()[0].first_row_id is None + + table.append(df.slice(0, 1)) + snapshot = table.snapshots()[-1] + # The pre-upgrade manifest has no first row id yet, so it is assigned one together with the new manifest + assert (snapshot.first_row_id, snapshot.added_rows) == (0, 3) + assert table.metadata.next_row_id == 3 + assert all(m.first_row_id is not None for m in snapshot.manifests(table.io)) + assert len(table.scan().to_arrow()) == 3 + + +def test_failed_append_removes_written_files( + catalog: Catalog, table_schema_simple: Schema, test_table_identifier: Identifier, monkeypatch: pytest.MonkeyPatch +) -> None: + import pyiceberg.table.update.snapshot as snapshot_module + + catalog.create_namespace(Catalog.namespace_from(test_table_identifier)) + table = catalog.create_table(test_table_identifier, table_schema_simple) + df = pa.Table.from_pydict( + {"foo": ["a"], "bar": [1], "baz": [True]}, + schema=schema_to_pyarrow(table_schema_simple), + ) + + written: list[str] = [] + original_append = snapshot_module._SnapshotProducer._append_written_data_file + + def _track(self: Any, data_file: Any) -> Any: + written.append(data_file.file_path) + return original_append(self, data_file) + + def _failing_write_manifest(*args: Any, **kwargs: Any) -> Any: + raise OSError("Simulated manifest write failure") + + monkeypatch.setattr(snapshot_module._SnapshotProducer, "_append_written_data_file", _track) + monkeypatch.setattr(snapshot_module, "write_manifest", _failing_write_manifest) + + with pytest.raises(OSError, match="Simulated manifest write failure"): + table.append(df) + + assert len(written) == 1 + assert not table.io.new_input(written[0]).exists() + assert catalog.load_table(test_table_identifier).current_snapshot() is None + + # Test writes def test_table_writes_metadata_to_custom_location( catalog: Catalog, @@ -803,7 +998,7 @@ def test_table_metadata_writes_reflect_latest_path( assert table.location_provider().new_metadata_location("metadata.json") == f"{new_metadata_path}/metadata.json" -@pytest.mark.parametrize("format_version", [1, 2]) +@pytest.mark.parametrize("format_version", [1, 2, 3]) def test_write_and_evolve(catalog: Catalog, format_version: int) -> None: identifier = f"default.arrow_write_data_and_evolve_schema_v{format_version}" @@ -849,7 +1044,7 @@ def test_write_and_evolve(catalog: Catalog, format_version: int) -> None: # Merge manifests -@pytest.mark.parametrize("format_version", [1, 2]) +@pytest.mark.parametrize("format_version", [1, 2, 3]) def test_merge_manifests_local_file_system(catalog: Catalog, arrow_table_with_null: pa.Table, format_version: int) -> None: # To catch manifest file name collision bug during merge: # https://github.com/apache/iceberg-python/pull/363#discussion_r1660691918 diff --git a/tests/catalog/test_scan_planning_models.py b/tests/catalog/test_scan_planning_models.py index caf571d322..09d2a0ad6d 100644 --- a/tests/catalog/test_scan_planning_models.py +++ b/tests/catalog/test_scan_planning_models.py @@ -39,7 +39,7 @@ ValueMap, ) from pyiceberg.expressions import AlwaysTrue, EqualTo, Reference -from pyiceberg.manifest import FileFormat +from pyiceberg.manifest import DataFileContent, FileFormat TEST_URI = "https://iceberg-test-catalog/" @@ -694,7 +694,7 @@ def test_plan_scan_cancelled(rest_scan_catalog: RestCatalog, requests_mock: Mock rest_scan_catalog.plan_scan(("db", "tbl"), request) -def test_plan_scan_equality_deletes_not_supported(rest_scan_catalog: RestCatalog, requests_mock: Mocker) -> None: +def test_plan_scan_equality_deletes(rest_scan_catalog: RestCatalog, requests_mock: Mocker) -> None: file_one = _rest_data_file(file_path="s3://bucket/tbl/data/file1.parquet") equality_delete = _rest_equality_delete_file(equality_ids=[1, 2]) requests_mock.post( @@ -715,8 +715,11 @@ def test_plan_scan_equality_deletes_not_supported(rest_scan_catalog: RestCatalog ) request = PlanTableScanRequest() - with pytest.raises(NotImplementedError, match="PyIceberg does not yet support equality deletes"): - rest_scan_catalog.plan_scan(("db", "tbl"), request) + tasks = rest_scan_catalog.plan_scan(("db", "tbl"), request) + assert len(tasks) == 1 + (delete_file,) = tasks[0].delete_files + assert delete_file.content == DataFileContent.EQUALITY_DELETES + assert delete_file.equality_ids == [1, 2] def _mock_load_table(requests_mock: Mocker, metadata: dict[str, Any], config: dict[str, str]) -> None: @@ -816,3 +819,22 @@ def test_scan_load_table_override_survives_invalid_catalog_mode( assert scan._should_use_server_side_planning() is True assert [task.file.file_path for task in scan.plan_files()] == ["s3://bucket/tbl/data/file1.parquet"] assert plan_mock.call_count == 1 + + +def test_file_scan_task_from_rest_response_carries_v3_fields() -> None: + from pyiceberg.table import FileScanTask + + data_file = {**_rest_data_file(), "first-row-id": 42} + delete_file = { + **_rest_position_delete_file(file_path="s3://bucket/table/deletes.puffin", file_format="puffin"), + "referenced-data-file": "s3://bucket/table/data/file.parquet", + } + scan_tasks = ScanTasks.model_validate( + {"delete-files": [delete_file], "file-scan-tasks": [{"data-file": data_file, "delete-file-references": [0]}]} + ) + task = FileScanTask.from_rest_response(scan_tasks.file_scan_tasks[0], scan_tasks.delete_files) + assert task.file.first_row_id == 42 + (dv,) = task.delete_files + assert dv.referenced_data_file == "s3://bucket/table/data/file.parquet" + assert dv.content_offset == 100 + assert dv.content_size_in_bytes == 200 diff --git a/tests/conftest.py b/tests/conftest.py index 316499ba53..16da3e947c 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -972,7 +972,7 @@ def generate_snapshot( {"id": 1, "name": "x", "required": True, "type": "long"}, {"id": 2, "name": "y", "required": True, "type": "long", "doc": "comment"}, {"id": 3, "name": "z", "required": True, "type": "long"}, - {"id": 4, "name": "u", "required": True, "type": "unknown"}, + {"id": 4, "name": "u", "required": False, "type": "unknown"}, {"id": 5, "name": "ns", "required": True, "type": "timestamp_ns"}, {"id": 6, "name": "nstz", "required": True, "type": "timestamptz_ns"}, ], diff --git a/tests/expressions/test_evaluator.py b/tests/expressions/test_evaluator.py index bba4156e99..60e25cbc8d 100644 --- a/tests/expressions/test_evaluator.py +++ b/tests/expressions/test_evaluator.py @@ -15,6 +15,7 @@ # specific language governing permissions and limitations # under the License. # pylint:disable=redefined-outer-name +import struct from collections.abc import Iterator, Mapping from concurrent.futures import ThreadPoolExecutor from threading import Event @@ -58,6 +59,8 @@ from pyiceberg.types import ( DoubleType, FloatType, + GeographyType, + GeometryType, IcebergType, IntegerType, LongType, @@ -1907,3 +1910,83 @@ def test_strict_metrics_eval_bounds_after_promotion( evaluator = _StrictMetricsEvaluator(schema, op("col", lit)) assert evaluator.eval(data_file) == expected + + +def _wkb_point(x: float, y: float) -> bytes: + return struct.pack(" Schema: + return Schema( + NestedField(1, "geom", GeometryType(), required=False), + NestedField(2, "geog", GeographyType("srid:4326", "vincenty"), required=False), + NestedField(3, "geom_nulls", GeometryType(), required=False), + ) + + +@pytest.fixture +def geo_data_file() -> DataFile: + # Geo bounds are the corner points of the bounding box, not the smallest and largest WKB values + return DataFile.from_args( + file_path="geo.parquet", + file_format=FileFormat.PARQUET, + partition={}, + record_count=10, + file_size_in_bytes=3, + value_counts={1: 10, 2: 10, 3: 10}, + null_value_counts={1: 0, 2: 0, 3: 10}, + lower_bounds={1: _wkb_point(0, 0), 2: _wkb_point(0, 0)}, + upper_bounds={1: _wkb_point(10, 10), 2: _wkb_point(10, 10)}, + ) + + +@pytest.mark.parametrize("column", ["geom", "geog"]) +@pytest.mark.parametrize( + "predicate", + [ + lambda col: EqualTo(col, GEO_LINESTRING), + lambda col: NotEqualTo(col, GEO_LINESTRING), + lambda col: In(col, {GEO_LINESTRING, _wkb_point(20, 20)}), + lambda col: NotIn(col, {GEO_LINESTRING}), + lambda col: EqualTo(col, _wkb_point(0, 0)), + lambda col: NotNull(col), + ], +) +def test_inclusive_metrics_geo_bounds_are_not_compared( + geo_schema: Schema, geo_data_file: DataFile, column: str, predicate: Any +) -> None: + assert _InclusiveMetricsEvaluator(geo_schema, predicate(column)).eval(geo_data_file) == ROWS_MIGHT_MATCH + + +@pytest.mark.parametrize("column", ["geom", "geog"]) +@pytest.mark.parametrize( + "predicate", + [ + lambda col: EqualTo(col, GEO_LINESTRING), + lambda col: NotEqualTo(col, GEO_LINESTRING), + lambda col: In(col, {_wkb_point(0, 0), _wkb_point(10, 10)}), + lambda col: NotIn(col, {GEO_LINESTRING}), + lambda col: EqualTo(col, _wkb_point(0, 0)), + lambda col: IsNull(col), + ], +) +def test_strict_metrics_geo_bounds_are_not_compared( + geo_schema: Schema, geo_data_file: DataFile, column: str, predicate: Any +) -> None: + assert _StrictMetricsEvaluator(geo_schema, predicate(column)).eval(geo_data_file) == ROWS_MIGHT_NOT_MATCH + + +def test_metrics_geo_null_counts_still_apply(geo_schema: Schema, geo_data_file: DataFile) -> None: + assert _InclusiveMetricsEvaluator(geo_schema, EqualTo("geom_nulls", GEO_LINESTRING)).eval(geo_data_file) == ROWS_CANNOT_MATCH + assert _InclusiveMetricsEvaluator(geo_schema, In("geom_nulls", {GEO_LINESTRING})).eval(geo_data_file) == ROWS_CANNOT_MATCH + assert _InclusiveMetricsEvaluator(geo_schema, NotNull("geom_nulls")).eval(geo_data_file) == ROWS_CANNOT_MATCH + assert _InclusiveMetricsEvaluator(geo_schema, IsNull("geom")).eval(geo_data_file) == ROWS_CANNOT_MATCH + assert _StrictMetricsEvaluator(geo_schema, NotEqualTo("geom_nulls", GEO_LINESTRING)).eval(geo_data_file) == ROWS_MUST_MATCH + assert _StrictMetricsEvaluator(geo_schema, NotIn("geom_nulls", {GEO_LINESTRING})).eval(geo_data_file) == ROWS_MUST_MATCH + assert _StrictMetricsEvaluator(geo_schema, IsNull("geom_nulls")).eval(geo_data_file) == ROWS_MUST_MATCH + assert _StrictMetricsEvaluator(geo_schema, NotNull("geom")).eval(geo_data_file) == ROWS_MUST_MATCH diff --git a/tests/expressions/test_expressions.py b/tests/expressions/test_expressions.py index 8ce48a6897..8d57086a93 100644 --- a/tests/expressions/test_expressions.py +++ b/tests/expressions/test_expressions.py @@ -20,6 +20,7 @@ import pickle import uuid from decimal import Decimal +from typing import Any import pytest from typing_extensions import assert_type @@ -63,7 +64,7 @@ UnboundPredicate, ) from pyiceberg.expressions.literals import Literal, literal -from pyiceberg.expressions.visitors import _from_byte_buffer +from pyiceberg.expressions.visitors import _from_byte_buffer, expression_evaluator from pyiceberg.schema import Accessor, Schema from pyiceberg.typedef import Record from pyiceberg.types import ( @@ -71,6 +72,8 @@ DecimalType, DoubleType, FloatType, + GeographyType, + GeometryType, IntegerType, ListType, LongType, @@ -1401,3 +1404,43 @@ def _assert_literal_predicate_type(expr: LiteralPredicate) -> None: assert_type(In("a", ("a", "b", "c")), In) assert_type(In("a", (1, 2, 3)), In) assert_type(NotIn("a", ("a", "b", "c")), NotIn) + + +WKB_POINT = bytes.fromhex("0101000000000000000000f03f0000000000000040") +WKB_OTHER_POINT = bytes.fromhex("010100000000000000000000400000000000000040") + + +@pytest.fixture +def geo_schema() -> Schema: + return Schema( + NestedField(1, "geom", GeometryType("srid:3857"), required=False), + NestedField(2, "geog", GeographyType("srid:4326", "vincenty"), required=False), + ) + + +@pytest.mark.parametrize("column", ["geom", "geog"]) +@pytest.mark.parametrize( + "predicate, bound_type, matches", + [ + (lambda col: EqualTo(col, WKB_POINT), BoundEqualTo, True), + (lambda col: EqualTo(col, WKB_OTHER_POINT), BoundEqualTo, False), + (lambda col: NotEqualTo(col, WKB_POINT), BoundNotEqualTo, False), + (lambda col: In(col, {WKB_POINT, WKB_OTHER_POINT}), BoundIn, True), + (lambda col: NotIn(col, {WKB_POINT, WKB_OTHER_POINT}), BoundNotIn, False), + (lambda col: IsNull(col), BoundIsNull, False), + (lambda col: NotNull(col), BoundNotNull, True), + ], +) +def test_bind_and_evaluate_geo_predicates( + geo_schema: Schema, column: str, predicate: Any, bound_type: type, matches: bool +) -> None: + expr = predicate(column) + assert isinstance(expr.bind(geo_schema), bound_type) + assert expression_evaluator(geo_schema, expr, case_sensitive=True)(Record(WKB_POINT, WKB_POINT)) is matches + + +@pytest.mark.parametrize("column", ["geom", "geog"]) +@pytest.mark.parametrize("predicate", [LessThan, LessThanOrEqual, GreaterThan, GreaterThanOrEqual]) +def test_bind_geo_ordering_predicate_fails(geo_schema: Schema, column: str, predicate: type[LiteralPredicate]) -> None: + with pytest.raises(ValueError, match=f"{predicate.__name__} is not supported for .* column {column}"): + predicate(column, WKB_POINT).bind(geo_schema) diff --git a/tests/expressions/test_literals.py b/tests/expressions/test_literals.py index 9251e79a7d..612b967db4 100644 --- a/tests/expressions/test_literals.py +++ b/tests/expressions/test_literals.py @@ -36,6 +36,8 @@ FloatAboveMax, FloatBelowMin, FloatLiteral, + GeographyLiteral, + GeometryLiteral, IntAboveMax, IntBelowMin, Literal, @@ -45,6 +47,8 @@ StringLiteral, TimeLiteral, TimestampLiteral, + TimestampNanoLiteral, + TimestamptzNanoLiteral, literal, ) from pyiceberg.types import ( @@ -55,11 +59,15 @@ DoubleType, FixedType, FloatType, + GeographyType, + GeometryType, IntegerType, LongType, PrimitiveType, StringType, + TimestampNanoType, TimestampType, + TimestamptzNanoType, TimestamptzType, TimeType, UUIDType, @@ -387,6 +395,84 @@ def test_invalid_timestamp_with_zone_in_literal() -> None: assert "Invalid timestamp without zone: abc (must be ISO-8601)" in str(e.value) +@pytest.mark.parametrize( + "value, iceberg_type, expected", + [ + ("2024-01-01T00:00:00.123456789", TimestampNanoType(), 1704067200123456789), + ("2024-01-01T00:00:00.1234567", TimestampNanoType(), 1704067200123456700), + ("2024-01-01T00:00:00.123456", TimestampNanoType(), 1704067200123456000), + ("2024-01-01T00:00:00", TimestampNanoType(), 1704067200000000000), + ("1969-12-31T23:59:59.999999999", TimestampNanoType(), -1), + ("2024-01-01T00:00:00.123456789+00:00", TimestamptzNanoType(), 1704067200123456789), + ("2024-01-01T01:00:00.123456789+01:00", TimestamptzNanoType(), 1704067200123456789), + ], +) +def test_string_to_timestamp_nano_literal(value: str, iceberg_type: PrimitiveType, expected: int) -> None: + lit = literal(value).to(iceberg_type) + expected_type = TimestampNanoLiteral if isinstance(iceberg_type, TimestampNanoType) else TimestamptzNanoLiteral + assert type(lit) is expected_type + assert lit.value == expected + + +def test_string_to_timestamp_nano_literal_zone_errors() -> None: + with pytest.raises(ValueError, match="Zone offset provided, but not expected"): + literal("2024-01-01T00:00:00.123456789+00:00").to(TimestampNanoType()) + with pytest.raises(ValueError, match="Missing zone offset"): + literal("2024-01-01T00:00:00.123456789").to(TimestamptzNanoType()) + with pytest.raises(ValueError, match="Invalid timestamp without zone"): + literal("abc").to(TimestampNanoType()) + with pytest.raises(ValueError, match="Invalid timestamp with zone"): + literal("abc").to(TimestamptzNanoType()) + + +@pytest.mark.parametrize("iceberg_type", [TimestampNanoType(), TimestamptzNanoType()]) +def test_timestamp_nano_literal_above_and_below_max(iceberg_type: PrimitiveType) -> None: + zone = "+00:00" if isinstance(iceberg_type, TimestamptzNanoType) else "" + assert literal(f"2262-04-12T00:00:00{zone}").to(iceberg_type) == LongAboveMax() + assert literal(f"1677-09-21T00:00:00{zone}").to(iceberg_type) == LongBelowMin() + assert LongLiteral(LongType.max + 1).to(iceberg_type) == LongAboveMax() + assert LongLiteral(LongType.min - 1).to(iceberg_type) == LongBelowMin() + assert TimestampLiteral(LongType.max // 1_000 + 1).to(iceberg_type) == LongAboveMax() + assert DateLiteral(-106752).to(iceberg_type) == LongBelowMin() + + +def test_timestamp_nano_literal_conversions() -> None: + nanos = 1704067200123456789 + assert LongLiteral(nanos).to(TimestampNanoType()) == TimestampNanoLiteral(nanos) + assert type(LongLiteral(nanos).to(TimestamptzNanoType())) is TimestamptzNanoLiteral + assert TimestampLiteral(1704067200123456).to(TimestampNanoType()) == TimestampNanoLiteral(1704067200123456000) + assert type(TimestampLiteral(1704067200123456).to(TimestamptzNanoType())) is TimestamptzNanoLiteral + assert DateLiteral(19723).to(TimestampNanoType()) == TimestampNanoLiteral(1704067200000000000) + assert type(DateLiteral(19723).to(TimestamptzNanoType())) is TimestamptzNanoLiteral + + lit = TimestampNanoLiteral(nanos) + assert lit.to(TimestampNanoType()) is lit + assert type(lit.to(TimestamptzNanoType())) is TimestamptzNanoLiteral + assert lit.to(TimestampType()) == TimestampLiteral(1704067200123456) + assert lit.to(TimestamptzType()) == TimestampLiteral(1704067200123456) + assert lit.to(DateType()) == DateLiteral(19723) + # Conversion to micros and days floors, including before the epoch + assert TimestampNanoLiteral(-1).to(TimestampType()) == TimestampLiteral(-1) + assert TimestampNanoLiteral(-1).to(DateType()) == DateLiteral(-1) + + tz_lit = TimestamptzNanoLiteral(nanos) + assert tz_lit.to(TimestamptzNanoType()) is tz_lit + assert type(tz_lit.to(TimestampNanoType())) is TimestampNanoLiteral + + assert lit.increment() == TimestampNanoLiteral(nanos + 1) + assert lit.decrement() == TimestampNanoLiteral(nanos - 1) + assert type(tz_lit.increment()) is TimestamptzNanoLiteral + assert type(tz_lit.decrement()) is TimestamptzNanoLiteral + + with pytest.raises(TypeError, match="Cannot convert TimestampNanoLiteral into string"): + lit.to(StringType()) + + +def test_timestamp_nano_literal_serialization() -> None: + assert TimestampNanoLiteral(1704067200123456789).model_dump() == "2024-01-01T00:00:00.123456789" + assert TimestamptzNanoLiteral(1704067200123456789).model_dump() == "2024-01-01T00:00:00.123456789+00:00" + + def test_string_to_uuid_literal() -> None: expected = uuid.uuid4() uuid_str = literal(str(expected)) @@ -1030,3 +1116,31 @@ def test_to_json() -> None: assert_type(literal(bytes([0x01, 0x02, 0x03])), Literal[bytes]) assert_type(literal(Decimal("19.25")), Literal[Decimal]) assert_type({literal(1), literal(2), literal(3)}, set[Literal[int]]) + + +WKB_POINT = bytes.fromhex("0101000000000000000000f03f0000000000000040") + + +def test_binary_to_geometry() -> None: + geometry_lit = literal(WKB_POINT).to(GeometryType("srid:3857")) + assert isinstance(geometry_lit, GeometryLiteral) + assert geometry_lit.value == WKB_POINT + assert geometry_lit.to(GeometryType()) is geometry_lit + assert geometry_lit.to(BinaryType()) == BinaryLiteral(WKB_POINT) + with pytest.raises(TypeError, match="Cannot convert GeometryLiteral into geography"): + geometry_lit.to(GeographyType()) + + +def test_binary_to_geography() -> None: + geography_lit = literal(WKB_POINT).to(GeographyType("srid:4326", "vincenty")) + assert isinstance(geography_lit, GeographyLiteral) + assert geography_lit.value == WKB_POINT + assert geography_lit.to(GeographyType()) is geography_lit + assert geography_lit.to(BinaryType()) == BinaryLiteral(WKB_POINT) + with pytest.raises(TypeError, match="Cannot convert GeographyLiteral into geometry"): + geography_lit.to(GeometryType()) + + +def test_string_to_geometry_is_not_supported() -> None: + with pytest.raises(TypeError, match="Cannot convert StringLiteral into geometry"): + literal("POINT (1 2)").to(GeometryType()) diff --git a/tests/expressions/test_projection.py b/tests/expressions/test_projection.py index 4d0c2c1346..9e4a92551b 100644 --- a/tests/expressions/test_projection.py +++ b/tests/expressions/test_projection.py @@ -49,7 +49,9 @@ LongType, NestedField, StringType, + TimestampNanoType, TimestampType, + TimestamptzNanoType, ) @@ -204,6 +206,38 @@ def test_hour_projection(schema: Schema, hour_spec: PartitionSpec) -> None: assert expected[index] == expr, predicate +@pytest.mark.parametrize("source_type, zone", [(TimestampNanoType(), ""), (TimestamptzNanoType(), "+00:00")]) +def test_day_projection_nano_timestamp(source_type: TimestampNanoType | TimestamptzNanoType, zone: str) -> None: + schema = Schema(NestedField(1, "ts_ns", source_type, required=False)) + spec = PartitionSpec(PartitionField(1, 1000, DayTransform(), "ts_day")) + predicates = [ + LessThan("ts_ns", f"2022-11-27T00:00:00{zone}"), + LessThanOrEqual("ts_ns", f"2022-11-27T00:00:00{zone}"), + GreaterThan("ts_ns", f"2022-11-26T23:59:59.999999999{zone}"), + GreaterThanOrEqual("ts_ns", f"2022-11-26T23:59:59.999999999{zone}"), + EqualTo("ts_ns", f"2022-11-27T10:00:00.000000001{zone}"), + NotEqualTo("ts_ns", f"2022-11-27T10:00:00{zone}"), + In("ts_ns", {f"2022-11-27T00:00:00{zone}", f"2022-11-26T23:59:59.999999999{zone}"}), + NotIn("ts_ns", {f"2022-11-27T00:00:00{zone}"}), + ] + + expected = [ + LessThanOrEqual("ts_day", 19322), + LessThanOrEqual("ts_day", 19323), + GreaterThanOrEqual("ts_day", 19323), + GreaterThanOrEqual("ts_day", 19322), + EqualTo("ts_day", 19323), + AlwaysTrue(), + In("ts_day", {19322, 19323}), + AlwaysTrue(), + ] + + project = inclusive_projection(schema, spec) + for index, predicate in enumerate(predicates): + expr = project(predicate) + assert expected[index] == expr, predicate + + def test_day_projection(schema: Schema, day_spec: PartitionSpec) -> None: predicates = [ NotNull("event_ts"), diff --git a/tests/integration/test_add_files.py b/tests/integration/test_add_files.py index a1d45451d8..d97d722c60 100644 --- a/tests/integration/test_add_files.py +++ b/tests/integration/test_add_files.py @@ -150,7 +150,14 @@ def _create_table( ) -@pytest.fixture(name="format_version", params=[pytest.param(1, id="format_version=1"), pytest.param(2, id="format_version=2")]) +@pytest.fixture( + name="format_version", + params=[ + pytest.param(1, id="format_version=1"), + pytest.param(2, id="format_version=2"), + pytest.param(3, id="format_version=3"), + ], +) def format_version_fixture(request: pytest.FixtureRequest) -> Iterator[int]: """Fixture to run tests with different table format versions.""" yield request.param @@ -716,6 +723,8 @@ def test_add_files_with_large_and_regular_schema(spark: SparkSession, session_ca @pytest.mark.integration def test_add_files_with_timestamp_tz_ns_fails(session_catalog: Catalog, format_version: int, mocker: MockerFixture) -> None: + if format_version >= 3: + pytest.skip("Nanosecond timestamps are supported from format version 3") nanoseconds_schema_iceberg = Schema(NestedField(1, "quux", TimestamptzType())) nanoseconds_schema = pa.schema( @@ -760,7 +769,7 @@ def test_add_files_with_timestamp_tz_ns_fails(session_catalog: Catalog, format_v @pytest.mark.integration -@pytest.mark.parametrize("format_version", [1, 2]) +@pytest.mark.parametrize("format_version", [1, 2, 3]) def test_add_file_with_valid_nullability_diff(spark: SparkSession, session_catalog: Catalog, format_version: int) -> None: identifier = f"default.test_table_with_valid_nullability_diff{format_version}" table_schema = Schema( @@ -799,7 +808,7 @@ def test_add_file_with_valid_nullability_diff(spark: SparkSession, session_catal @pytest.mark.integration -@pytest.mark.parametrize("format_version", [1, 2]) +@pytest.mark.parametrize("format_version", [1, 2, 3]) def test_add_files_with_valid_upcast( spark: SparkSession, session_catalog: Catalog, @@ -851,7 +860,7 @@ def test_add_files_with_valid_upcast( @pytest.mark.integration -@pytest.mark.parametrize("format_version", [1, 2]) +@pytest.mark.parametrize("format_version", [1, 2, 3]) def test_add_files_subset_of_schema(spark: SparkSession, session_catalog: Catalog, format_version: int) -> None: identifier = f"default.test_table_subset_of_schema{format_version}" tbl = _create_table(session_catalog, identifier, format_version) @@ -888,7 +897,7 @@ def test_add_files_subset_of_schema(spark: SparkSession, session_catalog: Catalo @pytest.mark.integration -@pytest.mark.parametrize("format_version", [1, 2]) +@pytest.mark.parametrize("format_version", [1, 2, 3]) def test_add_files_with_duplicate_files_in_file_paths(spark: SparkSession, session_catalog: Catalog, format_version: int) -> None: identifier = f"default.test_table_duplicate_add_files_v{format_version}" tbl = _create_table(session_catalog, identifier, format_version) @@ -902,7 +911,7 @@ def test_add_files_with_duplicate_files_in_file_paths(spark: SparkSession, sessi @pytest.mark.integration -@pytest.mark.parametrize("format_version", [1, 2]) +@pytest.mark.parametrize("format_version", [1, 2, 3]) def test_add_files_that_referenced_by_current_snapshot( spark: SparkSession, session_catalog: Catalog, format_version: int ) -> None: @@ -928,7 +937,7 @@ def test_add_files_that_referenced_by_current_snapshot( @pytest.mark.integration -@pytest.mark.parametrize("format_version", [1, 2]) +@pytest.mark.parametrize("format_version", [1, 2, 3]) def test_add_files_that_referenced_by_current_snapshot_with_check_duplicate_files_false( spark: SparkSession, session_catalog: Catalog, format_version: int ) -> None: @@ -959,7 +968,7 @@ def test_add_files_that_referenced_by_current_snapshot_with_check_duplicate_file @pytest.mark.integration -@pytest.mark.parametrize("format_version", [1, 2]) +@pytest.mark.parametrize("format_version", [1, 2, 3]) def test_add_files_that_referenced_by_current_snapshot_with_check_duplicate_files_true( spark: SparkSession, session_catalog: Catalog, format_version: int ) -> None: diff --git a/tests/integration/test_deletes.py b/tests/integration/test_deletes.py index 20205c59fb..3303718832 100644 --- a/tests/integration/test_deletes.py +++ b/tests/integration/test_deletes.py @@ -60,7 +60,7 @@ def test_table(session_catalog: RestCatalog) -> Generator[Table, None, None]: @pytest.mark.integration -@pytest.mark.parametrize("format_version", [1, 2]) +@pytest.mark.parametrize("format_version", [1, 2, 3]) def test_partitioned_table_delete_full_file(spark: SparkSession, session_catalog: RestCatalog, format_version: int) -> None: identifier = "default.table_partitioned_delete" @@ -95,7 +95,7 @@ def test_partitioned_table_delete_full_file(spark: SparkSession, session_catalog @pytest.mark.integration -@pytest.mark.parametrize("format_version", [1, 2]) +@pytest.mark.parametrize("format_version", [1, 2, 3]) def test_partitioned_table_rewrite(spark: SparkSession, session_catalog: RestCatalog, format_version: int) -> None: identifier = "default.table_partitioned_delete" @@ -130,7 +130,7 @@ def test_partitioned_table_rewrite(spark: SparkSession, session_catalog: RestCat @pytest.mark.integration -@pytest.mark.parametrize("format_version", [1, 2]) +@pytest.mark.parametrize("format_version", [1, 2, 3]) def test_rewrite_partitioned_table_with_null(spark: SparkSession, session_catalog: RestCatalog, format_version: int) -> None: identifier = "default.table_partitioned_delete" @@ -165,7 +165,7 @@ def test_rewrite_partitioned_table_with_null(spark: SparkSession, session_catalo @pytest.mark.integration -@pytest.mark.parametrize("format_version", [1, 2]) +@pytest.mark.parametrize("format_version", [1, 2, 3]) @pytest.mark.filterwarnings("ignore:Delete operation did not match any records") def test_partitioned_table_no_match(spark: SparkSession, session_catalog: RestCatalog, format_version: int) -> None: identifier = "default.table_partitioned_delete" @@ -197,7 +197,7 @@ def test_partitioned_table_no_match(spark: SparkSession, session_catalog: RestCa @pytest.mark.integration -@pytest.mark.filterwarnings("ignore:Merge on read is not yet supported, falling back to copy-on-write") +@pytest.mark.filterwarnings("ignore:Merge-on-read deletes require format version 3") def test_delete_partitioned_table_positional_deletes(spark: SparkSession, session_catalog: RestCatalog) -> None: identifier = "default.table_partitioned_delete" @@ -244,7 +244,7 @@ def test_delete_partitioned_table_positional_deletes(spark: SparkSession, sessio @pytest.mark.integration -@pytest.mark.filterwarnings("ignore:Merge on read is not yet supported, falling back to copy-on-write") +@pytest.mark.filterwarnings("ignore:Merge-on-read deletes require format version 3") def test_delete_partitioned_table_positional_deletes_empty_batch(spark: SparkSession, session_catalog: RestCatalog) -> None: identifier = "default.test_delete_partitioned_table_positional_deletes_empty_batch" @@ -312,7 +312,7 @@ def test_delete_partitioned_table_positional_deletes_empty_batch(spark: SparkSes @pytest.mark.integration -@pytest.mark.filterwarnings("ignore:Merge on read is not yet supported, falling back to copy-on-write") +@pytest.mark.filterwarnings("ignore:Merge-on-read deletes require format version 3") def test_read_multiple_batches_in_task_with_position_deletes(spark: SparkSession, session_catalog: RestCatalog) -> None: identifier = "default.test_read_multiple_batches_in_task_with_position_deletes" @@ -365,7 +365,7 @@ def test_read_multiple_batches_in_task_with_position_deletes(spark: SparkSession @pytest.mark.integration -@pytest.mark.filterwarnings("ignore:Merge on read is not yet supported, falling back to copy-on-write") +@pytest.mark.filterwarnings("ignore:Merge-on-read deletes require format version 3") def test_overwrite_partitioned_table(spark: SparkSession, session_catalog: RestCatalog) -> None: identifier = "default.table_partitioned_delete" @@ -415,7 +415,7 @@ def test_overwrite_partitioned_table(spark: SparkSession, session_catalog: RestC @pytest.mark.integration -@pytest.mark.filterwarnings("ignore:Merge on read is not yet supported, falling back to copy-on-write") +@pytest.mark.filterwarnings("ignore:Merge-on-read deletes require format version 3") def test_partitioned_table_positional_deletes_sequence_number(spark: SparkSession, session_catalog: RestCatalog) -> None: identifier = "default.table_partitioned_delete_sequence_number" @@ -899,7 +899,7 @@ def test_overwrite_with_filter_case_insensitive(test_table: Table) -> None: @pytest.mark.integration -@pytest.mark.parametrize("format_version", [1, 2]) +@pytest.mark.parametrize("format_version", [1, 2, 3]) @pytest.mark.filterwarnings("ignore:Delete operation did not match any records") def test_delete_on_empty_table(spark: SparkSession, session_catalog: RestCatalog, format_version: int) -> None: identifier = f"default.test_delete_on_empty_table_{format_version}" diff --git a/tests/integration/test_inspect_table.py b/tests/integration/test_inspect_table.py index 4d8dfbe9bb..b88f8b6f07 100644 --- a/tests/integration/test_inspect_table.py +++ b/tests/integration/test_inspect_table.py @@ -101,6 +101,10 @@ def _inspect_files_asserts(df: pa.Table, spark_df: DataFrame) -> None: "split_offsets", "equality_ids", "sort_order_id", + "first_row_id", + "referenced_data_file", + "content_offset", + "content_size_in_bytes", "readable_metrics", ] @@ -132,6 +136,10 @@ def _inspect_files_asserts(df: pa.Table, spark_df: DataFrame) -> None: "split_offsets", "equality_ids", "sort_order_id", + "first_row_id", + "referenced_data_file", + "content_offset", + "content_size_in_bytes", ] ] rhs_subset = rhs[ @@ -145,6 +153,10 @@ def _inspect_files_asserts(df: pa.Table, spark_df: DataFrame) -> None: "split_offsets", "equality_ids", "sort_order_id", + "first_row_id", + "referenced_data_file", + "content_offset", + "content_size_in_bytes", ] ] @@ -943,6 +955,10 @@ def inspect_files_asserts(df: pa.Table) -> None: "split_offsets", "equality_ids", "sort_order_id", + "first_row_id", + "referenced_data_file", + "content_offset", + "content_size_in_bytes", "readable_metrics", ] diff --git a/tests/integration/test_reads.py b/tests/integration/test_reads.py index a151d62b82..3dab3af87f 100644 --- a/tests/integration/test_reads.py +++ b/tests/integration/test_reads.py @@ -978,10 +978,21 @@ def test_upgrade_table_version(catalog: Catalog) -> None: transaction.upgrade_table_version(format_version=1) assert "Cannot downgrade v2 table to v1" in str(e.value) + with table_test_table_version.transaction() as transaction: + transaction.upgrade_table_version(format_version=3) + + assert table_test_table_version.format_version == 3 + assert table_test_table_version.metadata.next_row_id == 0 + + # The upgraded metadata must persist and reload from the catalog + reloaded = catalog.load_table("default.test_table_version") + assert reloaded.format_version == 3 + assert reloaded.metadata.next_row_id == 0 + with pytest.raises(ValueError) as e: with table_test_table_version.transaction() as transaction: - transaction.upgrade_table_version(format_version=3) - assert "Unsupported table format version: 3" in str(e.value) + transaction.upgrade_table_version(format_version=4) + assert "Unsupported table format version: 4" in str(e.value) @pytest.mark.integration @@ -1243,8 +1254,9 @@ def test_initial_default(catalog: Catalog, spark: SparkSession) -> None: tbl.append(one_column) - # Do the bump version through Spark, since PyIceberg does not support this (yet) + # Bump the version through Spark, and refresh since default values require format version 3 spark.sql(f"ALTER TABLE {identifier} SET TBLPROPERTIES('format-version'='3')") + tbl.refresh() with tbl.update_schema() as upd: upd.add_column("so_true", BooleanType(), required=False, default_value=True) diff --git a/tests/integration/test_v3_interop.py b/tests/integration/test_v3_interop.py new file mode 100644 index 0000000000..d005e89c22 --- /dev/null +++ b/tests/integration/test_v3_interop.py @@ -0,0 +1,732 @@ +# Licensed to the Apache Software Foundation (ASF) under one +# or more contributor license agreements. See the NOTICE file +# distributed with this work for additional information +# regarding copyright ownership. The ASF licenses this file +# to you under the Apache License, Version 2.0 (the +# "License"); you may not use this file except in compliance +# with the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, +# software distributed under the License is distributed on an +# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +# KIND, either express or implied. See the License for the +# specific language governing permissions and limitations +# under the License. +"""Format-version 3 interoperability tests against tables written by the Java reference implementation. + +Fixtures come from ``dev/provision.py`` (Spark SQL) and ``dev/provision_v3.scala`` (Java Table API, for the +features Spark SQL cannot produce). Each test asserts spec-correct behaviour; the ones marked ``xfail(strict=True)`` +document known gaps and must be un-marked when the gap is closed. +""" + +import json +import uuid +from collections.abc import Iterator +from datetime import datetime, timezone +from decimal import Decimal + +import pyarrow as pa +import pyarrow.parquet as pq +import pytest + +from pyiceberg.catalog import Catalog +from pyiceberg.expressions import EqualTo, GreaterThanOrEqual, LessThan +from pyiceberg.io.pyarrow import PyArrowFileIO +from pyiceberg.schema import Schema +from pyiceberg.table import Table +from pyiceberg.table.deletion_vector import DeletionVector +from pyiceberg.table.puffin import PuffinFile +from pyiceberg.table.snapshots import Operation +from pyiceberg.transforms import BucketTransform, MonthTransform, YearTransform +from pyiceberg.types import ( + GeographyType, + GeometryType, + IntegerType, + NestedField, + StringType, + TimestampNanoType, + TimestamptzNanoType, + UnknownType, + VariantType, +) +from pyiceberg.utils.geo import GeospatialBound +from pyiceberg.variant import to_json as variant_to_json +from pyiceberg.variant import to_python as variant_to_python + +NS_VALUES = [1704067200123456789, 1706832000000000001, None] +NS_PARTITIONS = [ + # (id, year(ts), month(tstz), bucket(tstz, 4)) as computed by the Java transforms + (1, 53, 647, 0), + (2, 54, 648, 0), + (3, 54, 652, 1), +] + + +@pytest.fixture +def probe_identifier(session_catalog: Catalog) -> Iterator[str]: + name = f"default.test_v3_probe_{uuid.uuid4().hex[:8]}" + yield name + try: + session_catalog.drop_table(name) + except Exception: # noqa: BLE001 + pass + + +# ---------------------------------------------------------------------------- nanosecond timestamps + + +@pytest.mark.integration +def test_v3_ns_timestamps_read(session_catalog: Catalog) -> None: + table = session_catalog.load_table("default.test_v3_ns_timestamps") + assert table.metadata.format_version == 3 + assert isinstance(table.schema().find_type("ts_ns"), TimestampNanoType) + assert isinstance(table.schema().find_type("tstz_ns"), TimestamptzNanoType) + arrow = table.scan().to_arrow().sort_by("id") + assert arrow.schema.field("ts_ns").type == pa.timestamp("ns") + assert arrow.schema.field("tstz_ns").type == pa.timestamp("ns", tz="UTC") + assert arrow.column("ts_ns").cast(pa.int64()).to_pylist() == NS_VALUES + assert arrow.column("tstz_ns").cast(pa.int64()).to_pylist() == NS_VALUES + assert arrow.column("ts_us").to_pylist() == [datetime(2024, 1, 1, 0, 0, 0, 123456), datetime(2024, 2, 2), None] + + +@pytest.mark.integration +@pytest.mark.parametrize( + "row_filter, expected", + [ + (GreaterThanOrEqual("ts_ns", "2024-02-01T00:00:00"), [2]), + (LessThan("ts_ns", "2024-01-01T00:00:00.500000"), [1]), + (GreaterThanOrEqual("ts_ns", datetime(2024, 2, 1)), [2]), + (GreaterThanOrEqual("tstz_ns", datetime(2024, 2, 1, tzinfo=timezone.utc)), [2]), + (GreaterThanOrEqual("tstz_ns", "2024-02-01T00:00:00+00:00"), [2]), + ], +) +def test_v3_ns_timestamp_filters(session_catalog: Catalog, row_filter: object, expected: list[int]) -> None: + table = session_catalog.load_table("default.test_v3_ns_timestamps") + assert sorted(r["id"] for r in table.scan(row_filter=row_filter).to_arrow().to_pylist()) == expected # type: ignore[arg-type] + + +@pytest.mark.integration +def test_v3_ns_partition_transforms_match_java(session_catalog: Catalog) -> None: + table = session_catalog.load_table("default.test_v3_ns_partitions") + schema = table.schema() + arrow = table.scan().to_arrow().sort_by("id") + ts = arrow.column("ts").cast(pa.int64()).to_pylist() + tstz = arrow.column("tstz").cast(pa.int64()).to_pylist() + for (row_id, year, month, bucket), ts_value, tstz_value in zip(NS_PARTITIONS, ts, tstz, strict=True): + assert YearTransform().transform(schema.find_type("ts"))(ts_value) == year, row_id + assert MonthTransform().transform(schema.find_type("tstz"))(tstz_value) == month, row_id + assert BucketTransform(4).transform(schema.find_type("tstz"))(tstz_value) == bucket, row_id + files = table.inspect.files().to_pylist() + partitions = sorted((f["partition"]["ts_year"], f["partition"]["tstz_month"], f["partition"]["tstz_bucket"]) for f in files) + assert partitions == sorted((y, m, b) for _, y, m, b in NS_PARTITIONS) + + +@pytest.mark.integration +def test_v3_ns_partition_pruning(session_catalog: Catalog) -> None: + table = session_catalog.load_table("default.test_v3_ns_partitions") + scan = table.scan(row_filter=GreaterThanOrEqual("ts", "2024-01-01T00:00:00")) + assert len(list(scan.plan_files())) == 2 + assert sorted(r["id"] for r in scan.to_arrow().to_pylist()) == [2, 3] + + +# ---------------------------------------------------------------------------- geometry / geography + + +@pytest.mark.integration +def test_v3_geo_schema_from_java(session_catalog: Catalog) -> None: + table = session_catalog.load_table("default.test_v3_geo") + schema = table.schema() + assert isinstance(schema.find_type("geom"), GeometryType) + assert schema.find_type("geom_srid") == GeometryType("srid:3857") + assert schema.find_type("geog") == GeographyType() + assert schema.find_type("geog_v") == GeographyType("srid:4326", "vincenty") + + +@pytest.mark.integration +def test_v3_geo_types_created_through_rest_match_java_serialization(session_catalog: Catalog, probe_identifier: str) -> None: + schema = Schema( + NestedField(1, "id", IntegerType(), required=True), + NestedField(2, "geom_srid", GeometryType("srid:3857"), required=False), + NestedField(3, "geog_v", GeographyType("srid:4326", "vincenty"), required=False), + ) + session_catalog.create_table(probe_identifier, schema, properties={"format-version": "3"}) + loaded = session_catalog.load_table(probe_identifier) + assert loaded.schema().find_type("geom_srid") == GeometryType("srid:3857") + assert loaded.schema().find_type("geog_v") == GeographyType("srid:4326", "vincenty") + + +@pytest.mark.integration +def test_v3_geo_metadata_written_by_java_server_uses_spec_type_strings(session_catalog: Catalog, probe_identifier: str) -> None: + schema = Schema( + NestedField(1, "id", IntegerType(), required=True), + NestedField(2, "geom", GeometryType(), required=False), + NestedField(3, "geom_srid", GeometryType("srid:3857"), required=False), + NestedField(4, "geog_v", GeographyType("srid:4326", "vincenty"), required=False), + ) + table = session_catalog.create_table(probe_identifier, schema, properties={"format-version": "3"}) + io = PyArrowFileIO( + properties={"s3.endpoint": "http://localhost:9000", "s3.access-key-id": "admin", "s3.secret-access-key": "password"} + ) + with io.new_input(table.metadata_location).open() as f: + metadata = json.loads(f.read()) + # The metadata file is written by the Java REST server, so it shows how Java understood the type strings + types = {field["name"]: field["type"] for field in metadata["schemas"][-1]["fields"]} + assert types == { + "id": "int", + "geom": "geometry", + "geom_srid": "geometry(srid:3857)", + "geog_v": "geography(srid:4326, vincenty)", + } + + +# ---------------------------------------------------------------------------- default values + + +@pytest.mark.integration +def test_v3_initial_defaults_applied_on_read(session_catalog: Catalog, spark: "SparkSession") -> None: # type: ignore[name-defined] # noqa: F821 + table = session_catalog.load_table("default.test_v3_defaults") + assert table.schema().find_field("color").initial_default == "blue" + assert table.schema().find_field("color").write_default == "green" + rows = table.scan().to_arrow().sort_by("id").to_pylist() + spark_rows = spark.sql("SELECT * FROM rest.default.test_v3_defaults ORDER BY id").collect() + assert len(rows) == len(spark_rows) == 2 + for row, spark_row in zip(rows, spark_rows, strict=True): + assert row["color"] == spark_row["color"] == "blue" + assert row["qty"] == spark_row["qty"] == 42 + assert row["ratio"] == spark_row["ratio"] == 1.5 + assert str(row["d"]) == str(spark_row["d"]) == "2024-03-04" + assert row["ts"] == datetime(2024, 3, 4, 5, 6, 7, tzinfo=timezone.utc) + assert str(row["dec"]) == "12.34" + assert row["b"] is True + assert row["bin"] == b"\x01\x02" + assert str(row["u"]) == "f79c3e09-677c-4bbd-a479-3f349cb785e7" + assert row["req"] == 7 + assert table.scan(row_filter=EqualTo("color", "blue")).count() == 2 + + +@pytest.mark.integration +def test_v3_nested_struct_default_applied_on_read(session_catalog: Catalog) -> None: + table = session_catalog.load_table("default.test_v3_nested_defaults") + assert table.scan().to_arrow().to_pylist() == [{"id": 1, "s": {"a": 1, "b": "x", "c": 99}}] + + +# ---------------------------------------------------------------------------- unknown + + +@pytest.mark.integration +def test_v3_unknown_column_reads_as_null(session_catalog: Catalog) -> None: + table = session_catalog.load_table("default.test_v3_unknown") + assert isinstance(table.schema().find_type("unk"), UnknownType) + assert table.scan().to_arrow().sort_by("id").to_pylist() == [{"id": 1, "unk": None}, {"id": 2, "unk": None}] + + +# ---------------------------------------------------------------------------- variant + + +@pytest.mark.integration +def test_v3_variant_table_loads(session_catalog: Catalog, spark: "SparkSession", probe_identifier: str) -> None: # type: ignore[name-defined] # noqa: F821 + spark.sql(f"CREATE TABLE rest.{probe_identifier} (id INT, v VARIANT) USING iceberg TBLPROPERTIES ('format-version'='3')") + spark.sql(f"""INSERT INTO rest.{probe_identifier} SELECT 1, parse_json('{{"a": 1}}')""") + table = session_catalog.load_table(probe_identifier) + assert isinstance(table.schema().find_type("v"), VariantType) + assert table.scan().count() == 1 + (row,) = table.scan().to_arrow().to_pylist() + assert json.loads(variant_to_json(row["v"]["metadata"], row["v"]["value"])) == {"a": 1} + + spark.sql( + f"""INSERT INTO rest.{probe_identifier} VALUES + (2, parse_json('[1, "two", null, 3.5]')), + (3, parse_json('"just a string"')), + (4, parse_json('null')), + (5, NULL), + (6, parse_json('{{"nested": {{"b": [true, false]}}, "s": "x"}}'))""" + ) + table = session_catalog.load_table(probe_identifier) + rows = {row["id"]: row["v"] for row in table.scan().to_arrow().to_pylist()} + decoded = {id_: None if v is None else variant_to_python(v["metadata"], v["value"]) for id_, v in rows.items()} + assert decoded == { + 1: {"a": 1}, + 2: [1, "two", None, Decimal("3.5")], + 3: "just a string", + 4: None, + 5: None, + 6: {"nested": {"b": [True, False]}, "s": "x"}, + } + # the JSON matches Spark's rendering of the same values + expected_json = {r["id"]: r["j"] for r in spark.sql(f"SELECT id, to_json(v) AS j FROM rest.{probe_identifier}").collect()} + for id_, v in rows.items(): + if v is not None: + assert json.loads(variant_to_json(v["metadata"], v["value"])) == json.loads(expected_json[id_]) + + # variant columns cannot be written until pyarrow can annotate the Parquet VARIANT logical type + with pytest.raises(NotImplementedError, match="Writing variant is not supported"): + table.append(table.scan().to_arrow()) + + +# ---------------------------------------------------------------------------- equality deletes + + +@pytest.mark.integration +def test_v3_equality_deletes_match_spark(session_catalog: Catalog, spark: "SparkSession") -> None: # type: ignore[name-defined] # noqa: F821 + identifier = "default.test_v3_equality_deletes" + table = session_catalog.load_table(identifier) + delete_files = table.inspect.delete_files().to_pylist() + assert {d["content"] for d in delete_files} == {2} + assert sorted(d["equality_ids"] for d in delete_files) == [[1], [2]] + + expected = sorted(tuple(r) for r in spark.sql(f"SELECT id, name, value FROM rest.{identifier}").collect()) + # id 2 and 5 are deleted by id, id 4 by its null name; the row added after the deletes is kept + assert expected == [(1, "a", 1.0), (2, "b2", 20.0), (3, "c", 3.0)] + actual = sorted((r["id"], r["name"], r["value"]) for r in table.scan().to_arrow().to_pylist()) + assert actual == expected + assert table.scan().count() == len(expected) + + # the delete columns are applied even when they are not projected + assert sorted(r["value"] for r in table.scan(selected_fields=("value",)).to_arrow().to_pylist()) == [1.0, 3.0, 20.0] + assert [r["id"] for r in table.scan(row_filter=EqualTo("id", 2)).to_arrow().to_pylist()] == [2] + + +# ---------------------------------------------------------------------------- row lineage and deletion vectors + + +@pytest.mark.integration +def test_v3_deletion_vectors_match_spark(session_catalog: Catalog, spark: "SparkSession") -> None: # type: ignore[name-defined] # noqa: F821 + table = session_catalog.load_table("default.test_positional_mor_deletes_v3") + assert table.metadata.format_version == 3 + delete_files = table.inspect.delete_files().to_pylist() + assert delete_files and {d["file_format"] for d in delete_files} == {"PUFFIN"} + expected = sorted(r["number"] for r in spark.sql("SELECT number FROM rest.default.test_positional_mor_deletes_v3").collect()) + assert sorted(r["number"] for r in table.scan().to_arrow().to_pylist()) == expected + assert table.metadata.next_row_id == sum(s.added_rows or 0 for s in table.snapshots()) + + +@pytest.mark.integration +@pytest.mark.parametrize( + "identifier", ["default.test_positional_mor_deletes_v3", "default.test_positional_mor_double_deletes_v3"] +) +def test_v3_row_lineage_columns_match_spark(session_catalog: Catalog, spark: "SparkSession", identifier: str) -> None: # type: ignore[name-defined] # noqa: F821 + table = session_catalog.load_table(identifier) + spark_rows = spark.sql(f"SELECT number, _row_id, _last_updated_sequence_number FROM rest.{identifier}").collect() + expected = {r["number"]: (r["_row_id"], r["_last_updated_sequence_number"]) for r in spark_rows} + scan = table.scan(selected_fields=("number", "_row_id", "_last_updated_sequence_number")) + got = {r["number"]: (r["_row_id"], r["_last_updated_sequence_number"]) for r in scan.to_arrow().to_pylist()} + assert got == expected + + +@pytest.mark.integration +def test_v3_first_row_id_inherited(session_catalog: Catalog) -> None: + table = session_catalog.load_table("default.test_positional_mor_deletes_v3") + snapshot = table.current_snapshot() + assert snapshot is not None + for manifest in snapshot.manifests(table.io): + if manifest.content.name == "DATA": + assert manifest.first_row_id is not None + for entry in manifest.fetch_manifest_entry(table.io): + assert entry.data_file.first_row_id is not None + + +@pytest.mark.integration +def test_v3_inspect_exposes_v3_fields(session_catalog: Catalog) -> None: + table = session_catalog.load_table("default.test_positional_mor_deletes_v3") + delete_files = table.inspect.delete_files().to_pylist() + assert {"referenced_data_file", "content_offset", "content_size_in_bytes"} <= set(delete_files[0]) + assert all(d["referenced_data_file"] and d["content_offset"] is not None for d in delete_files) + entries = table.inspect.entries().to_pylist() + assert "first_row_id" in entries[0]["data_file"] + + +# ---------------------------------------------------------------------------- writes + + +@pytest.mark.integration +def test_v3_append_to_rest_table(session_catalog: Catalog, probe_identifier: str) -> None: + schema = Schema(NestedField(1, "id", IntegerType(), required=True)) + table = session_catalog.create_table(probe_identifier, schema, properties={"format-version": "3"}) + assert table.metadata.next_row_id == 0 + df = pa.Table.from_arrays([pa.array([1, 2], pa.int32())], schema=pa.schema([pa.field("id", pa.int32(), nullable=False)])) + table.append(df) + table = session_catalog.load_table(probe_identifier) + assert table.scan().count() == 2 + snapshot = table.current_snapshot() + assert snapshot is not None and snapshot.first_row_id == 0 and snapshot.added_rows == 2 + assert table.metadata.next_row_id == 2 + + +@pytest.mark.integration +def test_v3_upgrade_from_v2(session_catalog: Catalog, spark: "SparkSession", probe_identifier: str) -> None: # type: ignore[name-defined] # noqa: F821 + schema = Schema(NestedField(1, "id", IntegerType(), required=True)) + session_catalog.create_table(probe_identifier, schema, properties={"format-version": "2"}) + spark.sql(f"INSERT INTO rest.{probe_identifier} VALUES (1)") + table = session_catalog.load_table(probe_identifier) + with table.transaction() as transaction: + transaction.upgrade_table_version(3) + table = session_catalog.load_table(probe_identifier) + assert table.metadata.format_version == 3 + assert table.metadata.next_row_id == 0 + spark.sql(f"INSERT INTO rest.{probe_identifier} VALUES (2)") + # The pre-upgrade manifest has no first-row-id, so the first v3 commit assigns row ids to its row as well + assert session_catalog.load_table(probe_identifier).metadata.next_row_id == 2 + + +def _ids(values: list[int]) -> pa.Table: + return pa.Table.from_arrays([pa.array(values, pa.int32())], schema=pa.schema([pa.field("id", pa.int32(), nullable=False)])) + + +@pytest.mark.integration +def test_v3_pyiceberg_appends_read_by_spark_with_row_lineage( + session_catalog: Catalog, + spark: "SparkSession", # type: ignore[name-defined] # noqa: F821 + probe_identifier: str, +) -> None: + schema = Schema(NestedField(1, "id", IntegerType(), required=True)) + table = session_catalog.create_table(probe_identifier, schema, properties={"format-version": "3"}) + table.append(_ids([10, 11, 12])) + table.append(_ids([13, 14])) + + table = session_catalog.load_table(probe_identifier) + assert table.metadata.next_row_id == 5 + first, second = table.snapshots() + assert (first.first_row_id, first.added_rows, second.first_row_id, second.added_rows) == (0, 3, 3, 2) + + rows = spark.sql(f"SELECT id, _row_id, _last_updated_sequence_number FROM rest.{probe_identifier} ORDER BY id").collect() + assert [(r["id"], r["_row_id"]) for r in rows] == [(10, 0), (11, 1), (12, 2), (13, 3), (14, 4)] + assert [r["_last_updated_sequence_number"] for r in rows] == [first.sequence_number] * 3 + [second.sequence_number] * 2 + + # Spark can keep writing to the table PyIceberg wrote + spark.sql(f"DELETE FROM rest.{probe_identifier} WHERE id = 11") + rows = spark.sql(f"SELECT id, _row_id FROM rest.{probe_identifier} ORDER BY id").collect() + assert [(r["id"], r["_row_id"]) for r in rows] == [(10, 0), (12, 2), (13, 3), (14, 4)] + table = session_catalog.load_table(probe_identifier) + assert sorted(table.scan().to_arrow()["id"].to_pylist()) == [10, 12, 13, 14] + assert table.metadata.next_row_id >= 5 + + +@pytest.mark.integration +def test_v3_pyiceberg_copy_on_write_read_by_spark( + session_catalog: Catalog, + spark: "SparkSession", # type: ignore[name-defined] # noqa: F821 + probe_identifier: str, +) -> None: + schema = Schema(NestedField(1, "id", IntegerType(), required=True)) + table = session_catalog.create_table(probe_identifier, schema, properties={"format-version": "3"}) + table.append(_ids([1, 2, 3])) + table.append(_ids([4, 5])) + first_sequence_number, second_sequence_number = (s.sequence_number for s in table.snapshots()) + # Rewrites the first data file: the copied rows keep their row ids and last updated sequence numbers + table.delete("id = 2") + rows = spark.sql(f"SELECT id, _row_id, _last_updated_sequence_number FROM rest.{probe_identifier} ORDER BY id").collect() + assert [(r["id"], r["_row_id"]) for r in rows] == [(1, 0), (3, 2), (4, 3), (5, 4)] + assert [r["_last_updated_sequence_number"] for r in rows] == [first_sequence_number] * 2 + [second_sequence_number] * 2 + + table.overwrite(_ids([6]), overwrite_filter="id = 5") + table = session_catalog.load_table(probe_identifier) + assert table.metadata.next_row_id == sum(s.added_rows or 0 for s in table.snapshots()) + rows = spark.sql(f"SELECT id, _row_id, _last_updated_sequence_number FROM rest.{probe_identifier} ORDER BY id").collect() + assert [r["id"] for r in rows] == [1, 3, 4, 6] + assert [r["_row_id"] for r in rows[:3]] == [0, 2, 3] + assert [r["_last_updated_sequence_number"] for r in rows[:3]] == [first_sequence_number] * 2 + [second_sequence_number] + assert rows[3]["_row_id"] >= 5 + assert all(r["_row_id"] < table.metadata.next_row_id for r in rows) + + +@pytest.mark.integration +def test_v3_hive_table_written_by_pyiceberg_read_by_spark( + session_catalog_hive: Catalog, + spark: "SparkSession", # type: ignore[name-defined] # noqa: F821 +) -> None: + # Hive stores the metadata PyIceberg serializes, unlike REST where the server builds it + identifier = f"default.test_v3_hive_{uuid.uuid4().hex[:8]}" + schema = Schema(NestedField(1, "id", IntegerType(), required=True)) + table = session_catalog_hive.create_table(identifier, schema, properties={"format-version": "3"}) + try: + assert table.metadata.next_row_id == 0 + table.append(_ids([1, 2])) + table.append(_ids([3])) + assert session_catalog_hive.load_table(identifier).metadata.next_row_id == 3 + + properties = {row.key: row.value for row in spark.sql(f"SHOW TBLPROPERTIES hive.{identifier}").collect()} + assert properties["format-version"] == "3" + rows = spark.sql(f"SELECT id, _row_id FROM hive.{identifier} ORDER BY id").collect() + assert [(r["id"], r["_row_id"]) for r in rows] == [(1, 0), (2, 1), (3, 2)] + finally: + session_catalog_hive.drop_table(identifier) + + +@pytest.mark.integration +def test_v3_upgrade_then_pyiceberg_append_read_by_spark( + session_catalog: Catalog, + spark: "SparkSession", # type: ignore[name-defined] # noqa: F821 + probe_identifier: str, +) -> None: + schema = Schema(NestedField(1, "id", IntegerType(), required=True)) + table = session_catalog.create_table(probe_identifier, schema, properties={"format-version": "2"}) + table.append(_ids([1, 2])) + with table.transaction() as transaction: + transaction.upgrade_table_version(3) + + table = session_catalog.load_table(probe_identifier) + table.append(_ids([3])) + table = session_catalog.load_table(probe_identifier) + assert table.metadata.format_version == 3 + snapshot = table.current_snapshot() + assert snapshot is not None and (snapshot.first_row_id, snapshot.added_rows) == (0, 3) + assert table.metadata.next_row_id == 3 + + rows = spark.sql(f"SELECT id, _row_id FROM rest.{probe_identifier} ORDER BY id").collect() + assert [r["id"] for r in rows] == [1, 2, 3] + assert sorted(r["_row_id"] for r in rows) == [0, 1, 2] + + +# ---------------------------------------------------------------------------- PyIceberg-written geo, defaults, unknown + +WKB_POINT_1_2 = bytes.fromhex("0101000000000000000000f03f0000000000000040") +WKB_POINT_NEG_3_5 = bytes.fromhex("010100000000000000000008c00000000000001440") + + +@pytest.mark.integration +def test_v3_pyiceberg_geo_write_round_trip(session_catalog: Catalog, probe_identifier: str) -> None: + ga = pytest.importorskip("geoarrow.pyarrow") + # Spark 4.0 with Iceberg 1.11 cannot read geometry columns, so the files are verified through PyIceberg and Parquet + schema = Schema( + NestedField(1, "id", IntegerType(), required=True), + NestedField(2, "geom", GeometryType(), required=False), + NestedField(3, "geom_srid", GeometryType("srid:3857"), required=False), + NestedField(4, "geog", GeographyType(), required=False), + ) + table = session_catalog.create_table(probe_identifier, schema, properties={"format-version": "3"}) + id_field = pa.field("id", pa.int32(), nullable=False) + plain = pa.array([WKB_POINT_1_2, None], pa.binary()) + plain_schema = pa.schema( + [id_field, pa.field("geom", pa.binary()), pa.field("geom_srid", pa.binary()), pa.field("geog", pa.binary())] + ) + table.append(pa.Table.from_arrays([pa.array([1, 2], pa.int32()), plain, plain, plain], schema=plain_schema)) + + wkb = pa.array([WKB_POINT_1_2, WKB_POINT_NEG_3_5], pa.binary()) + geo_arrays = [ + pa.ExtensionArray.from_storage(ga.wkb(), wkb), + pa.ExtensionArray.from_storage(ga.wkb().with_crs("srid:3857"), wkb), + pa.ExtensionArray.from_storage(ga.wkb().with_edge_type(ga.EdgeType.SPHERICAL), wkb), + ] + geo_schema = pa.schema( + [id_field, *[pa.field(name, array.type) for name, array in zip(["geom", "geom_srid", "geog"], geo_arrays, strict=True)]] + ) + table.append(pa.Table.from_arrays([pa.array([3, 4], pa.int32()), *geo_arrays], schema=geo_schema)) + + table = session_catalog.load_table(probe_identifier) + result = table.scan().to_arrow().sort_by("id") + for name in ("geom", "geom_srid", "geog"): + column = result.column(name) + assert column.cast(column.type.storage_type).to_pylist() == [WKB_POINT_1_2, None, WKB_POINT_1_2, WKB_POINT_NEG_3_5] + + for task in table.scan().plan_files(): + with table.io.new_input(task.file.file_path).open() as f: + parquet_schema = str(pq.ParquetFile(f).schema) + # The Parquet logical types carry the CRS; the default OGC:CRS84 is written as an empty CRS + assert "field_id=2 geom (Geometry(crs=))" in parquet_schema + assert "field_id=3 geom_srid (Geometry(crs=srid:3857))" in parquet_schema + assert "field_id=4 geog (Geography(crs=, algorithm=spherical))" in parquet_schema + + bounds = [ + ( + GeospatialBound.from_bytes(row["readable_metrics"]["geom_srid"]["lower_bound"]), + GeospatialBound.from_bytes(row["readable_metrics"]["geom_srid"]["upper_bound"]), + ) + for row in table.inspect.files().to_pylist() + ] + bounds.sort(key=lambda bound: bound[0].x) + assert bounds == [ + (GeospatialBound(-3.0, 2.0), GeospatialBound(1.0, 5.0)), + (GeospatialBound(1.0, 2.0), GeospatialBound(1.0, 2.0)), + ] + + +@pytest.mark.integration +def test_v3_pyiceberg_write_default_read_by_spark( + session_catalog: Catalog, + spark: "SparkSession", # type: ignore[name-defined] # noqa: F821 + probe_identifier: str, +) -> None: + schema = Schema(NestedField(1, "id", IntegerType(), required=True)) + table = session_catalog.create_table(probe_identifier, schema, properties={"format-version": "3"}) + table.append(_ids([1])) + with table.update_schema() as update: + update.add_column("color", StringType(), default_value="blue") + with table.update_schema() as update: + update.set_default_value("color", "green") + # The dataframe does not have the column, so the write-default is written + table.append(_ids([2])) + + table = session_catalog.load_table(probe_identifier) + assert [row["color"] for row in table.scan().to_arrow().sort_by("id").to_pylist()] == ["blue", "green"] + rows = spark.sql(f"SELECT id, color FROM rest.{probe_identifier} ORDER BY id").collect() + assert [(r["id"], r["color"]) for r in rows] == [(1, "blue"), (2, "green")] + + +@pytest.mark.integration +def test_v3_pyiceberg_unknown_column_append_read_by_spark( + session_catalog: Catalog, + spark: "SparkSession", # type: ignore[name-defined] # noqa: F821 + probe_identifier: str, +) -> None: + schema = Schema(NestedField(1, "id", IntegerType(), required=True), NestedField(2, "unk", UnknownType(), required=False)) + table = session_catalog.create_table(probe_identifier, schema, properties={"format-version": "3"}) + table.append(_ids([1, 2])) + + table = session_catalog.load_table(probe_identifier) + assert table.scan().to_arrow().sort_by("id").to_pylist() == [{"id": 1, "unk": None}, {"id": 2, "unk": None}] + (task,) = table.scan().plan_files() + with table.io.new_input(task.file.file_path).open() as f: + # Columns of type unknown are not stored in data files + assert pq.ParquetFile(f).schema_arrow.names == ["id"] + rows = spark.sql(f"SELECT id, unk FROM rest.{probe_identifier} ORDER BY id").collect() + assert [(r["id"], r["unk"]) for r in rows] == [(1, None), (2, None)] + + +# ---------------------------------------------------------------------------- deletion vector writes + +_DELETE_FILE_COLUMNS = ["content", "file_path", "file_format", "record_count", "referenced_data_file", "content_offset"] + + +def _spark_delete_files(spark: "SparkSession", identifier: str) -> list[dict[str, object]]: # type: ignore[name-defined] # noqa: F821 + rows = spark.sql(f"SELECT * FROM rest.{identifier}.delete_files").collect() + return sorted(({column: row[column] for column in _DELETE_FILE_COLUMNS + ["content_size_in_bytes"]} for row in rows), key=str) + + +def _pyiceberg_delete_files(table: Table) -> list[dict[str, object]]: + rows = table.inspect.delete_files().to_pylist() + return sorted(({column: row[column] for column in _DELETE_FILE_COLUMNS + ["content_size_in_bytes"]} for row in rows), key=str) + + +@pytest.mark.integration +def test_v3_pyiceberg_merge_on_read_deletes_read_by_spark( + session_catalog: Catalog, + spark: "SparkSession", # type: ignore[name-defined] # noqa: F821 + probe_identifier: str, +) -> None: + schema = Schema(NestedField(1, "id", IntegerType(), required=True)) + table = session_catalog.create_table( + probe_identifier, schema, properties={"format-version": "3", "write.delete.mode": "merge-on-read"} + ) + table.append(_ids([1, 2, 3, 4, 5])) + data_file = next(iter(table.scan().plan_files())).file + row_ids_before = {r["id"]: r["_row_id"] for r in table.scan(selected_fields=("id", "_row_id")).to_arrow().to_pylist()} + + table.delete(EqualTo("id", 2)) + + table = session_catalog.load_table(probe_identifier) + snapshot = table.current_snapshot() + assert snapshot is not None and snapshot.summary is not None + assert snapshot.summary.operation == Operation.DELETE + assert snapshot.summary["added-dvs"] == "1" + assert snapshot.summary["added-position-deletes"] == "1" + (dv,) = table.inspect.delete_files().to_pylist() + assert dv["file_format"] == "PUFFIN" + assert dv["referenced_data_file"] == data_file.file_path + assert dv["content_offset"] == 4 + assert dv["content_size_in_bytes"] > 0 + assert dv["record_count"] == 1 + assert _spark_delete_files(spark, probe_identifier) == _pyiceberg_delete_files(table) + + spark_rows = spark.sql(f"SELECT id, _row_id FROM rest.{probe_identifier} ORDER BY id").collect() + assert [r["id"] for r in spark_rows] == [1, 3, 4, 5] + # The data file is not rewritten, so the surviving rows keep their row ids + assert {r["id"]: r["_row_id"] for r in spark_rows} == {i: row_ids_before[i] for i in [1, 3, 4, 5]} + row_ids_after = {r["id"]: r["_row_id"] for r in table.scan(selected_fields=("id", "_row_id")).to_arrow().to_pylist()} + assert row_ids_after == {i: row_ids_before[i] for i in [1, 3, 4, 5]} + assert table.scan().count() == spark.sql(f"SELECT COUNT(*) AS n FROM rest.{probe_identifier}").collect()[0]["n"] == 4 + + # A second delete on the same data file replaces its DV with the union of the deleted positions + table.delete("id in (4, 5)") + + table = session_catalog.load_table(probe_identifier) + snapshot = table.current_snapshot() + assert snapshot is not None and snapshot.summary is not None + assert snapshot.summary["added-dvs"] == "1" + assert snapshot.summary["removed-dvs"] == "1" + assert snapshot.summary["removed-delete-files"] == "1" + assert snapshot.summary["total-delete-files"] == "1" + assert snapshot.summary["total-position-deletes"] == "3" + (replacement,) = table.inspect.delete_files().to_pylist() + assert replacement["file_path"] != dv["file_path"] + assert replacement["referenced_data_file"] == data_file.file_path + assert replacement["record_count"] == 3 + assert _spark_delete_files(spark, probe_identifier) == _pyiceberg_delete_files(table) + assert [r["id"] for r in spark.sql(f"SELECT id FROM rest.{probe_identifier} ORDER BY id").collect()] == [1, 3] + assert table.scan().count() == spark.sql(f"SELECT COUNT(*) AS n FROM rest.{probe_identifier}").collect()[0]["n"] == 2 + + +@pytest.mark.integration +def test_v3_pyiceberg_delete_merges_spark_deletion_vector( + session_catalog: Catalog, + spark: "SparkSession", # type: ignore[name-defined] # noqa: F821 + probe_identifier: str, +) -> None: + spark.sql( + f""" + CREATE TABLE rest.{probe_identifier} (id BIGINT) + USING iceberg + TBLPROPERTIES ('format-version' = '3', 'write.delete.mode' = 'merge-on-read') + """ + ) + spark.range(1, 21).coalesce(1).writeTo(f"rest.{probe_identifier}").append() + spark.sql(f"DELETE FROM rest.{probe_identifier} WHERE id IN (3, 4)") + + table = session_catalog.load_table(probe_identifier) + (spark_dv,) = table.inspect.delete_files().to_pylist() + table.delete("id = 10 or id = 11") + + table = session_catalog.load_table(probe_identifier) + (dv,) = table.inspect.delete_files().to_pylist() + assert dv["file_path"] != spark_dv["file_path"] + assert dv["referenced_data_file"] == spark_dv["referenced_data_file"] + assert dv["record_count"] == 4 + snapshot = table.current_snapshot() + assert snapshot is not None and snapshot.summary is not None + assert snapshot.summary["removed-dvs"] == "1" + assert snapshot.summary["total-position-deletes"] == "4" + + expected = [i for i in range(1, 21) if i not in (3, 4, 10, 11)] + assert [r["id"] for r in spark.sql(f"SELECT id FROM rest.{probe_identifier} ORDER BY id").collect()] == expected + assert sorted(table.scan().to_arrow()["id"].to_pylist()) == expected + assert _spark_delete_files(spark, probe_identifier) == _pyiceberg_delete_files(table) + + # Spark keeps deleting on top of the PyIceberg DV + spark.sql(f"DELETE FROM rest.{probe_identifier} WHERE id = 20") + table = session_catalog.load_table(probe_identifier) + (dv,) = table.inspect.delete_files().to_pylist() + assert dv["record_count"] == 5 + assert sorted(table.scan().to_arrow()["id"].to_pylist()) == expected[:-1] + assert table.scan().count() == spark.sql(f"SELECT COUNT(*) AS n FROM rest.{probe_identifier}").collect()[0]["n"] == 15 + + +@pytest.mark.integration +def test_v3_deletion_vector_blob_matches_spark( + session_catalog: Catalog, + spark: "SparkSession", # type: ignore[name-defined] # noqa: F821 + probe_identifier: str, +) -> None: + spark.sql( + f""" + CREATE TABLE rest.{probe_identifier} (id BIGINT) + USING iceberg + TBLPROPERTIES ('format-version' = '3', 'write.delete.mode' = 'merge-on-read') + """ + ) + spark.range(0, 100).coalesce(1).writeTo(f"rest.{probe_identifier}").append() + # A contiguous run and scattered positions, so the Java side run-length encodes part of the bitmap + spark.sql(f"DELETE FROM rest.{probe_identifier} WHERE id BETWEEN 10 AND 60 OR id IN (3, 70, 99)") + + table = session_catalog.load_table(probe_identifier) + (delete_file,) = [task_delete for task in table.scan().plan_files() for task_delete in task.delete_files] + with table.io.new_input(delete_file.file_path).open() as f: + puffin = PuffinFile(f.read()) + (blob,) = puffin.footer.blobs + + positions = [3, 70, 99, *range(10, 61)] + expected = DeletionVector.from_positions(blob.properties["referenced-data-file"], positions).to_blob() + assert puffin.get_blob_payload(blob) == expected.payload + assert blob.type == expected.metadata.type + assert blob.fields == expected.metadata.fields + assert blob.properties == expected.metadata.properties diff --git a/tests/integration/test_writes/test_partitioned_writes.py b/tests/integration/test_writes/test_partitioned_writes.py index 1d1488255f..93497c8da6 100644 --- a/tests/integration/test_writes/test_partitioned_writes.py +++ b/tests/integration/test_writes/test_partitioned_writes.py @@ -49,7 +49,7 @@ @pytest.mark.parametrize( "part_col", ["int", "bool", "string", "string_long", "long", "float", "double", "date", "timestamp", "timestamptz", "binary"] ) -@pytest.mark.parametrize("format_version", [1, 2]) +@pytest.mark.parametrize("format_version", [1, 2, 3]) def test_query_filter_null_partitioned( session_catalog: Catalog, spark: SparkSession, arrow_table_with_null: pa.Table, part_col: str, format_version: int ) -> None: @@ -82,7 +82,7 @@ def test_query_filter_null_partitioned( @pytest.mark.parametrize( "part_col", ["int", "bool", "string", "string_long", "long", "float", "double", "date", "timestamp", "timestamptz", "binary"] ) -@pytest.mark.parametrize("format_version", [1, 2]) +@pytest.mark.parametrize("format_version", [1, 2, 3]) def test_query_filter_without_data_partitioned( session_catalog: Catalog, spark: SparkSession, @@ -119,7 +119,7 @@ def test_query_filter_without_data_partitioned( @pytest.mark.parametrize( "part_col", ["int", "bool", "string", "string_long", "long", "float", "double", "date", "timestamp", "timestamptz", "binary"] ) -@pytest.mark.parametrize("format_version", [1, 2]) +@pytest.mark.parametrize("format_version", [1, 2, 3]) def test_query_filter_only_nulls_partitioned( session_catalog: Catalog, spark: SparkSession, arrow_table_with_only_nulls: pa.Table, part_col: str, format_version: int ) -> None: @@ -151,7 +151,7 @@ def test_query_filter_only_nulls_partitioned( @pytest.mark.parametrize( "part_col", ["int", "bool", "string", "string_long", "long", "float", "double", "date", "timestamptz", "timestamp", "binary"] ) -@pytest.mark.parametrize("format_version", [1, 2]) +@pytest.mark.parametrize("format_version", [1, 2, 3]) def test_query_filter_appended_null_partitioned( session_catalog: Catalog, spark: SparkSession, arrow_table_with_null: pa.Table, part_col: str, format_version: int ) -> None: @@ -204,7 +204,7 @@ def test_query_filter_appended_null_partitioned( ) @pytest.mark.parametrize( "format_version", - [1, 2], + [1, 2, 3], ) def test_query_filter_dynamic_partition_overwrite_null_partitioned( session_catalog: Catalog, spark: SparkSession, arrow_table_with_null: pa.Table, part_col: str, format_version: int @@ -285,7 +285,7 @@ def test_query_filter_v1_v2_append_null( @pytest.mark.parametrize( "part_col", ["int", "bool", "string", "string_long", "long", "float", "double", "date", "timestamp", "timestamptz", "binary"] ) -@pytest.mark.parametrize("format_version", [1, 2]) +@pytest.mark.parametrize("format_version", [1, 2, 3]) def test_object_storage_location_provider_excludes_partition_path( session_catalog: Catalog, spark: SparkSession, arrow_table_with_null: pa.Table, part_col: str, format_version: int ) -> None: @@ -353,7 +353,7 @@ def test_object_storage_location_provider_excludes_partition_path( ) @pytest.mark.parametrize( "format_version", - [1, 2], + [1, 2, 3], ) def test_dynamic_partition_overwrite_non_identity_transform( session_catalog: Catalog, arrow_table_with_null: pa.Table, spec: PartitionSpec, format_version: int @@ -409,7 +409,7 @@ def test_dynamic_partition_overwrite_invalid_on_unpartitioned_table( ) @pytest.mark.parametrize( "format_version", - [1, 2], + [1, 2, 3], ) def test_dynamic_partition_overwrite_unpartitioned_evolve_to_identity_transform( spark: SparkSession, session_catalog: Catalog, arrow_table_with_null: pa.Table, part_col: str, format_version: int @@ -655,7 +655,7 @@ def test_data_files_with_table_partitioned_with_null( @pytest.mark.integration @pytest.mark.parametrize( "format_version", - [1, 2], + [1, 2, 3], ) def test_dynamic_partition_overwrite_rename_column(spark: SparkSession, session_catalog: Catalog, format_version: int) -> None: arrow_table = pa.Table.from_pydict( @@ -702,7 +702,7 @@ def test_dynamic_partition_overwrite_rename_column(spark: SparkSession, session_ @pytest.mark.integration @pytest.mark.parametrize( "format_version", - [1, 2], + [1, 2, 3], ) @pytest.mark.filterwarnings("ignore") def test_dynamic_partition_overwrite_evolve_partition(spark: SparkSession, session_catalog: Catalog, format_version: int) -> None: @@ -781,7 +781,7 @@ def test_invalid_arguments(spark: SparkSession, session_catalog: Catalog) -> Non (PartitionSpec(PartitionField(source_id=2, field_id=1001, transform=TruncateTransform(2), name="string_trunc"))), ], ) -@pytest.mark.parametrize("format_version", [1, 2]) +@pytest.mark.parametrize("format_version", [1, 2, 3]) def test_truncate_transform( spec: PartitionSpec, spark: SparkSession, @@ -834,7 +834,7 @@ def test_truncate_transform( ), ], ) -@pytest.mark.parametrize("format_version", [1, 2]) +@pytest.mark.parametrize("format_version", [1, 2, 3]) def test_identity_and_bucket_transform_spec( spec: PartitionSpec, spark: SparkSession, @@ -919,7 +919,7 @@ def test_unsupported_transform( (PartitionSpec(PartitionField(source_id=11, field_id=1001, transform=BucketTransform(2), name="binary_bucket")), 2), ], ) -@pytest.mark.parametrize("format_version", [1, 2]) +@pytest.mark.parametrize("format_version", [1, 2, 3]) def test_bucket_transform( spark: SparkSession, session_catalog: Catalog, @@ -970,7 +970,7 @@ def test_bucket_transform( ], ) @pytest.mark.parametrize("part_col", ["date", "timestamp", "timestamptz"]) -@pytest.mark.parametrize("format_version", [1, 2]) +@pytest.mark.parametrize("format_version", [1, 2, 3]) def test_append_ymd_transform_partitioned( session_catalog: Catalog, spark: SparkSession, @@ -1034,7 +1034,7 @@ def test_append_ymd_transform_partitioned( pytest.param(HourTransform(), {473328, 473352, 474072, 474096, 474102, None}, id="hour_transform"), ], ) -@pytest.mark.parametrize("format_version", [1, 2]) +@pytest.mark.parametrize("format_version", [1, 2, 3]) def test_append_transform_partition_verify_partitions_count( session_catalog: Catalog, spark: SparkSession, @@ -1093,7 +1093,7 @@ def test_append_transform_partition_verify_partitions_count( @pytest.mark.integration -@pytest.mark.parametrize("format_version", [1, 2]) +@pytest.mark.parametrize("format_version", [1, 2, 3]) def test_append_multiple_partitions( session_catalog: Catalog, spark: SparkSession, @@ -1157,7 +1157,7 @@ def test_append_multiple_partitions( @pytest.mark.integration -@pytest.mark.parametrize("format_version", [1, 2]) +@pytest.mark.parametrize("format_version", [1, 2, 3]) def test_stage_only_dynamic_partition_overwrite_files( spark: SparkSession, session_catalog: Catalog, arrow_table_with_null: pa.Table, format_version: int ) -> None: diff --git a/tests/integration/test_writes/test_writes.py b/tests/integration/test_writes/test_writes.py index 30fdd76ab7..ed04deab20 100644 --- a/tests/integration/test_writes/test_writes.py +++ b/tests/integration/test_writes/test_writes.py @@ -125,6 +125,48 @@ def table_v2_appended_with_null(session_catalog: Catalog, arrow_table_with_null: assert tbl.format_version == 2, f"Expected v2, got: v{tbl.format_version}" +@pytest.fixture(scope="session", autouse=True) +def table_v3_with_null(session_catalog: Catalog, arrow_table_with_null: pa.Table) -> None: + identifier = "default.arrow_table_v3_with_null" + tbl = _create_table(session_catalog, identifier, {"format-version": "3"}, [arrow_table_with_null]) + assert tbl.format_version == 3, f"Expected v3, got: v{tbl.format_version}" + + +@pytest.fixture(scope="session", autouse=True) +def table_v3_without_data(session_catalog: Catalog, arrow_table_without_data: pa.Table) -> None: + identifier = "default.arrow_table_v3_without_data" + tbl = _create_table(session_catalog, identifier, {"format-version": "3"}, [arrow_table_without_data]) + assert tbl.format_version == 3, f"Expected v3, got: v{tbl.format_version}" + + +@pytest.fixture(scope="session", autouse=True) +def table_v3_with_only_nulls(session_catalog: Catalog, arrow_table_with_only_nulls: pa.Table) -> None: + identifier = "default.arrow_table_v3_with_only_nulls" + tbl = _create_table(session_catalog, identifier, {"format-version": "3"}, [arrow_table_with_only_nulls]) + assert tbl.format_version == 3, f"Expected v3, got: v{tbl.format_version}" + + +@pytest.fixture(scope="session", autouse=True) +def table_v3_appended_with_null(session_catalog: Catalog, arrow_table_with_null: pa.Table) -> None: + identifier = "default.arrow_table_v3_appended_with_null" + tbl = _create_table(session_catalog, identifier, {"format-version": "3"}, 2 * [arrow_table_with_null]) + assert tbl.format_version == 3, f"Expected v3, got: v{tbl.format_version}" + + +@pytest.fixture(scope="session", autouse=True) +def table_v2_v3_appended_with_null(session_catalog: Catalog, arrow_table_with_null: pa.Table) -> None: + identifier = "default.arrow_table_v2_v3_appended_with_null" + tbl = _create_table(session_catalog, identifier, {"format-version": "2"}, [arrow_table_with_null]) + assert tbl.format_version == 2, f"Expected v2, got: v{tbl.format_version}" + + with tbl.transaction() as tx: + tx.upgrade_table_version(format_version=3) + + tbl.append(arrow_table_with_null) + + assert tbl.format_version == 3, f"Expected v3, got: v{tbl.format_version}" + + @pytest.fixture(scope="session", autouse=True) def table_v1_v2_appended_with_null(session_catalog: Catalog, arrow_table_with_null: pa.Table) -> None: identifier = "default.arrow_table_v1_v2_appended_with_null" @@ -140,14 +182,14 @@ def table_v1_v2_appended_with_null(session_catalog: Catalog, arrow_table_with_nu @pytest.mark.integration -@pytest.mark.parametrize("format_version", [1, 2]) +@pytest.mark.parametrize("format_version", [1, 2, 3]) def test_query_count(spark: SparkSession, format_version: int) -> None: df = spark.table(f"default.arrow_table_v{format_version}_with_null") assert df.count() == 3, "Expected 3 rows" @pytest.mark.integration -@pytest.mark.parametrize("format_version", [1, 2]) +@pytest.mark.parametrize("format_version", [1, 2, 3]) def test_query_filter_null(spark: SparkSession, arrow_table_with_null: pa.Table, format_version: int) -> None: identifier = f"default.arrow_table_v{format_version}_with_null" df = spark.table(identifier) @@ -157,7 +199,7 @@ def test_query_filter_null(spark: SparkSession, arrow_table_with_null: pa.Table, @pytest.mark.integration -@pytest.mark.parametrize("format_version", [1, 2]) +@pytest.mark.parametrize("format_version", [1, 2, 3]) def test_query_filter_without_data(spark: SparkSession, arrow_table_with_null: pa.Table, format_version: int) -> None: identifier = f"default.arrow_table_v{format_version}_without_data" df = spark.table(identifier) @@ -167,7 +209,7 @@ def test_query_filter_without_data(spark: SparkSession, arrow_table_with_null: p @pytest.mark.integration -@pytest.mark.parametrize("format_version", [1, 2]) +@pytest.mark.parametrize("format_version", [1, 2, 3]) def test_query_filter_only_nulls(spark: SparkSession, arrow_table_with_null: pa.Table, format_version: int) -> None: identifier = f"default.arrow_table_v{format_version}_with_only_nulls" df = spark.table(identifier) @@ -177,7 +219,7 @@ def test_query_filter_only_nulls(spark: SparkSession, arrow_table_with_null: pa. @pytest.mark.integration -@pytest.mark.parametrize("format_version", [1, 2]) +@pytest.mark.parametrize("format_version", [1, 2, 3]) def test_query_filter_appended_null(spark: SparkSession, arrow_table_with_null: pa.Table, format_version: int) -> None: identifier = f"default.arrow_table_v{format_version}_appended_with_null" df = spark.table(identifier) @@ -187,11 +229,14 @@ def test_query_filter_appended_null(spark: SparkSession, arrow_table_with_null: @pytest.mark.integration +@pytest.mark.parametrize( + "identifier", ["default.arrow_table_v1_v2_appended_with_null", "default.arrow_table_v2_v3_appended_with_null"] +) def test_query_filter_v1_v2_append_null( spark: SparkSession, arrow_table_with_null: pa.Table, + identifier: str, ) -> None: - identifier = "default.arrow_table_v1_v2_appended_with_null" df = spark.table(identifier) for col in arrow_table_with_null.column_names: assert df.where(f"{col} is null").count() == 2, f"Expected 1 row for {col}" @@ -395,7 +440,7 @@ def test_data_files(spark: SparkSession, session_catalog: Catalog, arrow_table_w @pytest.mark.integration -@pytest.mark.parametrize("format_version", [1, 2]) +@pytest.mark.parametrize("format_version", [1, 2, 3]) def test_object_storage_data_files( spark: SparkSession, session_catalog: Catalog, arrow_table_with_null: pa.Table, format_version: int ) -> None: @@ -444,7 +489,7 @@ def get_current_snapshot_id(identifier: str) -> int: @pytest.mark.integration -@pytest.mark.parametrize("format_version", [1, 2]) +@pytest.mark.parametrize("format_version", [1, 2, 3]) def test_python_writes_special_character_column_with_spark_reads( spark: SparkSession, session_catalog: Catalog, format_version: int ) -> None: @@ -488,7 +533,7 @@ def test_python_writes_special_character_column_with_spark_reads( @pytest.mark.integration -@pytest.mark.parametrize("format_version", [1, 2]) +@pytest.mark.parametrize("format_version", [1, 2, 3]) def test_python_writes_dictionary_encoded_column_with_spark_reads( spark: SparkSession, session_catalog: Catalog, format_version: int ) -> None: @@ -521,7 +566,7 @@ def test_python_writes_dictionary_encoded_column_with_spark_reads( @pytest.mark.integration -@pytest.mark.parametrize("format_version", [1, 2]) +@pytest.mark.parametrize("format_version", [1, 2, 3]) def test_python_writes_with_small_and_large_types_spark_reads( spark: SparkSession, session_catalog: Catalog, format_version: int ) -> None: @@ -624,7 +669,7 @@ def get_data_files_count(identifier: str) -> int: @pytest.mark.integration -@pytest.mark.parametrize("format_version", [1, 2]) +@pytest.mark.parametrize("format_version", [1, 2, 3]) @pytest.mark.parametrize( "properties, expected_compression_name", [ @@ -713,7 +758,7 @@ def test_write_parquet_unsupported_properties( @pytest.mark.integration -@pytest.mark.parametrize("format_version", [1, 2]) +@pytest.mark.parametrize("format_version", [1, 2, 3]) def test_spark_writes_orc_pyiceberg_reads(spark: SparkSession, session_catalog: Catalog, format_version: int) -> None: """Test that ORC files written by Spark can be read by PyIceberg.""" identifier = f"default.spark_writes_orc_pyiceberg_reads_v{format_version}" @@ -924,7 +969,7 @@ def test_duckdb_url_import(warehouse: Path, arrow_table_with_null: pa.Table) -> @pytest.mark.integration -@pytest.mark.parametrize("format_version", [1, 2]) +@pytest.mark.parametrize("format_version", [1, 2, 3]) def test_write_and_evolve(session_catalog: Catalog, format_version: int) -> None: identifier = f"default.arrow_write_data_and_evolve_schema_v{format_version}" @@ -967,7 +1012,7 @@ def test_write_and_evolve(session_catalog: Catalog, format_version: int) -> None @pytest.mark.integration -@pytest.mark.parametrize("format_version", [1, 2]) +@pytest.mark.parametrize("format_version", [1, 2, 3]) @pytest.mark.parametrize("catalog", [lf("session_catalog_hive"), lf("session_catalog")]) def test_create_table_transaction(catalog: Catalog, format_version: int) -> None: identifier = f"default.arrow_create_table_transaction_{catalog.name}_{format_version}" @@ -1019,7 +1064,7 @@ def test_create_table_transaction(catalog: Catalog, format_version: int) -> None @pytest.mark.integration -@pytest.mark.parametrize("format_version", [1, 2]) +@pytest.mark.parametrize("format_version", [1, 2, 3]) @pytest.mark.parametrize("catalog", [lf("session_catalog_hive"), lf("session_catalog")]) def test_create_table_with_non_default_values(catalog: Catalog, table_schema_with_all_types: Schema, format_version: int) -> None: identifier = f"default.arrow_create_table_transaction_with_non_default_values_{catalog.name}_{format_version}" @@ -1070,7 +1115,7 @@ def test_create_table_with_non_default_values(catalog: Catalog, table_schema_wit @pytest.mark.integration -@pytest.mark.parametrize("format_version", [1, 2]) +@pytest.mark.parametrize("format_version", [1, 2, 3]) def test_table_properties_int_value( session_catalog: Catalog, arrow_table_with_null: pa.Table, @@ -1087,7 +1132,7 @@ def test_table_properties_int_value( @pytest.mark.integration -@pytest.mark.parametrize("format_version", [1, 2]) +@pytest.mark.parametrize("format_version", [1, 2, 3]) def test_table_properties_raise_for_none_value( session_catalog: Catalog, arrow_table_with_null: pa.Table, @@ -1104,7 +1149,7 @@ def test_table_properties_raise_for_none_value( @pytest.mark.integration -@pytest.mark.parametrize("format_version", [1, 2]) +@pytest.mark.parametrize("format_version", [1, 2, 3]) def test_inspect_snapshots( spark: SparkSession, session_catalog: Catalog, arrow_table_with_null: pa.Table, format_version: int ) -> None: @@ -1215,7 +1260,7 @@ def get_metadata_entries_count(identifier: str) -> int: @pytest.mark.integration -@pytest.mark.parametrize("format_version", [1, 2]) +@pytest.mark.parametrize("format_version", [1, 2, 3]) def test_hive_catalog_storage_descriptor( session_catalog_hive: HiveCatalog, pa_schema: pa.Schema, @@ -1234,7 +1279,7 @@ def test_hive_catalog_storage_descriptor( @pytest.mark.integration -@pytest.mark.parametrize("format_version", [1, 2]) +@pytest.mark.parametrize("format_version", [1, 2, 3]) def test_hive_catalog_storage_descriptor_has_changed( session_catalog_hive: HiveCatalog, pa_schema: pa.Schema, @@ -1331,7 +1376,7 @@ def test_sanitize_character_partitioned_avro_bug(catalog: Catalog) -> None: @pytest.mark.integration -@pytest.mark.parametrize("format_version", [1, 2]) +@pytest.mark.parametrize("format_version", [1, 2, 3]) def test_cross_platform_special_character_compatibility( spark: SparkSession, session_catalog: Catalog, format_version: int ) -> None: @@ -1411,7 +1456,7 @@ def test_cross_platform_special_character_compatibility( @pytest.mark.integration -@pytest.mark.parametrize("format_version", [1, 2]) +@pytest.mark.parametrize("format_version", [1, 2, 3]) def test_table_write_subset_of_schema(session_catalog: Catalog, arrow_table_with_null: pa.Table, format_version: int) -> None: identifier = "default.test_table_write_subset_of_schema" tbl = _create_table(session_catalog, identifier, {"format-version": format_version}, [arrow_table_with_null]) @@ -1424,7 +1469,7 @@ def test_table_write_subset_of_schema(session_catalog: Catalog, arrow_table_with @pytest.mark.integration -@pytest.mark.parametrize("format_version", [1, 2]) +@pytest.mark.parametrize("format_version", [1, 2, 3]) @pytest.mark.filterwarnings("ignore:Delete operation did not match any records") def test_table_write_out_of_order_schema(session_catalog: Catalog, arrow_table_with_null: pa.Table, format_version: int) -> None: identifier = "default.test_table_write_out_of_order_schema" @@ -1443,7 +1488,7 @@ def test_table_write_out_of_order_schema(session_catalog: Catalog, arrow_table_w @pytest.mark.integration -@pytest.mark.parametrize("format_version", [1, 2]) +@pytest.mark.parametrize("format_version", [1, 2, 3]) def test_table_write_schema_with_valid_nullability_diff( spark: SparkSession, session_catalog: Catalog, arrow_table_with_null: pa.Table, format_version: int ) -> None: @@ -1475,7 +1520,7 @@ def test_table_write_schema_with_valid_nullability_diff( @pytest.mark.integration -@pytest.mark.parametrize("format_version", [1, 2]) +@pytest.mark.parametrize("format_version", [1, 2, 3]) def test_table_write_schema_with_valid_upcast( spark: SparkSession, session_catalog: Catalog, @@ -1525,7 +1570,7 @@ def test_table_write_schema_with_valid_upcast( @pytest.mark.integration -@pytest.mark.parametrize("format_version", [1, 2]) +@pytest.mark.parametrize("format_version", [1, 2, 3]) def test_write_all_timestamp_precision( mocker: MockerFixture, spark: SparkSession, @@ -1566,7 +1611,7 @@ def test_write_all_timestamp_precision( @pytest.mark.integration -@pytest.mark.parametrize("format_version", [1, 2]) +@pytest.mark.parametrize("format_version", [1, 2, 3]) def test_merge_manifests(session_catalog: Catalog, arrow_table_with_null: pa.Table, format_version: int) -> None: tbl_a = _create_table( session_catalog, @@ -1618,7 +1663,7 @@ def test_merge_manifests(session_catalog: Catalog, arrow_table_with_null: pa.Tab @pytest.mark.integration -@pytest.mark.parametrize("format_version", [1, 2]) +@pytest.mark.parametrize("format_version", [1, 2, 3]) def test_merge_manifests_file_content(session_catalog: Catalog, arrow_table_with_null: pa.Table, format_version: int) -> None: tbl_a = _create_table( session_catalog, @@ -1641,14 +1686,14 @@ def test_merge_manifests_file_content(session_catalog: Catalog, arrow_table_with # verify the sequence number of tbl_a's only manifest file tbl_a_manifest = tbl_a.current_snapshot().manifests(tbl_a.io)[0] # type: ignore - assert tbl_a_manifest.sequence_number == (3 if format_version == 2 else 0) - assert tbl_a_manifest.min_sequence_number == (1 if format_version == 2 else 0) + assert tbl_a_manifest.sequence_number == (3 if format_version >= 2 else 0) + assert tbl_a_manifest.min_sequence_number == (1 if format_version >= 2 else 0) # verify the manifest entries of tbl_a, in which the manifests are merged tbl_a_entries = tbl_a.inspect.entries().to_pydict() assert tbl_a_entries["status"] == [1, 0, 0] - assert tbl_a_entries["sequence_number"] == [3, 2, 1] if format_version == 2 else [0, 0, 0] - assert tbl_a_entries["file_sequence_number"] == [3, 2, 1] if format_version == 2 else [0, 0, 0] + assert tbl_a_entries["sequence_number"] == ([3, 2, 1] if format_version >= 2 else [0, 0, 0]) + assert tbl_a_entries["file_sequence_number"] == ([3, 2, 1] if format_version >= 2 else [0, 0, 0]) for i in range(3): tbl_a_data_file = tbl_a_entries["data_file"][i] assert tbl_a_data_file["column_sizes"] == [ @@ -1989,7 +2034,7 @@ def test_writing_null_structs(session_catalog: Catalog) -> None: @pytest.mark.integration -@pytest.mark.parametrize("format_version", [1, 2]) +@pytest.mark.parametrize("format_version", [1, 2, 3]) def test_abort_table_transaction_on_exception( spark: SparkSession, session_catalog: Catalog, arrow_table_with_null: pa.Table, format_version: int ) -> None: @@ -2048,7 +2093,7 @@ def test_write_optional_list(session_catalog: Catalog) -> None: @pytest.mark.integration -@pytest.mark.parametrize("format_version", [1, 2]) +@pytest.mark.parametrize("format_version", [1, 2, 3]) def test_double_commit_transaction( spark: SparkSession, session_catalog: Catalog, arrow_table_with_null: pa.Table, format_version: int ) -> None: @@ -2065,7 +2110,7 @@ def test_double_commit_transaction( @pytest.mark.integration -@pytest.mark.parametrize("format_version", [1, 2]) +@pytest.mark.parametrize("format_version", [1, 2, 3]) def test_evolve_and_write( spark: SparkSession, session_catalog: Catalog, arrow_table_with_null: pa.Table, format_version: int ) -> None: @@ -2373,10 +2418,11 @@ def test_nanosecond_support_on_catalog( _create_table(session_catalog, identifier, {"format-version": "3"}, schema=arrow_table_schema_with_all_timestamp_precisions) - with pytest.raises(NotImplementedError, match="Writing V3 is not yet supported"): - catalog.create_table( - "ns.table1", schema=arrow_table_schema_with_all_timestamp_precisions, properties={"format-version": "3"} - ) + table_v3 = catalog.create_table( + "ns.table1", schema=arrow_table_schema_with_all_timestamp_precisions, properties={"format-version": "3"} + ) + assert table_v3.format_version == 3 + assert table_v3.metadata.next_row_id == 0 with pytest.raises( UnsupportedPyArrowTypeException, match=re.escape("Column 'timestamp_ns' has an unsupported type: timestamp[ns]") @@ -2387,7 +2433,7 @@ def test_nanosecond_support_on_catalog( @pytest.mark.integration -@pytest.mark.parametrize("format_version", [1, 2]) +@pytest.mark.parametrize("format_version", [1, 2, 3]) def test_stage_only_delete( spark: SparkSession, session_catalog: Catalog, arrow_table_with_null: pa.Table, format_version: int ) -> None: @@ -2438,7 +2484,7 @@ def test_stage_only_delete( @pytest.mark.integration -@pytest.mark.parametrize("format_version", [1, 2]) +@pytest.mark.parametrize("format_version", [1, 2, 3]) def test_stage_only_append( spark: SparkSession, session_catalog: Catalog, arrow_table_with_null: pa.Table, format_version: int ) -> None: @@ -2485,7 +2531,7 @@ def test_stage_only_append( @pytest.mark.integration -@pytest.mark.parametrize("format_version", [1, 2]) +@pytest.mark.parametrize("format_version", [1, 2, 3]) def test_stage_only_overwrite_files( spark: SparkSession, session_catalog: Catalog, arrow_table_with_null: pa.Table, format_version: int ) -> None: @@ -2532,7 +2578,6 @@ def test_stage_only_overwrite_files( assert parent_snapshot_id == [None, first_snapshot, second_snapshot, second_snapshot, second_snapshot] -@pytest.mark.skip("V3 writer support is not enabled.") @pytest.mark.integration def test_v3_write_and_read_row_lineage(spark: SparkSession, session_catalog: Catalog) -> None: """Test writing to a v3 table and reading with Spark.""" @@ -2569,6 +2614,9 @@ def test_v3_write_and_read_row_lineage(spark: SparkSession, session_catalog: Cat "Expected next_row_id to be incremented by the number of added rows" ) + rows = spark.sql(f"SELECT int, _row_id FROM {identifier} ORDER BY int").collect() + assert [(row["int"], row["_row_id"]) for row in rows] == [(1, 0), (2, 1), (3, 2)] + # RecordBatchReader streaming append/overwrite — see https://github.com/apache/iceberg-python/issues/2152 # @@ -2579,7 +2627,7 @@ def test_v3_write_and_read_row_lineage(spark: SparkSession, session_catalog: Cat @pytest.mark.integration -@pytest.mark.parametrize("format_version", [1, 2]) +@pytest.mark.parametrize("format_version", [1, 2, 3]) def test_append_record_batch_reader( spark: SparkSession, session_catalog: Catalog, arrow_table_with_null: pa.Table, format_version: int ) -> None: @@ -2602,7 +2650,7 @@ def test_append_record_batch_reader( @pytest.mark.integration -@pytest.mark.parametrize("format_version", [1, 2]) +@pytest.mark.parametrize("format_version", [1, 2, 3]) def test_overwrite_record_batch_reader( spark: SparkSession, session_catalog: Catalog, arrow_table_with_null: pa.Table, format_version: int ) -> None: @@ -2623,7 +2671,7 @@ def test_overwrite_record_batch_reader( @pytest.mark.integration -@pytest.mark.parametrize("format_version", [1, 2]) +@pytest.mark.parametrize("format_version", [1, 2, 3]) def test_append_record_batch_reader_multifile( spark: SparkSession, session_catalog: Catalog, arrow_table_with_null: pa.Table, format_version: int ) -> None: diff --git a/tests/io/test_pyarrow.py b/tests/io/test_pyarrow.py index 892d8e54eb..985f010834 100644 --- a/tests/io/test_pyarrow.py +++ b/tests/io/test_pyarrow.py @@ -15,15 +15,19 @@ # specific language governing permissions and limitations # under the License. # pylint: disable=protected-access,unused-argument,redefined-outer-name +import json import logging import os +import struct import sys import tempfile import uuid import warnings +import zlib from collections.abc import Iterator from datetime import date, datetime, timezone from pathlib import Path +from types import SimpleNamespace from typing import Any from unittest.mock import MagicMock, patch from uuid import uuid4 @@ -36,6 +40,7 @@ import pytest from packaging import version from pyarrow.fs import AwsDefaultS3RetryStrategy, FileType, LocalFileSystem, S3FileSystem +from pyroaring import BitMap from pyiceberg.exceptions import ResolveError from pyiceberg.expressions import ( @@ -58,8 +63,14 @@ BoundNotStartsWith, BoundReference, BoundStartsWith, + EqualTo, GreaterThan, + GreaterThanOrEqual, + In, + IsNull, Not, + NotEqualTo, + NotNull, Or, ) from pyiceberg.expressions.literals import literal @@ -68,6 +79,7 @@ ICEBERG_SCHEMA, PYARROW_PARQUET_FIELD_ID_KEY, ArrowScan, + GeospatialStatsAggregator, PyArrowFile, PyArrowFileIO, StatsAggregator, @@ -75,6 +87,8 @@ _ConvertToArrowSchema, _determine_partitions, _primitive_to_physical, + _pyarrow_to_schema_without_ids, + _read_all_delete_files, _read_deletes, _task_to_record_batches, _to_requested_schema, @@ -84,15 +98,23 @@ data_file_statistics_from_parquet_metadata, expression_to_pyarrow, parquet_path_to_id_mapping, + pyarrow_to_schema, schema_to_pyarrow, write_file, ) from pyiceberg.manifest import DataFile, DataFileContent, FileFormat from pyiceberg.partitioning import PartitionField, PartitionSpec from pyiceberg.schema import Schema, make_compatible_name, visit -from pyiceberg.table import FileScanTask, TableProperties, WriteTask +from pyiceberg.table import FileScanTask, Table, TableProperties, WriteTask +from pyiceberg.table.deletion_vector import ( + _DV_BLOB_MAGIC_NUMBER, + PROPERTY_REFERENCED_DATA_FILE, + deletion_vectors_from_puffin_file, +) from pyiceberg.table.metadata import TableMetadataV2 +from pyiceberg.table.metadata_columns import LAST_UPDATED_SEQUENCE_NUMBER, ROW_ID from pyiceberg.table.name_mapping import create_mapping_from_schema +from pyiceberg.table.puffin import MAGIC_BYTES, PuffinFile from pyiceberg.transforms import HourTransform, IdentityTransform from pyiceberg.typedef import UTF8, Properties, Record, TableVersion from pyiceberg.types import ( @@ -115,9 +137,13 @@ StructType, TimestampNanoType, TimestampType, + TimestamptzNanoType, TimestamptzType, TimeType, + UnknownType, + VariantType, ) +from pyiceberg.utils.geo import GeospatialBound from tests.catalog.test_base import InMemoryCatalog from tests.conftest import UNIFIED_AWS_SESSION_PROPERTIES @@ -1160,6 +1186,290 @@ def _set_spec_id(datafile: DataFile) -> DataFile: ) +WKB_POINT_1_2 = bytes.fromhex("0101000000000000000000f03f0000000000000040") +WKB_POINT_3_4 = bytes.fromhex("010100000000000000000008400000000000001040") +GEO_TABLE_SCHEMA = Schema( + NestedField(1, "id", IntegerType(), required=True), + NestedField(2, "geom", GeometryType("srid:3857"), required=False), + NestedField(3, "geog", GeographyType("srid:4326", "vincenty"), required=False), +) + + +def _geo_parquet_file(tmp_path: Path, geom: pa.Array, geog: pa.Array) -> str: + """Write a Parquet file with geo columns the way another engine would, with field ids but no Iceberg schema.""" + arrow_schema = pa.schema( + [ + pa.field("id", pa.int32(), nullable=False, metadata={PYARROW_PARQUET_FIELD_ID_KEY: "1"}), + pa.field("geom", geom.type, nullable=True, metadata={PYARROW_PARQUET_FIELD_ID_KEY: "2"}), + pa.field("geog", geog.type, nullable=True, metadata={PYARROW_PARQUET_FIELD_ID_KEY: "3"}), + ] + ) + table = pa.Table.from_arrays([pa.array([1, 2, 3], pa.int32()), geom, geog], schema=arrow_schema) + return _write_table_to_file(str(tmp_path / "geo.parquet"), arrow_schema, table) + + +def _assert_geo_result(result: pa.Table, expected_types: tuple[pa.DataType, pa.DataType]) -> None: + assert result.column("id").to_pylist() == [1, 2, 3] + assert result.schema.field("geom").type == expected_types[0] + assert result.schema.field("geog").type == expected_types[1] + for name in ("geom", "geog"): + column = result.column(name) + storage = column.cast(column.type.storage_type) if isinstance(column.type, pa.ExtensionType) else column + assert storage.to_pylist() == [WKB_POINT_1_2, None, WKB_POINT_3_4] + + +@pytest.mark.parametrize("binary_type", [pa.binary(), pa.large_binary()]) +def test_projection_geo_from_plain_binary(tmp_path: Path, binary_type: pa.DataType) -> None: + values = pa.array([WKB_POINT_1_2, None, WKB_POINT_3_4], binary_type) + file = _geo_parquet_file(tmp_path, values, values) + result = project(GEO_TABLE_SCHEMA, [file]) + _assert_geo_result( + result, + (schema_to_pyarrow(GeometryType("srid:3857")), schema_to_pyarrow(GeographyType("srid:4326", "vincenty"))), + ) + + +def test_projection_geo_from_plain_binary_without_geoarrow(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> None: + values = pa.array([WKB_POINT_1_2, None, WKB_POINT_3_4], pa.binary()) + file = _geo_parquet_file(tmp_path, values, values) + # Importing a module set to None raises ImportError, as if geoarrow-pyarrow was not installed + monkeypatch.setitem(sys.modules, "geoarrow.pyarrow", None) + result = project(GEO_TABLE_SCHEMA, [file]) + _assert_geo_result(result, (pa.large_binary(), pa.large_binary())) + + +def test_projection_geo_from_geoarrow(tmp_path: Path) -> None: + ga = pytest.importorskip("geoarrow.pyarrow") + wkb = ga.as_wkb(pa.array([WKB_POINT_1_2, None, WKB_POINT_3_4], pa.binary())) + geom_type = ga.wkb().with_crs("srid:3857") + # The PyArrow Parquet writer only supports planar and spherical edges + geog_type = ga.wkb().with_crs("srid:4326").with_edge_type(ga.EdgeType.SPHERICAL) + geom = pa.ExtensionArray.from_storage(geom_type, wkb.storage) + geog = pa.ExtensionArray.from_storage(geog_type, wkb.storage) + file = _geo_parquet_file(tmp_path, geom, geog) + table_schema = Schema( + NestedField(1, "id", IntegerType(), required=True), + NestedField(2, "geom", GeometryType("srid:3857"), required=False), + NestedField(3, "geog", GeographyType("srid:4326"), required=False), + ) + result = project(table_schema, [file]) + _assert_geo_result(result, (geom_type, geog_type)) + + +@pytest.mark.parametrize( + "expr, expected_ids", + [ + (EqualTo("geom", WKB_POINT_3_4), [3]), + (NotEqualTo("geog", WKB_POINT_3_4), [1]), + (In("geom", {WKB_POINT_1_2, WKB_POINT_3_4}), [1, 3]), + (IsNull("geom"), [2]), + (NotNull("geog"), [1, 3]), + ], +) +def test_projection_geo_equality_filter(tmp_path: Path, expr: BooleanExpression, expected_ids: list[int]) -> None: + values = pa.array([WKB_POINT_1_2, None, WKB_POINT_3_4], pa.binary()) + file = _geo_parquet_file(tmp_path, values, values) + assert project(GEO_TABLE_SCHEMA, [file], expr=expr).column("id").to_pylist() == expected_ids + + +def test_projection_geo_filter_on_other_column(tmp_path: Path) -> None: + values = pa.array([WKB_POINT_1_2, None, WKB_POINT_3_4], pa.binary()) + file = _geo_parquet_file(tmp_path, values, values) + result = project(GEO_TABLE_SCHEMA, [file], expr=GreaterThan("id", 1)) + assert result.column("id").to_pylist() == [2, 3] + geom = result.column("geom") + storage = geom.cast(geom.type.storage_type) if isinstance(geom.type, pa.ExtensionType) else geom + assert storage.to_pylist() == [None, WKB_POINT_3_4] + + +def _v3_table(tmp_path: Path, schema: Schema, name: str = "tbl") -> Table: + from pyiceberg.catalog.memory import InMemoryCatalog as MemoryCatalog + + memory_catalog = MemoryCatalog("memory", warehouse=f"file://{tmp_path}") + memory_catalog.create_namespace("default") + return memory_catalog.create_table(f"default.{name}", schema, properties={"format-version": "3"}) + + +def _wkb_storage(column: pa.ChunkedArray) -> list[Any]: + storage = column.cast(column.type.storage_type) if isinstance(column.type, pa.ExtensionType) else column + return storage.to_pylist() + + +WKB_POINT_NEG = struct.pack(" None: + pytest.importorskip("geoarrow.pyarrow") + tbl = _v3_table(tmp_path, GEO_WRITE_SCHEMA) + values = pa.array([WKB_POINT_1_2, None, WKB_POINT_NEG], binary_type) + tbl.append(pa.table({"id": pa.array([1, 2, 3], pa.int32()), "geom": values, "geom_default": values, "geog": values})) + + result = tbl.scan().to_arrow() + for name in ("geom", "geom_default", "geog"): + assert _wkb_storage(result.column(name)) == [WKB_POINT_1_2, None, WKB_POINT_NEG] + + (data_file,) = [task.file for task in tbl.scan().plan_files()] + parquet_schema = str(pq.ParquetFile(data_file.file_path.removeprefix("file://")).schema) + assert "field_id=2 geom (Geometry(crs=srid:3857))" in parquet_schema + # The default CRS OGC:CRS84 is written as an empty CRS + assert "field_id=3 geom_default (Geometry(crs=))" in parquet_schema + assert "field_id=4 geog (Geography(crs=srid:4326, algorithm=spherical))" in parquet_schema + + # Geometry bounds are the corners of the bounding box + for field_id in (2, 3): + assert GeospatialBound.from_bytes(data_file.lower_bounds[field_id]) == GeospatialBound(-3.0, 2.0) + assert GeospatialBound.from_bytes(data_file.upper_bounds[field_id]) == GeospatialBound(1.0, 5.0) + assert data_file.value_counts[2] == 3 + + readable_metrics = tbl.inspect.files().to_pylist()[0]["readable_metrics"] + assert GeospatialBound.from_bytes(readable_metrics["geom"]["lower_bound"]) == GeospatialBound(-3.0, 2.0) + assert readable_metrics["geom"]["value_count"] == 3 + + +def test_write_geo_from_geoarrow(tmp_path: Path) -> None: + ga = pytest.importorskip("geoarrow.pyarrow") + tbl = _v3_table(tmp_path, GEO_WRITE_SCHEMA) + wkb = pa.array([WKB_POINT_1_2, WKB_POINT_NEG], pa.binary()) + geom = pa.ExtensionArray.from_storage(ga.wkb().with_crs("srid:3857"), wkb) + geom_default = pa.ExtensionArray.from_storage(ga.wkb(), wkb) + geog = pa.ExtensionArray.from_storage(ga.wkb().with_crs("srid:4326").with_edge_type(ga.EdgeType.SPHERICAL), wkb) + ids = pa.array([1, 2], pa.int32()) + + # A GeoArrow CRS that differs from the table CRS is a type mismatch + with pytest.raises(ValueError, match="Mismatch in fields"): + tbl.append(pa.table({"id": ids, "geom": geom_default, "geom_default": geom_default, "geog": geog})) + + tbl.append(pa.table({"id": ids, "geom": geom, "geom_default": geom_default, "geog": geog})) + + result = tbl.scan().to_arrow() + assert result.schema.field("geom").type == schema_to_pyarrow(GeometryType("srid:3857")) + for name in ("geom", "geom_default", "geog"): + assert _wkb_storage(result.column(name)) == [WKB_POINT_1_2, WKB_POINT_NEG] + + +def test_write_geography_unsupported_edges_as_plain_binary(tmp_path: Path) -> None: + pytest.importorskip("geoarrow.pyarrow") + schema = Schema(NestedField(1, "geog", GeographyType("srid:4326", "vincenty"), required=False)) + tbl = _v3_table(tmp_path, schema) + with pytest.warns(UserWarning, match="vincenty edge algorithm as plain WKB binary"): + tbl.append(pa.table({"geog": pa.array([WKB_POINT_1_2], pa.binary())})) + + assert _wkb_storage(tbl.scan().to_arrow().column("geog")) == [WKB_POINT_1_2] + (data_file,) = [task.file for task in tbl.scan().plan_files()] + parquet_schema = str(pq.ParquetFile(data_file.file_path.removeprefix("file://")).schema) + assert "field_id=1 geog;" in parquet_schema + assert not data_file.lower_bounds + + +def test_write_geo_without_geoarrow(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> None: + tbl = _v3_table(tmp_path, GEO_WRITE_SCHEMA) + # Importing a module set to None raises ImportError, as if geoarrow-pyarrow was not installed + monkeypatch.setitem(sys.modules, "geoarrow.pyarrow", None) + values = pa.array([WKB_POINT_1_2, None], pa.binary()) + with pytest.warns(UserWarning, match="install pyiceberg\\[geoarrow\\]"): + tbl.append(pa.table({"id": pa.array([1, 2], pa.int32()), "geom": values, "geom_default": values, "geog": values})) + + result = tbl.scan().to_arrow() + assert result.schema.field("geom").type == pa.large_binary() + assert result.column("geom").to_pylist() == [WKB_POINT_1_2, None] + (data_file,) = [task.file for task in tbl.scan().plan_files()] + # Lexicographic min/max of WKB are not geo bounds, so they are not recorded + assert 2 not in data_file.lower_bounds and 2 not in data_file.upper_bounds + assert data_file.null_value_counts[2] == 1 + + +def test_add_files_geo_bounds_from_parquet_statistics(tmp_path: Path) -> None: + ga = pytest.importorskip("geoarrow.pyarrow") + tbl = _v3_table(tmp_path, Schema(NestedField(1, "geom", GeometryType(), required=False))) + points = [struct.pack(" SimpleNamespace: + """Stand-in for the pyarrow GeoStatistics of a row group.""" + return SimpleNamespace(**{key: box.get(key) for key in ("xmin", "ymin", "zmin", "mmin", "xmax", "ymax", "zmax", "mmax")}) + + +def test_geography_bounds_wrapping_antimeridian_are_omitted() -> None: + aggregator = GeospatialStatsAggregator(GeographyType()) + aggregator.update_box(_geo_statistics(xmin=170.0, ymin=-10.0, xmax=-170.0, ymax=10.0)) + assert aggregator.min_as_bytes() is None and aggregator.max_as_bytes() is None + + aggregator = GeospatialStatsAggregator(GeographyType()) + aggregator.update_box(_geo_statistics(xmin=-170.0, ymin=-10.0, xmax=170.0, ymax=10.0)) + assert aggregator.min_as_bytes() == GeospatialBound(-170.0, -10.0).to_bytes() + + +def test_geo_bounds_merge_row_groups() -> None: + aggregator = GeospatialStatsAggregator(GeometryType()) + aggregator.update_box(_geo_statistics(xmin=0.0, ymin=0.0, xmax=1.0, ymax=1.0, zmin=5.0, zmax=6.0)) + aggregator.update_box(_geo_statistics(xmin=-1.0, ymin=0.5, xmax=0.5, ymax=2.0)) + # Z is only kept when every row group has it + assert aggregator.min_as_bytes() == GeospatialBound(-1.0, 0.0).to_bytes() + assert aggregator.max_as_bytes() == GeospatialBound(1.0, 2.0).to_bytes() + + # A row group without a bounding box (e.g. only nulls) invalidates the bounds + aggregator.update_box(None) + assert aggregator.min_as_bytes() is None + + +def test_write_default_fills_missing_column(tmp_path: Path) -> None: + tbl = _v3_table(tmp_path, Schema(NestedField(1, "id", IntegerType(), required=False))) + tbl.append(pa.table({"id": pa.array([1], pa.int32())})) + with tbl.update_schema() as update: + update.add_column("color", StringType(), required=True, default_value="blue") + with tbl.update_schema() as update: + update.set_default_value("color", "green") + + # The required column is missing from the dataframe, which is allowed because it has a write-default + tbl.append(pa.table({"id": pa.array([2], pa.int32())})) + + result = tbl.scan().to_arrow().sort_by("id") + # Rows written before the column existed read the initial-default, new rows were written with the write-default + assert result.column("color").to_pylist() == ["blue", "green"] + + +@pytest.mark.parametrize("read_type, unit", [(TimestampType(), "us"), (TimestampNanoType(), "ns")]) +def test_projection_date_promoted_to_timestamp(tmp_path: Path, read_type: PrimitiveType, unit: str) -> None: + arrow_schema = pa.schema([pa.field("d", pa.date32(), metadata={PYARROW_PARQUET_FIELD_ID_KEY: "1"})]) + file = _write_table_to_file( + str(tmp_path / "date.parquet"), arrow_schema, pa.table({"d": [date(2024, 1, 2), None]}, schema=arrow_schema) + ) + result = project(Schema(NestedField(1, "d", read_type, required=False)), [file]) + assert result.schema.field("d").type == pa.timestamp(unit) + assert result.column("d").to_pylist() == [datetime(2024, 1, 2), None] + + +def test_unknown_columns_are_not_written(tmp_path: Path) -> None: + schema = Schema(NestedField(1, "id", IntegerType(), required=False), NestedField(2, "u", UnknownType(), required=False)) + tbl = _v3_table(tmp_path, schema) + tbl.append(pa.table({"id": pa.array([1], pa.int32())})) + tbl.append(pa.table({"id": pa.array([2], pa.int32()), "u": pa.nulls(1)})) + + result = tbl.scan().to_arrow().sort_by("id") + assert result.column("u").to_pylist() == [None, None] + for task in tbl.scan().plan_files(): + assert pq.ParquetFile(task.file.file_path.removeprefix("file://")).schema_arrow.names == ["id"] + + def test_projection_add_column(file_int: str) -> None: schema = Schema( # All new IDs @@ -1822,6 +2132,193 @@ def test_read_deletes(deletes_file: str, request: pytest.FixtureRequest) -> None assert list(deletes.values())[0] == pa.chunked_array([[1, 3, 5]]) +def _deletion_vector_bitmap_payload(values: list[int] | None = None) -> bytes: + return (1).to_bytes(8, byteorder="little") + (0).to_bytes(4, byteorder="little") + BitMap(values or [1, 3, 5]).serialize() + + +def _deletion_vector_blob(bitmap_payload: bytes) -> bytes: + bitmap_data = struct.pack("I", len(bitmap_data)) + bitmap_data + struct.pack(">I", zlib.crc32(bitmap_data) & 0xFFFFFFFF) + + +def _deletion_vector_puffin_payload(referenced_data_file: str) -> bytes: + dv_blob = _deletion_vector_blob(_deletion_vector_bitmap_payload()) + footer_payload = json.dumps( + { + "blobs": [ + { + "type": "deletion-vector-v1", + "fields": [2147483546], + "snapshot-id": 1, + "sequence-number": 1, + "offset": len(MAGIC_BYTES) + 4, + "length": len(dv_blob), + "properties": {PROPERTY_REFERENCED_DATA_FILE: referenced_data_file}, + } + ], + "properties": {}, + } + ).encode() + + return ( + MAGIC_BYTES + + b"\x00\x00\x00\x00" + + dv_blob + + footer_payload + + len(footer_payload).to_bytes(4, byteorder="little") + + b"\x00\x00\x00\x00" + + MAGIC_BYTES + ) + + +def test_deletion_vectors_from_puffin_file(tmp_path: Path) -> None: + referenced_data_file = f"{tmp_path}/data.parquet" + puffin_payload = _deletion_vector_puffin_payload(referenced_data_file) + delete_file_path = f"{tmp_path}/deletes.puffin" + + with open(delete_file_path, "wb") as f: + f.write(puffin_payload) + + with open(delete_file_path, "rb") as f: + deletion_vectors = deletion_vectors_from_puffin_file(PuffinFile(f.read())) + + assert {dv.referenced_data_file: dv.to_vector() for dv in deletion_vectors} == { + referenced_data_file: pa.chunked_array([[1, 3, 5]]) + } + + +def test_read_deletes_from_whole_puffin_file(tmp_path: Path) -> None: + referenced_data_file = f"{tmp_path}/data.parquet" + delete_file_path = f"{tmp_path}/deletes.puffin" + + with open(delete_file_path, "wb") as f: + f.write(_deletion_vector_puffin_payload(referenced_data_file)) + + deletes = _read_deletes( + PyArrowFileIO(), + DataFile.from_args( + _table_format_version=3, + content=DataFileContent.POSITION_DELETES, + file_path=delete_file_path, + file_format=FileFormat.PUFFIN, + record_count=3, + ), + ) + + assert deletes == {referenced_data_file: pa.chunked_array([[1, 3, 5]])} + + +def test_read_deletion_vector_blob_from_content_range(tmp_path: Path) -> None: + referenced_data_file = f"{tmp_path}/data.parquet" + dv_blob = _deletion_vector_blob(_deletion_vector_bitmap_payload()) + prefix = b"\x01not-a-puffin-file" + delete_file_path = f"{tmp_path}/deletes.bin" + + with open(delete_file_path, "wb") as f: + f.write(prefix + dv_blob + b"trailing-bytes") + + deletes = _read_deletes( + PyArrowFileIO(), + DataFile.from_args( + _table_format_version=3, + content=DataFileContent.POSITION_DELETES, + file_path=delete_file_path, + file_format=FileFormat.PUFFIN, + record_count=3, + referenced_data_file=referenced_data_file, + content_offset=len(prefix), + content_size_in_bytes=len(dv_blob), + ), + ) + + assert deletes == {referenced_data_file: pa.chunked_array([[1, 3, 5]])} + + +def test_read_all_delete_files_keeps_multiple_dv_content_ranges_for_same_path(tmp_path: Path) -> None: + referenced_data_file = f"{tmp_path}/data.parquet" + first_dv_blob = _deletion_vector_blob(_deletion_vector_bitmap_payload([1, 3])) + second_dv_blob = _deletion_vector_blob(_deletion_vector_bitmap_payload([5])) + delete_file_path = f"{tmp_path}/deletes.bin" + + with open(delete_file_path, "wb") as f: + f.write(first_dv_blob + second_dv_blob) + + first_dv = DataFile.from_args( + _table_format_version=3, + content=DataFileContent.POSITION_DELETES, + file_path=delete_file_path, + file_format=FileFormat.PUFFIN, + record_count=2, + referenced_data_file=referenced_data_file, + content_offset=0, + content_size_in_bytes=len(first_dv_blob), + ) + second_dv = DataFile.from_args( + _table_format_version=3, + content=DataFileContent.POSITION_DELETES, + file_path=delete_file_path, + file_format=FileFormat.PUFFIN, + record_count=1, + referenced_data_file=referenced_data_file, + content_offset=len(first_dv_blob), + content_size_in_bytes=len(second_dv_blob), + ) + data_file = DataFile.from_args( + content=DataFileContent.DATA, + file_path=referenced_data_file, + file_format=FileFormat.PARQUET, + record_count=10, + file_size_in_bytes=100, + ) + + deletes = _read_all_delete_files( + PyArrowFileIO(), + [FileScanTask(data_file=data_file, delete_files=[first_dv, second_dv])], + ) + + assert sorted(delete.to_pylist() for delete in deletes[referenced_data_file]) == [[1, 3], [5]] + + +def test_read_all_delete_files_applies_delete_files_only_to_planned_data_files(tmp_path: Path) -> None: + data_a, data_b = f"{tmp_path}/a.parquet", f"{tmp_path}/b.parquet" + position_deletes_path = f"{tmp_path}/pos-deletes.parquet" + pq.write_table(pa.table({"file_path": [data_a, data_b], "pos": [0, 1]}), position_deletes_path) + position_deletes = DataFile.from_args( + content=DataFileContent.POSITION_DELETES, file_path=position_deletes_path, file_format=FileFormat.PARQUET + ) + + dv_blob = _deletion_vector_blob(_deletion_vector_bitmap_payload([2])) + dv_path = f"{tmp_path}/deletes.puffin" + with open(dv_path, "wb") as f: + f.write(dv_blob) + dv = DataFile.from_args( + _table_format_version=3, + content=DataFileContent.POSITION_DELETES, + file_path=dv_path, + file_format=FileFormat.PUFFIN, + record_count=1, + referenced_data_file=data_a, + content_offset=0, + content_size_in_bytes=len(dv_blob), + ) + + def data_file(path: str) -> DataFile: + return DataFile.from_args( + content=DataFileContent.DATA, file_path=path, file_format=FileFormat.PARQUET, record_count=10, file_size_in_bytes=1 + ) + + # The DV supersedes the position delete file for a.parquet, which still applies to b.parquet + deletes = _read_all_delete_files( + PyArrowFileIO(), + [ + FileScanTask(data_file=data_file(data_a), delete_files=[dv]), + FileScanTask(data_file=data_file(data_b), delete_files=[position_deletes]), + ], + ) + + assert {path: [arr.to_pylist() for arr in arrays] for path, arrays in deletes.items()} == {data_a: [[2]], data_b: [[1]]} + + def test_delete(deletes_file: str, request: pytest.FixtureRequest, table_schema_simple: Schema) -> None: # Determine file format from the file extension file_format = FileFormat.PARQUET if deletes_file.endswith(".parquet") else FileFormat.ORC @@ -1867,6 +2364,121 @@ def test_delete(deletes_file: str, request: pytest.FixtureRequest, table_schema_ assert str(with_deletes) == expected_str +_ROW_LINEAGE_TABLE_SCHEMA = Schema(NestedField(1, "id", LongType(), required=False), schema_id=0) + + +def _row_lineage_scan( + tmp_path: Path, + first_row_id: int | None, + sequence_number: int | None, + deleted_positions: list[int] | None = None, + stored_lineage: dict[str, list[int | None]] | None = None, + row_filter: BooleanExpression | None = None, +) -> tuple[ArrowScan, FileScanTask]: + """Write a data file with ids 0..4 and return a scan projecting ``id`` and the row lineage columns.""" + columns: dict[str, pa.Array] = { + "id": pa.array(range(5), type=pa.int64()), + } + fields = [pa.field("id", pa.int64(), metadata={PYARROW_PARQUET_FIELD_ID_KEY: "1"})] + for column, field in ((ROW_ID.name, ROW_ID), (LAST_UPDATED_SEQUENCE_NUMBER.name, LAST_UPDATED_SEQUENCE_NUMBER)): + if stored_lineage and column in stored_lineage: + columns[column] = pa.array(stored_lineage[column], type=pa.int64()) + fields.append(pa.field(column, pa.int64(), metadata={PYARROW_PARQUET_FIELD_ID_KEY: str(field.field_id)})) + data_file_path = str(tmp_path / "data.parquet") + pq.write_table(pa.Table.from_pydict(columns, schema=pa.schema(fields)), data_file_path) + + data_file = DataFile.from_args( + _table_format_version=3, + content=DataFileContent.DATA, + file_path=data_file_path, + file_format=FileFormat.PARQUET, + partition=Record(), + record_count=5, + file_size_in_bytes=os.path.getsize(data_file_path), + first_row_id=first_row_id, + ) + data_file.spec_id = 0 + + delete_files = [] + if deleted_positions: + deletes_path = str(tmp_path / "deletes.parquet") + pq.write_table(pa.table({"file_path": [data_file_path] * len(deleted_positions), "pos": deleted_positions}), deletes_path) + delete_files.append( + DataFile.from_args(content=DataFileContent.POSITION_DELETES, file_path=deletes_path, file_format=FileFormat.PARQUET) + ) + + task = FileScanTask(data_file=data_file, delete_files=delete_files, sequence_number=sequence_number) + scan = ArrowScan( + table_metadata=TableMetadataV2( + location="file://a/b/c.json", + last_column_id=1, + format_version=2, + current_schema_id=0, + schemas=[_ROW_LINEAGE_TABLE_SCHEMA], + partition_specs=[PartitionSpec()], + ), + io=PyArrowFileIO(), + projected_schema=Schema(*_ROW_LINEAGE_TABLE_SCHEMA.fields, ROW_ID, LAST_UPDATED_SEQUENCE_NUMBER), + row_filter=row_filter if row_filter is not None else AlwaysTrue(), + ) + return scan, task + + +def test_row_lineage_inherited_with_positional_deletes(tmp_path: Path) -> None: + scan, task = _row_lineage_scan(tmp_path, first_row_id=10, sequence_number=5, deleted_positions=[1]) + + result = scan.to_table([task]) + + assert result.schema.names == ["id", "_row_id", "_last_updated_sequence_number"] + assert result.schema.field("_row_id").type == pa.int64() + assert result.to_pydict() == { + "id": [0, 2, 3, 4], + "_row_id": [10, 12, 13, 14], + "_last_updated_sequence_number": [5, 5, 5, 5], + } + batches = list(scan.to_record_batches([task])) + assert pa.Table.from_batches(batches).to_pydict() == result.to_pydict() + + +def test_row_lineage_positions_survive_row_filter(tmp_path: Path) -> None: + scan, task = _row_lineage_scan(tmp_path, first_row_id=10, sequence_number=5, row_filter=GreaterThanOrEqual("id", 3)) + assert scan.to_table([task]).to_pydict() == {"id": [3, 4], "_row_id": [13, 14], "_last_updated_sequence_number": [5, 5]} + + scan, task = _row_lineage_scan( + tmp_path, first_row_id=10, sequence_number=5, deleted_positions=[3], row_filter=GreaterThanOrEqual("id", 2) + ) + assert scan.to_table([task]).to_pydict() == {"id": [2, 4], "_row_id": [12, 14], "_last_updated_sequence_number": [5, 5]} + + +def test_row_lineage_null_without_first_row_id(tmp_path: Path) -> None: + scan, task = _row_lineage_scan(tmp_path, first_row_id=None, sequence_number=5, deleted_positions=[1]) + + assert scan.to_table([task]).to_pydict() == { + "id": [0, 2, 3, 4], + "_row_id": [None, None, None, None], + "_last_updated_sequence_number": [None, None, None, None], + } + + +def test_row_lineage_keeps_values_stored_in_data_file(tmp_path: Path) -> None: + scan, task = _row_lineage_scan( + tmp_path, + first_row_id=10, + sequence_number=5, + deleted_positions=[1], + stored_lineage={ + "_row_id": [100, None, 102, None, 104], + "_last_updated_sequence_number": [2, None, 3, None, None], + }, + ) + + assert scan.to_table([task]).to_pydict() == { + "id": [0, 2, 3, 4], + "_row_id": [100, 102, 13, 104], + "_last_updated_sequence_number": [2, 3, 5, 5], + } + + def test_delete_duplicates(deletes_file: str, request: pytest.FixtureRequest, table_schema_simple: Schema) -> None: # Determine file format from the file extension file_format = FileFormat.PARQUET if deletes_file.endswith(".parquet") else FileFormat.ORC @@ -3392,6 +4004,39 @@ def _expected_batch(unit: str) -> pa.RecordBatch: assert _expected_batch("ns" if format_version > 2 else "us").equals(actual_result) +@pytest.mark.parametrize( + "iceberg_type, tz", + [(TimestampNanoType(), None), (TimestamptzNanoType(), "UTC")], +) +def test_task_to_record_batches_micros_file_for_nanos_table(iceberg_type: PrimitiveType, tz: str | None, tmpdir: str) -> None: + # A file written with downcast-ns-timestamp-to-us-on-write stores microseconds for a nanosecond column + arrow_type = pa.timestamp("us", tz=tz) + arrow_table = pa.table( + [pa.array([1755172800000000, 1755176400000001], type=arrow_type)], + pa.schema((pa.field("ts_field", arrow_type, nullable=True, metadata={PYARROW_PARQUET_FIELD_ID_KEY: "1"}),)), + ) + data_file = _write_table_to_data_file(f"{tmpdir}/test_micros_file_for_nanos_table.parquet", arrow_table.schema, arrow_table) + table_schema = Schema(NestedField(1, "ts_field", iceberg_type, required=False)) + zone = "+00:00" if tz else "" + + actual_result = list( + _task_to_record_batches( + PyArrowFileIO(), + FileScanTask(data_file), + bound_row_filter=GreaterThanOrEqual("ts_field", f"2025-08-14T13:00:00{zone}").bind(table_schema), + projected_schema=table_schema, + table_schema=table_schema, + projected_field_ids={1}, + positional_deletes=None, + case_sensitive=True, + format_version=3, + ) + )[0] + + expected = pa.record_batch([pa.array([1755176400000001000], type=pa.timestamp("ns", tz=tz))], names=["ts_field"]) + assert expected.equals(actual_result) + + def test_task_to_record_batches_scanner_filter_not_set_with_positional_deletes(tmpdir: str) -> None: """Regression test for https://github.com/apache/iceberg-python/issues/3272. @@ -3516,6 +4161,391 @@ def test_task_to_record_batches_filter_after_positional_deletes_empty_result(tmp assert result_batches == [] +def _write_equality_delete_file(filepath: str, table: pa.Table, equality_ids: list[int]) -> DataFile: + _write_table_to_file(filepath, table.schema, table) + return DataFile.from_args( + content=DataFileContent.EQUALITY_DELETES, + file_path=filepath, + file_format=FileFormat.PARQUET, + partition={}, + record_count=len(table), + file_size_in_bytes=22, + equality_ids=equality_ids, + ) + + +def _id_name_file(tmpdir: str) -> tuple[DataFile, Schema]: + arrow_schema = pa.schema( + ( + pa.field("id", pa.int32(), nullable=True, metadata={PYARROW_PARQUET_FIELD_ID_KEY: "1"}), + pa.field("name", pa.string(), nullable=True, metadata={PYARROW_PARQUET_FIELD_ID_KEY: "2"}), + ) + ) + arrow_table = pa.table( + [pa.array([1, 2, 3, 4, None, 6], type=pa.int32()), pa.array(["a", "b", None, "d", "e", None])], schema=arrow_schema + ) + data_file = _write_table_to_data_file(f"{tmpdir}/eq_data.parquet", arrow_schema, arrow_table) + table_schema = Schema( + NestedField(1, "id", IntegerType(), required=False), + NestedField(2, "name", StringType(), required=False), + ) + return data_file, table_schema + + +def test_read_equality_deletes_projects_delete_columns(tmpdir: str) -> None: + from pyiceberg.io.pyarrow import _read_equality_deletes + + delete_schema = pa.schema( + ( + pa.field("id", pa.int64(), nullable=True, metadata={PYARROW_PARQUET_FIELD_ID_KEY: "1"}), + pa.field("extra", pa.string(), nullable=True, metadata={PYARROW_PARQUET_FIELD_ID_KEY: "3"}), + ) + ) + delete_file = _write_equality_delete_file( + f"{tmpdir}/eq_deletes.parquet", pa.table([pa.array([2, None]), pa.array(["x", "y"])], schema=delete_schema), [1] + ) + + deletes = _read_equality_deletes(PyArrowFileIO(), delete_file) + + assert deletes.field_ids == (1,) + # columns that are not delete columns are ignored + assert deletes.rows.column_names == ["1"] + assert deletes.rows.column("1").to_pylist() == [2, None] + + +def test_read_equality_deletes_missing_delete_column(tmpdir: str) -> None: + from pyiceberg.io.pyarrow import _read_equality_deletes + + delete_schema = pa.schema((pa.field("id", pa.int32(), nullable=True, metadata={PYARROW_PARQUET_FIELD_ID_KEY: "1"}),)) + delete_file = _write_equality_delete_file( + f"{tmpdir}/eq_deletes.parquet", pa.table([pa.array([2], type=pa.int32())], schema=delete_schema), [1, 2] + ) + + with pytest.raises(ValueError, match="missing the delete column with id 2"): + _read_equality_deletes(PyArrowFileIO(), delete_file) + + +def test_task_to_record_batches_equality_deletes_single_column(tmpdir: str) -> None: + from pyiceberg.io.pyarrow import _read_equality_deletes + + data_file, table_schema = _id_name_file(tmpdir) + # a long delete column matches the int data column; the null delete row deletes the row with a null id + delete_schema = pa.schema((pa.field("id", pa.int64(), nullable=True, metadata={PYARROW_PARQUET_FIELD_ID_KEY: "1"}),)) + delete_file = _write_equality_delete_file( + f"{tmpdir}/eq_deletes.parquet", pa.table([pa.array([2, 4, None, 100])], schema=delete_schema), [1] + ) + deletes = _read_equality_deletes(PyArrowFileIO(), delete_file) + + result = pa.Table.from_batches( + _task_to_record_batches( + PyArrowFileIO(), + FileScanTask(data_file, delete_files=[delete_file]), + bound_row_filter=AlwaysTrue(), + projected_schema=table_schema.select("name"), + table_schema=table_schema, + projected_field_ids={2}, + positional_deletes=None, + case_sensitive=True, + equality_deletes=[deletes], + ) + ) + + # the delete column is not projected, but it is still used to match rows + assert result.column_names == ["name"] + assert result.column("name").to_pylist() == ["a", None, None] + + +def test_task_to_record_batches_equality_deletes_multi_column_null_equal(tmpdir: str) -> None: + from pyiceberg.io.pyarrow import _read_equality_deletes + + data_file, table_schema = _id_name_file(tmpdir) + delete_schema = pa.schema( + ( + pa.field("id", pa.int32(), nullable=True, metadata={PYARROW_PARQUET_FIELD_ID_KEY: "1"}), + pa.field("name", pa.string(), nullable=True, metadata={PYARROW_PARQUET_FIELD_ID_KEY: "2"}), + ) + ) + delete_rows = pa.table( + [pa.array([1, 3, 4, 6], type=pa.int32()), pa.array(["a", None, "x", None])], + schema=delete_schema, + ) + delete_file = _write_equality_delete_file(f"{tmpdir}/eq_deletes.parquet", delete_rows, [1, 2]) + deletes = _read_equality_deletes(PyArrowFileIO(), delete_file) + + result = pa.Table.from_batches( + _task_to_record_batches( + PyArrowFileIO(), + FileScanTask(data_file, delete_files=[delete_file]), + bound_row_filter=AlwaysTrue(), + projected_schema=table_schema, + table_schema=table_schema, + projected_field_ids={1, 2}, + positional_deletes=None, + case_sensitive=True, + equality_deletes=[deletes], + ) + ) + + # (1, a) matches, (3, null) and (6, null) match with null == null, (4, d) does not match (4, x) + assert result.to_pylist() == [{"id": 2, "name": "b"}, {"id": 4, "name": "d"}, {"id": None, "name": "e"}] + + +def test_task_to_record_batches_equality_deletes_after_positional_deletes_with_filter(tmpdir: str) -> None: + from pyiceberg.expressions.visitors import bind + from pyiceberg.io.pyarrow import _read_equality_deletes + + data_file, table_schema = _id_name_file(tmpdir) + delete_schema = pa.schema((pa.field("name", pa.string(), nullable=True, metadata={PYARROW_PARQUET_FIELD_ID_KEY: "2"}),)) + delete_file = _write_equality_delete_file( + f"{tmpdir}/eq_deletes.parquet", pa.table([pa.array(["d"])], schema=delete_schema), [2] + ) + deletes = _read_equality_deletes(PyArrowFileIO(), delete_file) + + result = pa.Table.from_batches( + _task_to_record_batches( + PyArrowFileIO(), + FileScanTask(data_file, delete_files=[delete_file]), + bound_row_filter=bind(table_schema, GreaterThan("id", 1), case_sensitive=True), + projected_schema=table_schema, + table_schema=table_schema, + projected_field_ids={1, 2}, + # position 1 holds id 2 + positional_deletes=[pa.chunked_array([pa.array([1], type=pa.int64())])], + case_sensitive=True, + equality_deletes=[deletes], + ) + ) + + assert result.column("id").to_pylist() == [3, 6] + + +def test_task_to_record_batches_equality_deletes_column_missing_from_data_file(tmpdir: str) -> None: + from pyiceberg.io.pyarrow import _read_equality_deletes + + data_file, _ = _id_name_file(tmpdir) + # column 3 was added to the table after the data file was written, so its values read as null + table_schema = Schema( + NestedField(1, "id", IntegerType(), required=False), + NestedField(2, "name", StringType(), required=False), + NestedField(3, "added", StringType(), required=False), + ) + delete_schema = pa.schema( + ( + pa.field("id", pa.int32(), nullable=True, metadata={PYARROW_PARQUET_FIELD_ID_KEY: "1"}), + pa.field("added", pa.string(), nullable=True, metadata={PYARROW_PARQUET_FIELD_ID_KEY: "3"}), + ) + ) + delete_rows = pa.table([pa.array([1, 2], type=pa.int32()), pa.array([None, "x"])], schema=delete_schema) + delete_file = _write_equality_delete_file(f"{tmpdir}/eq_deletes.parquet", delete_rows, [1, 3]) + deletes = _read_equality_deletes(PyArrowFileIO(), delete_file) + + result = pa.Table.from_batches( + _task_to_record_batches( + PyArrowFileIO(), + FileScanTask(data_file, delete_files=[delete_file]), + bound_row_filter=AlwaysTrue(), + projected_schema=table_schema, + table_schema=table_schema, + projected_field_ids={1, 2, 3}, + positional_deletes=None, + case_sensitive=True, + equality_deletes=[deletes], + ) + ) + + assert result.column("id").to_pylist() == [2, 3, 4, None, 6] + + +def test_arrow_scan_applies_equality_deletes(tmpdir: str) -> None: + from pyiceberg.io.pyarrow import ArrowScan + + data_file, table_schema = _id_name_file(tmpdir) + data_file.spec_id = 0 + delete_schema = pa.schema((pa.field("id", pa.int32(), nullable=True, metadata={PYARROW_PARQUET_FIELD_ID_KEY: "1"}),)) + delete_file = _write_equality_delete_file( + f"{tmpdir}/eq_deletes.parquet", pa.table([pa.array([1, 6], type=pa.int32())], schema=delete_schema), [1] + ) + metadata = TableMetadataV2( + location="file://a/b/", + last_column_id=2, + format_version=2, + schemas=[table_schema], + partition_specs=[PartitionSpec()], + ) + + result = ArrowScan(metadata, PyArrowFileIO(), table_schema, AlwaysTrue()).to_table( + [FileScanTask(data_file, delete_files=[delete_file])] + ) + + assert result.column("id").to_pylist() == [2, 3, 4, None] + + +def _variant_group(*extra: pa.Field) -> pa.StructType: + return pa.struct([pa.field("metadata", pa.binary(), nullable=False), pa.field("value", pa.binary(), nullable=False), *extra]) + + +def _variant_data_file(tmpdir: str, variant_type: pa.StructType, values: list[Any]) -> DataFile: + arrow_schema = pa.schema( + ( + pa.field("id", pa.int32(), nullable=False, metadata={PYARROW_PARQUET_FIELD_ID_KEY: "1"}), + pa.field("v", variant_type, nullable=True, metadata={PYARROW_PARQUET_FIELD_ID_KEY: "2"}), + ) + ) + arrow_table = pa.table( + [pa.array(range(len(values)), type=pa.int32()), pa.array(values, type=variant_type)], schema=arrow_schema + ) + return _write_table_to_data_file(f"{tmpdir}/variant.parquet", arrow_schema, arrow_table) + + +VARIANT_TABLE_SCHEMA = Schema( + NestedField(1, "id", IntegerType(), required=True), NestedField(2, "v", VariantType(), required=False) +) + + +def test_variant_schema_to_pyarrow() -> None: + assert schema_to_pyarrow(VariantType()) == pa.struct( + [pa.field("metadata", pa.binary(), nullable=False), pa.field("value", pa.binary(), nullable=False)] + ) + + +def test_pyarrow_variant_group_to_iceberg() -> None: + arrow_schema = pa.schema([pa.field("v", _variant_group(), metadata={PYARROW_PARQUET_FIELD_ID_KEY: "2"})]) + assert pyarrow_to_schema(arrow_schema) == Schema(NestedField(2, "v", VariantType(), required=False)) + + # an Iceberg struct always has field ids on its fields, so it is never mistaken for a variant + struct_with_ids = pa.struct( + [ + pa.field("metadata", pa.binary(), metadata={PYARROW_PARQUET_FIELD_ID_KEY: "3"}), + pa.field("value", pa.binary(), metadata={PYARROW_PARQUET_FIELD_ID_KEY: "4"}), + ] + ) + converted = pyarrow_to_schema(pa.schema([pa.field("s", struct_with_ids, metadata={PYARROW_PARQUET_FIELD_ID_KEY: "2"})])) + assert isinstance(converted.find_type(2), StructType) + + # without field ids at all, a struct with metadata and value fields stays a struct + without_ids = _pyarrow_to_schema_without_ids(pa.schema([pa.field("s", _variant_group())])) + assert isinstance(without_ids.find_type("s"), StructType) + + +def test_task_to_record_batches_reads_variant(tmpdir: str) -> None: + from pyiceberg.variant import to_json + + values = [ + {"metadata": b"\x01\x01\x00\x01a", "value": b"\x02\x01\x00\x00\x02\x0c\x01"}, + {"metadata": b"\x01\x00\x00", "value": b"\x35just a string"}, + None, + ] + data_file = _variant_data_file(tmpdir, _variant_group(), values) + + result = pa.Table.from_batches( + _task_to_record_batches( + PyArrowFileIO(), + FileScanTask(data_file), + bound_row_filter=AlwaysTrue(), + projected_schema=VARIANT_TABLE_SCHEMA, + table_schema=VARIANT_TABLE_SCHEMA, + projected_field_ids={1, 2}, + positional_deletes=None, + case_sensitive=True, + format_version=3, + ) + ) + + assert result.schema.field("v").type == schema_to_pyarrow(VariantType()) + assert result.column("v").to_pylist() == values + assert [to_json(v["metadata"], v["value"]) for v in values if v is not None] == ['{"a":1}', '"just a string"'] + + +def test_task_to_record_batches_shredded_variant_not_supported(tmpdir: str) -> None: + variant_type = pa.struct( + [ + pa.field("metadata", pa.binary(), nullable=False), + pa.field("value", pa.binary(), nullable=True), + pa.field("typed_value", pa.int64(), nullable=True), + ] + ) + data_file = _variant_data_file(tmpdir, variant_type, [{"metadata": b"\x01\x00\x00", "value": None, "typed_value": 1}]) + + batches = _task_to_record_batches( + PyArrowFileIO(), + FileScanTask(data_file), + bound_row_filter=AlwaysTrue(), + projected_schema=VARIANT_TABLE_SCHEMA, + table_schema=VARIANT_TABLE_SCHEMA, + projected_field_ids={1, 2}, + positional_deletes=None, + case_sensitive=True, + format_version=3, + ) + with pytest.raises(NotImplementedError, match="shredded variant"): + list(batches) + + # a scan that does not project the variant column still works + id_schema = VARIANT_TABLE_SCHEMA.select("id") + result = pa.Table.from_batches( + _task_to_record_batches( + PyArrowFileIO(), + FileScanTask(data_file), + bound_row_filter=AlwaysTrue(), + projected_schema=id_schema, + table_schema=VARIANT_TABLE_SCHEMA, + projected_field_ids={1}, + positional_deletes=None, + case_sensitive=True, + format_version=3, + ) + ) + assert result.column("id").to_pylist() == [0] + + +def test_variant_read_from_struct_with_field_ids(tmpdir: str) -> None: + # a writer that assigned field ids to the variant fields surfaces a struct, which the table schema resolves + variant_type = pa.struct( + [ + pa.field("metadata", pa.binary(), nullable=False, metadata={PYARROW_PARQUET_FIELD_ID_KEY: "3"}), + pa.field("value", pa.binary(), nullable=False, metadata={PYARROW_PARQUET_FIELD_ID_KEY: "4"}), + ] + ) + values = [{"metadata": b"\x01\x00\x00", "value": b"\x00"}] + data_file = _variant_data_file(tmpdir, variant_type, values) + + result = pa.Table.from_batches( + _task_to_record_batches( + PyArrowFileIO(), + FileScanTask(data_file), + bound_row_filter=AlwaysTrue(), + projected_schema=VARIANT_TABLE_SCHEMA, + table_schema=VARIANT_TABLE_SCHEMA, + projected_field_ids={1, 2}, + positional_deletes=None, + case_sensitive=True, + format_version=3, + ) + ) + assert result.column("v").to_pylist() == values + + +def test_write_variant_not_supported() -> None: + with pytest.raises(NotImplementedError, match="Writing variant is not supported"): + _check_pyarrow_schema_compatible(VARIANT_TABLE_SCHEMA, schema_to_pyarrow(VARIANT_TABLE_SCHEMA), format_version=3) + + +def test_variant_statistics_are_counts_only(tmpdir: str) -> None: + values = [{"metadata": b"\x01\x00\x00", "value": b"\x00"}, None, {"metadata": b"\x01\x00\x00", "value": b"\x0c\x01"}] + data_file = _variant_data_file(tmpdir, _variant_group(), values) + parquet_metadata = pq.read_metadata(data_file.file_path) + + statistics = data_file_statistics_from_parquet_metadata( + parquet_metadata=parquet_metadata, + stats_columns=compute_statistics_plan(VARIANT_TABLE_SCHEMA, {}), + parquet_column_mapping=parquet_path_to_id_mapping(VARIANT_TABLE_SCHEMA), + ) + + assert statistics.value_counts[2] == 3 + assert statistics.null_value_counts[2] == 1 + assert 2 not in statistics.column_aggregates + + def test_parse_location_defaults() -> None: """Test that parse_location uses defaults.""" diff --git a/tests/io/test_pyarrow_stats.py b/tests/io/test_pyarrow_stats.py index 0e628829eb..a42d97e6d0 100644 --- a/tests/io/test_pyarrow_stats.py +++ b/tests/io/test_pyarrow_stats.py @@ -43,6 +43,7 @@ STRUCT_INT32, STRUCT_INT64, ) +from pyiceberg.conversions import from_bytes from pyiceberg.io.pyarrow import ( MetricModeTypes, MetricsMode, @@ -65,7 +66,10 @@ BooleanType, FloatType, IntegerType, + NestedField, StringType, + TimestampNanoType, + TimestamptzNanoType, ) from pyiceberg.utils.datetime import date_to_days, datetime_to_micros, time_to_micros @@ -748,6 +752,43 @@ def test_read_missing_statistics() -> None: assert datafile.null_value_counts[string_col_idx] == 1 +def test_nano_timestamp_bounds_keep_nanosecond_precision() -> None: + schema = Schema( + NestedField(1, "ts_ns", TimestampNanoType(), required=False), + NestedField(2, "tstz_ns", TimestamptzNanoType(), required=False), + NestedField(3, "ts_ns_as_us", TimestampNanoType(), required=False), + ) + arrow_schema = schema_to_pyarrow(schema) + # A file written with downcast-ns-timestamp-to-us-on-write stores microseconds for a timestamp_ns column + arrow_schema = arrow_schema.set(2, arrow_schema.field(2).with_type(pa.timestamp("us"))) + table = pa.Table.from_pydict( + { + "ts_ns": [1704067200123456789, None, 1706832000000000001], + "tstz_ns": [1704067200123456789, None, 1706832000000000001], + "ts_ns_as_us": [1704067200123456, None, 1706832000000001], + }, + schema=arrow_schema, + ) + + metadata_collector: list[Any] = [] + with pa.BufferOutputStream() as f: + with pq.ParquetWriter(f, table.schema, metadata_collector=metadata_collector) as writer: + writer.write_table(table) + + statistics = data_file_statistics_from_parquet_metadata( + parquet_metadata=metadata_collector[0], + stats_columns=compute_statistics_plan(schema, {}), + parquet_column_mapping=parquet_path_to_id_mapping(schema), + ) + datafile = DataFile.from_args(**statistics.to_serialized_dict()) + + for field_id in (1, 2): + assert from_bytes(TimestampNanoType(), datafile.lower_bounds[field_id]) == 1704067200123456789 + assert from_bytes(TimestampNanoType(), datafile.upper_bounds[field_id]) == 1706832000000000001 + assert from_bytes(TimestampNanoType(), datafile.lower_bounds[3]) == 1704067200123456000 + assert from_bytes(TimestampNanoType(), datafile.upper_bounds[3]) == 1706832000000001000 + + # This is commented out for now because write_to_dataset drops the partition # columns making it harder to calculate the mapping from the column index to # datatype id diff --git a/tests/io/test_pyarrow_visitor.py b/tests/io/test_pyarrow_visitor.py index e98d76e262..3919ffbdba 100644 --- a/tests/io/test_pyarrow_visitor.py +++ b/tests/io/test_pyarrow_visitor.py @@ -37,6 +37,7 @@ _ConvertToIceberg, _ConvertToIcebergWithoutIDs, _expression_to_complementary_pyarrow, + _geoarrow_crs_to_iceberg, _HasIds, _NullNaNUnmentionedTermsCollector, _pyarrow_schema_ensure_large_types, @@ -55,6 +56,8 @@ DoubleType, FixedType, FloatType, + GeographyType, + GeometryType, IcebergType, IntegerType, ListType, @@ -78,6 +81,65 @@ def test_pyarrow_binary_to_iceberg() -> None: assert visit(converted_iceberg_type, _ConvertToArrowSchema()) == pyarrow_type +class _GeoArrowWkbLike(pa.ExtensionType): + """Unregistered stand-in for the geoarrow.wkb extension type, usable without geoarrow-pyarrow.""" + + def __init__(self, metadata: bytes) -> None: + self._metadata = metadata + super().__init__(pa.binary(), "geoarrow.wkb") + + def __arrow_ext_serialize__(self) -> bytes: + return self._metadata + + @classmethod + def __arrow_ext_deserialize__(cls, storage_type: pa.DataType, serialized: bytes) -> "_GeoArrowWkbLike": + return cls(serialized) + + +@pytest.mark.parametrize( + "metadata, expected", + [ + (b"", GeometryType()), + (b"{}", GeometryType()), + (b'{"crs": "srid:3857"}', GeometryType("srid:3857")), + (b'{"crs": "OGC:CRS84", "edges": "planar"}', GeometryType()), + (b'{"crs": {"id": {"authority": "EPSG", "code": 3857}}}', GeometryType("EPSG:3857")), + (b'{"edges": "spherical"}', GeographyType()), + (b'{"crs": "srid:4326", "edges": "vincenty"}', GeographyType("srid:4326", "vincenty")), + ], +) +def test_pyarrow_geoarrow_wkb_to_iceberg(metadata: bytes, expected: IcebergType) -> None: + assert visit_pyarrow(_GeoArrowWkbLike(metadata), _ConvertToIceberg()) == expected + + +def test_pyarrow_geoarrow_wkb_unsupported_edges() -> None: + with pytest.raises(TypeError, match="Unsupported GeoArrow edge type: bogus"): + visit_pyarrow(_GeoArrowWkbLike(b'{"edges": "bogus"}'), _ConvertToIceberg()) + + +def test_geoarrow_projjson_crs_without_id() -> None: + with pytest.raises(TypeError, match="Unsupported GeoArrow CRS"): + _geoarrow_crs_to_iceberg({"type": "GeographicCRS"}) + + +@pytest.mark.parametrize( + "iceberg_type", + [ + GeometryType(), + GeometryType("srid:3857"), + GeographyType(), + GeographyType("srid:4326"), + GeographyType("srid:4326", "vincenty"), + GeographyType("OGC:CRS84", "karney"), + ], +) +def test_geoarrow_round_trip(iceberg_type: IcebergType) -> None: + pytest.importorskip("geoarrow.pyarrow") + arrow_type = schema_to_pyarrow(iceberg_type) + assert arrow_type.extension_name == "geoarrow.wkb" + assert visit_pyarrow(arrow_type, _ConvertToIceberg()) == iceberg_type + + def test_pyarrow_decimal128_to_iceberg() -> None: precision = 26 scale = 20 diff --git a/tests/table/test_commit_retry.py b/tests/table/test_commit_retry.py index ab95457e23..65b79fe84b 100644 --- a/tests/table/test_commit_retry.py +++ b/tests/table/test_commit_retry.py @@ -452,6 +452,69 @@ def test_snapshot_isolation_allows_concurrent_append_delete(catalog: Catalog) -> assert len(result) == 5 +def test_snapshot_isolation_merge_on_read_delete_retries_after_concurrent_append(catalog: Catalog) -> None: + """A deletion vector commit is rebuilt on top of a concurrent append under snapshot isolation.""" + catalog.create_namespace("default") + catalog.create_table( + "default.mor_retry_test", + schema=_test_schema(), + properties={ + "format-version": "3", + "write.delete.mode": "merge-on-read", + "write.delete.isolation-level": "snapshot", + }, + ) + + import pyarrow as pa + + df = pa.table({"x": [1, 2, 3]}) + + tbl = catalog.load_table("default.mor_retry_test") + tbl.append(df) + tbl.delete("x == 3") + + tbl1 = catalog.load_table("default.mor_retry_test") + tbl2 = catalog.load_table("default.mor_retry_test") + + tbl1.append(df) + tbl2.delete("x == 1") + + refreshed = catalog.load_table("default.mor_retry_test") + # The delete did not see the concurrently appended rows, so their x == 1 row remains + assert sorted(refreshed.scan().to_arrow()["x"].to_pylist()) == [1, 2, 2, 3] + snapshot = refreshed.current_snapshot() + assert snapshot is not None and snapshot.summary is not None + assert snapshot.summary["added-dvs"] == "1" + assert snapshot.summary["removed-dvs"] == "1" + assert snapshot.summary["total-data-files"] == "2" + assert snapshot.summary["total-delete-files"] == "1" + assert snapshot.summary["total-position-deletes"] == "2" + + +def test_concurrent_merge_on_read_deletes_on_same_file_conflict(catalog: Catalog) -> None: + """Two deletion vectors for the same data file cannot be committed concurrently.""" + catalog.create_namespace("default") + catalog.create_table( + "default.mor_conflict_test", + schema=_test_schema(), + properties={"format-version": "3", "write.delete.mode": "merge-on-read"}, + ) + + import pyarrow as pa + + tbl = catalog.load_table("default.mor_conflict_test") + tbl.append(pa.table({"x": [1, 2, 3]})) + + tbl1 = catalog.load_table("default.mor_conflict_test") + tbl2 = catalog.load_table("default.mor_conflict_test") + + tbl1.delete("x == 1") + with pytest.raises(ValidationException): + tbl2.delete("x == 2") + + assert sorted(catalog.load_table("default.mor_conflict_test").scan().to_arrow()["x"].to_pylist()) == [2, 3] + + def test_uncommitted_manifests_tracked_correctly(catalog: Catalog) -> None: """Verify that uncommitted manifests are moved to _uncommitted_manifests on retry.""" catalog.create_namespace("default") diff --git a/tests/table/test_delete_file_index.py b/tests/table/test_delete_file_index.py index 09dd9ac81b..9b93417545 100644 --- a/tests/table/test_delete_file_index.py +++ b/tests/table/test_delete_file_index.py @@ -17,6 +17,7 @@ import pytest from pyiceberg.manifest import DataFile, DataFileContent, FileFormat, ManifestEntry, ManifestEntryStatus +from pyiceberg.table.delete_file import DeleteFileSet from pyiceberg.table.delete_file_index import PATH_FIELD_ID, DeleteFileIndex, PositionDeletes from pyiceberg.typedef import Record @@ -65,17 +66,26 @@ def _create_partition_delete(sequence_number: int = 1, spec_id: int = 0, partiti def _create_deletion_vector( - sequence_number: int = 1, file_path: str = "s3://bucket/data.parquet", spec_id: int = 0 + sequence_number: int = 1, + file_path: str = "s3://bucket/data.parquet", + spec_id: int = 0, + delete_file_path: str | None = None, + content_offset: int | None = None, + content_size_in_bytes: int | None = None, ) -> ManifestEntry: delete_file = DataFile.from_args( + _table_format_version=3, content=DataFileContent.POSITION_DELETES, - file_path=f"s3://bucket/deletion-vector-{sequence_number}.puffin", + file_path=delete_file_path or f"s3://bucket/deletion-vector-{sequence_number}.puffin", file_format=FileFormat.PUFFIN, partition=Record(), record_count=10, file_size_in_bytes=100, lower_bounds={PATH_FIELD_ID: file_path.encode()}, upper_bounds={PATH_FIELD_ID: file_path.encode()}, + referenced_data_file=file_path, + content_offset=content_offset, + content_size_in_bytes=content_size_in_bytes, ) delete_file._spec_id = spec_id return ManifestEntry.from_args(status=ManifestEntryStatus.ADDED, sequence_number=sequence_number, data_file=delete_file) @@ -148,17 +158,106 @@ def test_mix_path_and_partition_deletes() -> None: assert len(result) == 2 -def test_dvs_treated_as_position_deletes() -> None: +def test_dv_supersedes_older_position_deletes() -> None: index = DeleteFileIndex() + partition = Record() index.add_delete_file(_create_positional_delete(sequence_number=2, file_path="s3://bucket/a.parquet")) + index.add_delete_file(_create_partition_delete(sequence_number=2), partition) index.add_delete_file(_create_deletion_vector(sequence_number=3, file_path="s3://bucket/a.parquet")) - data_file = _create_data_file(file_path="s3://bucket/a.parquet") + result = index.for_data_file(1, _create_data_file(file_path="s3://bucket/a.parquet"), partition) + assert [(d.file_format, d.referenced_data_file) for d in result] == [(FileFormat.PUFFIN, "s3://bucket/a.parquet")] - result = index.for_data_file(1, data_file) - assert len(result) == 2 - assert all(d.content == DataFileContent.POSITION_DELETES for d in result) + # Position deletes for other data files are unaffected by the DV + other = index.for_data_file(1, _create_data_file(file_path="s3://bucket/b.parquet"), partition) + assert [d.file_path for d in other] == ["s3://bucket/pos-delete-2.parquet"] + + +def test_position_deletes_newer_than_dv_are_kept() -> None: + index = DeleteFileIndex() + + index.add_delete_file(_create_deletion_vector(sequence_number=2, file_path="s3://bucket/a.parquet")) + index.add_delete_file(_create_positional_delete(sequence_number=3, file_path="s3://bucket/a.parquet")) + + result = index.for_data_file(1, _create_data_file(file_path="s3://bucket/a.parquet")) + assert {d.file_format for d in result} == {FileFormat.PUFFIN, FileFormat.PARQUET} + + +def test_only_latest_dv_applies_to_data_file() -> None: + index = DeleteFileIndex() + + index.add_delete_file(_create_deletion_vector(sequence_number=2, file_path="s3://bucket/a.parquet")) + index.add_delete_file(_create_deletion_vector(sequence_number=4, file_path="s3://bucket/a.parquet")) + + result = index.for_data_file(1, _create_data_file(file_path="s3://bucket/a.parquet")) + assert [d.file_path for d in result] == ["s3://bucket/deletion-vector-4.puffin"] + + # A DV older than the data file does not apply + assert len(index.for_data_file(5, _create_data_file(file_path="s3://bucket/a.parquet"))) == 0 + + +def test_dv_applies_only_to_referenced_data_file() -> None: + index = DeleteFileIndex() + dv = _create_deletion_vector(sequence_number=2, file_path="s3://bucket/a.parquet") + # Without path bounds, only referenced_data_file identifies the target data file + dv.data_file._data[10] = None + dv.data_file._data[11] = None + index.add_delete_file(dv) + + assert len(index.for_data_file(1, _create_data_file(file_path="s3://bucket/a.parquet"))) == 1 + assert len(index.for_data_file(1, _create_data_file(file_path="s3://bucket/b.parquet"))) == 0 + assert index.referenced_delete_files() == [dv.data_file] + + +def test_delete_file_set_uses_content_range_identity() -> None: + shared_file_path = "s3://bucket/deletion-vectors.bin" + first_dv = _create_deletion_vector( + sequence_number=2, + delete_file_path=shared_file_path, + content_offset=4, + content_size_in_bytes=10, + ).data_file + second_dv = _create_deletion_vector( + sequence_number=3, + delete_file_path=shared_file_path, + content_offset=40, + content_size_in_bytes=12, + ).data_file + + assert first_dv == second_dv + assert len(DeleteFileSet([first_dv, second_dv])) == 2 + + +def test_dvs_in_one_puffin_file_apply_to_their_own_data_files() -> None: + index = DeleteFileIndex() + shared_file_path = "s3://bucket/deletion-vectors.puffin" + + index.add_delete_file( + _create_deletion_vector( + sequence_number=2, + file_path="s3://bucket/a.parquet", + delete_file_path=shared_file_path, + content_offset=4, + content_size_in_bytes=10, + ) + ) + index.add_delete_file( + _create_deletion_vector( + sequence_number=2, + file_path="s3://bucket/b.parquet", + delete_file_path=shared_file_path, + content_offset=14, + content_size_in_bytes=12, + ) + ) + + result_a = index.for_data_file(1, _create_data_file(file_path="s3://bucket/a.parquet")) + result_b = index.for_data_file(1, _create_data_file(file_path="s3://bucket/b.parquet")) + + assert [(dv.file_path, dv.content_offset, dv.content_size_in_bytes) for dv in result_a] == [(shared_file_path, 4, 10)] + assert [(dv.file_path, dv.content_offset, dv.content_size_in_bytes) for dv in result_b] == [(shared_file_path, 14, 12)] + assert len(DeleteFileSet([*result_a, *result_b])) == 2 def test_cannot_add_after_indexing() -> None: @@ -187,3 +286,63 @@ def test_record_equality_for_partition_lookup() -> None: assert len(index.for_data_file(1, data_file, partition_b)) == 1 assert len(index.for_data_file(1, data_file, partition_c)) == 0 + + +def _create_equality_delete(sequence_number: int = 1, spec_id: int = 0, partition: Record | None = None) -> ManifestEntry: + delete_file = DataFile.from_args( + content=DataFileContent.EQUALITY_DELETES, + file_path=f"s3://bucket/eq-delete-{spec_id}-{sequence_number}.parquet", + file_format=FileFormat.PARQUET, + partition=partition or Record(), + record_count=10, + file_size_in_bytes=100, + equality_ids=[1], + ) + delete_file._spec_id = spec_id + return ManifestEntry.from_args(status=ManifestEntryStatus.ADDED, sequence_number=sequence_number, data_file=delete_file) + + +def test_equality_deletes_apply_to_strictly_older_data() -> None: + index = DeleteFileIndex() + index.add_delete_file(_create_equality_delete(sequence_number=2)) + index.add_delete_file(_create_equality_delete(sequence_number=4)) + data_file = _create_data_file() + + assert not index.is_empty() + assert {f.file_path for f in index.for_data_file(1, data_file)} == { + "s3://bucket/eq-delete-0-2.parquet", + "s3://bucket/eq-delete-0-4.parquet", + } + # an equality delete with the same data sequence number as the data file does not apply + assert {f.file_path for f in index.for_data_file(2, data_file)} == {"s3://bucket/eq-delete-0-4.parquet"} + assert len(index.for_data_file(4, data_file)) == 0 + assert len(index.referenced_delete_files()) == 2 + + +def test_partitioned_equality_deletes_apply_to_same_partition() -> None: + partition = Record(1) + index = DeleteFileIndex() + index.add_delete_file(_create_equality_delete(sequence_number=2, spec_id=1, partition=partition), partition_key=partition) + data_file = _create_data_file(spec_id=1) + + assert len(index.for_data_file(1, data_file, partition_key=Record(1))) == 1 + assert len(index.for_data_file(1, data_file, partition_key=Record(2))) == 0 + assert len(index.for_data_file(1, _create_data_file(spec_id=0), partition_key=Record(1))) == 0 + + +def test_unpartitioned_equality_deletes_are_global() -> None: + index = DeleteFileIndex() + # stored with an unpartitioned spec, so it applies to data files in every spec and partition + index.add_delete_file(_create_equality_delete(sequence_number=2, spec_id=0)) + data_file = _create_data_file(spec_id=1) + + assert len(index.for_data_file(1, data_file, partition_key=Record(7))) == 1 + + +def test_equality_deletes_are_kept_alongside_dvs() -> None: + index = DeleteFileIndex() + index.add_delete_file(_create_deletion_vector(sequence_number=3)) + index.add_delete_file(_create_equality_delete(sequence_number=2)) + + delete_files = index.for_data_file(1, _create_data_file()) + assert {f.content for f in delete_files} == {DataFileContent.POSITION_DELETES, DataFileContent.EQUALITY_DELETES} diff --git a/tests/table/test_deletion_vector.py b/tests/table/test_deletion_vector.py index 788216f8b3..388b6cee7c 100644 --- a/tests/table/test_deletion_vector.py +++ b/tests/table/test_deletion_vector.py @@ -14,12 +14,31 @@ # KIND, either express or implied. See the License for the # specific language governing permissions and limitations # under the License. +import struct +import zlib from os import path +from pathlib import Path import pytest from pyroaring import BitMap -from pyiceberg.table.deletion_vector import DeletionVector +from pyiceberg.io.pyarrow import PyArrowFileIO +from pyiceberg.manifest import DataFile, DataFileContent, FileFormat +from pyiceberg.table.deletion_vector import ( + _DV_BLOB_MAGIC_NUMBER, + DELETION_VECTOR_MAGIC, + DELETION_VECTOR_V1_BLOB_TYPE, + MAX_POSITION, + PROPERTY_REFERENCED_DATA_FILE, + ROW_POSITION_FIELD_ID, + DeletionVector, + _deserialize_dv_blob, + has_deletion_vector_content_reference, + read_deletion_vectors, + write_deletion_vectors, +) +from pyiceberg.table.puffin import PuffinFile +from pyiceberg.typedef import Record def _open_file(file: str) -> bytes: @@ -28,6 +47,97 @@ def _open_file(file: str) -> bytes: return f.read() +def _dv_blob(bitmap_payload: bytes) -> bytes: + bitmap_data = struct.pack("I", len(bitmap_data)) + bitmap_data + struct.pack(">I", zlib.crc32(bitmap_data) & 0xFFFFFFFF) + + +def _bitmap_payload() -> bytes: + return (1).to_bytes(8, byteorder="little") + (0).to_bytes(4, byteorder="little") + BitMap([1, 3, 5]).serialize() + + +def test_deserialize_deletion_vector_blob() -> None: + actual = _deserialize_dv_blob(_dv_blob(_bitmap_payload()), record_count=3) + + assert actual == [BitMap([1, 3, 5])] + + +def test_deserialize_deletion_vector_blob_invalid_length() -> None: + with pytest.raises(ValueError, match="Invalid bitmap data length"): + _deserialize_dv_blob(_dv_blob(_bitmap_payload())[:-1]) + + +def test_deserialize_deletion_vector_blob_invalid_magic() -> None: + bitmap_data = struct.pack("I", len(bitmap_data)) + bitmap_data + struct.pack(">I", zlib.crc32(bitmap_data) & 0xFFFFFFFF) + + with pytest.raises(ValueError, match="Invalid magic number"): + _deserialize_dv_blob(blob) + + +def test_deserialize_deletion_vector_blob_invalid_crc() -> None: + blob = bytearray(_dv_blob(_bitmap_payload())) + blob[-1] ^= 1 + + with pytest.raises(ValueError, match="Invalid CRC"): + _deserialize_dv_blob(bytes(blob)) + + +def test_deserialize_deletion_vector_blob_invalid_cardinality() -> None: + with pytest.raises(ValueError, match="Invalid cardinality"): + _deserialize_dv_blob(_dv_blob(_bitmap_payload()), record_count=4) + + +def _data_file( + content_offset: int | None = None, + content_size_in_bytes: int | None = None, + referenced_data_file: str | None = None, +) -> DataFile: + return DataFile.from_args( + _table_format_version=3, + content=DataFileContent.POSITION_DELETES, + file_path="deletes.puffin", + file_format=FileFormat.PUFFIN, + record_count=1, + content_offset=content_offset, + content_size_in_bytes=content_size_in_bytes, + referenced_data_file=referenced_data_file, + ) + + +def test_has_deletion_vector_content_reference() -> None: + assert not has_deletion_vector_content_reference(_data_file()) + assert has_deletion_vector_content_reference(_data_file(content_offset=0)) + assert has_deletion_vector_content_reference(_data_file(content_size_in_bytes=1)) + assert has_deletion_vector_content_reference(_data_file(referenced_data_file="data.parquet")) + + +def test_read_deletion_vector_rejects_truncated_range(tmp_path: Path) -> None: + delete_file_path = str(tmp_path / "deletes.puffin") + blob = _dv_blob(_bitmap_payload()) + with open(delete_file_path, "wb") as f: + f.write(blob) + + dv = DataFile.from_args( + _table_format_version=3, + content=DataFileContent.POSITION_DELETES, + file_path=delete_file_path, + file_format=FileFormat.PUFFIN, + record_count=3, + referenced_data_file="data.parquet", + content_offset=0, + content_size_in_bytes=len(blob) + 8, + ) + + with pytest.raises(ValueError, match="Could not read deletion vector"): + read_deletion_vectors(PyArrowFileIO(), dv) + + +def test_read_deletion_vector_requires_complete_content_reference() -> None: + with pytest.raises(ValueError, match="content size is missing"): + read_deletion_vectors(PyArrowFileIO(), _data_file(content_offset=0, referenced_data_file="data.parquet")) + + def test_map_empty() -> None: puffin = _open_file("64mapempty.bin") @@ -66,8 +176,174 @@ def test_map_spread_vals() -> None: assert expected == actual +def test_map_declared_count_exceeds_payload() -> None: + # A truncated payload that claims a large number of bitmaps must be rejected, + # rather than driving the deserialization loop on data that is not there. + puffin = (2**32).to_bytes(8, byteorder="little") + b"\x00\x00\x00\x00" + + with pytest.raises(ValueError, match="Payload declares 4294967296 bitmaps, but only holds 4 bytes"): + _ = DeletionVector._deserialize_bitmap(puffin) + + def test_map_high_vals() -> None: puffin = _open_file("64maphighvals.bin") with pytest.raises(ValueError, match="Key 4022190063 is too large, max 2147483647 to maintain compatibility with Java impl"): _ = DeletionVector._deserialize_bitmap(puffin) + + +@pytest.mark.parametrize( + "positions", + [ + [1, 2, 3], + [0], + [3, 1, 2, 1, 3], # unordered with duplicates + [0, 1, 5, (1 << 32) + 7, (3 << 32) + 4], # spread across bitmap keys, with a gap + ], +) +def test_serialize_round_trips(positions: list[int]) -> None: + dv = DeletionVector.from_positions("file.parquet", positions) + + serialized = dv.serialize() + deserialized = DeletionVector("file.parquet", DeletionVector._deserialize_bitmap(serialized)) + assert deserialized.serialize() == serialized + assert deserialized.to_vector().to_pylist() == sorted(set(positions)) + assert dv.cardinality == len(set(positions)) + + +def test_serialize_skips_empty_bitmaps() -> None: + serialized = DeletionVector.from_positions("file.parquet", [1, (2 << 32) + 1]).serialize() + + # Two non-empty bitmaps (keys 0 and 2); the empty bitmap for key 1 is omitted + assert int.from_bytes(serialized[0:8], "little") == 2 + assert int.from_bytes(serialized[8:12], "little") == 0 + + +def test_serialize_run_optimizes_contiguous_run() -> None: + dv = DeletionVector.from_positions("file.parquet", range(100_000)) + + serialized = dv.serialize() + assert len(serialized) < 1_000 # ~16 KB without run-length encoding + assert DeletionVector._deserialize_bitmap(serialized) == dv._bitmaps + + +def test_from_positions_rejects_negative() -> None: + with pytest.raises(ValueError, match=f"Invalid position: -1, must be between 0 and {MAX_POSITION}"): + DeletionVector.from_positions("file.parquet", [1, -1, 2]) + + +def test_from_positions_rejects_out_of_range() -> None: + too_large = MAX_POSITION + 1 + with pytest.raises(ValueError, match=f"Invalid position: {too_large}, must be between 0 and {MAX_POSITION}"): + DeletionVector.from_positions("file.parquet", [1, too_large, 2]) + + +def test_from_positions_rejects_empty() -> None: + with pytest.raises(ValueError, match="Deletion vector must contain at least one position"): + DeletionVector.from_positions("file.parquet", []) + + +def test_union() -> None: + left = DeletionVector.from_positions("file.parquet", [1, 2]) + right = DeletionVector.from_positions("file.parquet", [2, 3, (1 << 32) + 1]) + + union = left.union(right) + + assert union.to_vector().to_pylist() == [1, 2, 3, (1 << 32) + 1] + assert union.cardinality == 4 + # The operands are unchanged + assert left.to_vector().to_pylist() == [1, 2] + + +def test_union_rejects_different_data_files() -> None: + with pytest.raises(ValueError, match="Cannot union deletion vectors of a.parquet and b.parquet"): + DeletionVector.from_positions("a.parquet", [1]).union(DeletionVector.from_positions("b.parquet", [1])) + + +def test_to_blob_metadata() -> None: + blob = DeletionVector.from_positions("s3://bucket/file.parquet", [1, 2, 3, 3]).to_blob() + + assert blob.metadata.type == DELETION_VECTOR_V1_BLOB_TYPE + assert blob.metadata.fields == [ROW_POSITION_FIELD_ID] + assert blob.metadata.properties == {PROPERTY_REFERENCED_DATA_FILE: "s3://bucket/file.parquet", "cardinality": "3"} + assert blob.metadata.compression_codec is None + # The snapshot id and sequence number are inherited at commit time, as in Java's BaseDVFileWriter + assert blob.metadata.snapshot_id == -1 + assert blob.metadata.sequence_number == -1 + + +def test_to_blob_payload_layout() -> None: + dv = DeletionVector.from_positions("file.parquet", [1, 2, 3]) + payload = dv.to_blob().payload + + # length (4 bytes, big-endian) | magic (4 bytes) | vector | CRC-32 (4 bytes, big-endian), + # where the length and the CRC-32 cover the magic and the vector + length = int.from_bytes(payload[0:4], "big") + assert payload[4:8] == DELETION_VECTOR_MAGIC == b"\xd1\xd3\x39\x64" + vector = payload[8 : 4 + length] + assert vector == dv.serialize() + assert length == 4 + len(vector) + assert len(payload) == 4 + length + 4 + assert int.from_bytes(payload[-4:], "big") == zlib.crc32(DELETION_VECTOR_MAGIC + vector) + + # The reader accepts what the writer produces + assert _deserialize_dv_blob(payload, record_count=3) == dv._bitmaps + + +def _referenced_data_file(file_path: str, partition: Record, spec_id: int) -> DataFile: + return DataFile.from_args( + _table_format_version=3, + spec_id=spec_id, + content=DataFileContent.DATA, + file_path=file_path, + file_format=FileFormat.PARQUET, + partition=partition, + record_count=100, + file_size_in_bytes=1000, + ) + + +def test_write_deletion_vectors(tmp_path: Path) -> None: + io = PyArrowFileIO() + location = str(tmp_path / "deletes.puffin") + first = _referenced_data_file("s3://bucket/a.parquet", Record(1), spec_id=1) + second = _referenced_data_file("s3://bucket/b.parquet", Record(2), spec_id=1) + + delete_files = write_deletion_vectors( + io, + location, + [ + (first, DeletionVector.from_positions(first.file_path, [0, 7])), + (second, DeletionVector.from_positions(second.file_path, [1, 2, 3])), + ], + ) + + file_size = Path(location).stat().st_size + footer_blobs = PuffinFile(Path(location).read_bytes()).footer.blobs + assert len(delete_files) == 2 + for delete_file, data_file, blob, cardinality in zip(delete_files, [first, second], footer_blobs, [2, 3], strict=True): + assert delete_file.content == DataFileContent.POSITION_DELETES + assert delete_file.file_format == FileFormat.PUFFIN + assert delete_file.file_path == location + assert delete_file.file_size_in_bytes == file_size + assert delete_file.partition == data_file.partition + assert delete_file.spec_id == data_file.spec_id + assert delete_file.referenced_data_file == data_file.file_path + assert delete_file.content_offset == blob.offset + assert delete_file.content_size_in_bytes == blob.length + assert delete_file.record_count == cardinality + assert not delete_file.column_sizes + assert not delete_file.lower_bounds + + assert [dv.to_vector().to_pylist() for delete_file in delete_files for dv in read_deletion_vectors(io, delete_file)] == [ + [0, 7], + [1, 2, 3], + ] + + +def test_write_deletion_vectors_rejects_mismatched_data_file(tmp_path: Path) -> None: + data_file = _referenced_data_file("s3://bucket/a.parquet", Record(), spec_id=0) + with pytest.raises(ValueError, match="does not reference s3://bucket/a.parquet"): + write_deletion_vectors( + PyArrowFileIO(), str(tmp_path / "deletes.puffin"), [(data_file, DeletionVector.from_positions("other", [1]))] + ) diff --git a/tests/table/test_init.py b/tests/table/test_init.py index 3d160781e3..eb8995bd3c 100644 --- a/tests/table/test_init.py +++ b/tests/table/test_init.py @@ -27,6 +27,7 @@ from pyiceberg.catalog import Catalog from pyiceberg.catalog.noop import NoopCatalog from pyiceberg.exceptions import CommitFailedException +from pyiceberg.exceptions import ValidationError as IcebergValidationError from pyiceberg.expressions import ( AlwaysTrue, And, @@ -45,6 +46,7 @@ TableIdentifier, ) from pyiceberg.table.metadata import TableMetadataUtil, TableMetadataV1, TableMetadataV2, _generate_snapshot_id +from pyiceberg.table.metadata_columns import LAST_UPDATED_SEQUENCE_NUMBER, ROW_ID from pyiceberg.table.refs import MAIN_BRANCH, SnapshotRef, SnapshotRefType from pyiceberg.table.snapshots import ( MetadataLogEntry, @@ -62,6 +64,7 @@ from pyiceberg.table.statistics import BlobMetadata, PartitionStatisticsFile, StatisticsFile from pyiceberg.table.update import ( AddPartitionSpecUpdate, + AddSchemaUpdate, AddSnapshotUpdate, AddSortOrderUpdate, AssertCreate, @@ -84,6 +87,7 @@ SetPropertiesUpdate, SetSnapshotRefUpdate, SetStatisticsUpdate, + UpgradeFormatVersionUpdate, _apply_table_update, _TableMetadataUpdateContext, update_table_metadata, @@ -102,6 +106,8 @@ DoubleType, FixedType, FloatType, + GeographyType, + GeometryType, IntegerType, ListType, LongType, @@ -110,9 +116,11 @@ PrimitiveType, StringType, StructType, + TimestampNanoType, TimestampType, TimestamptzType, TimeType, + UnknownType, UUIDType, ) @@ -365,6 +373,26 @@ def test_table_scan_projection_unknown_column(table_v2: Table) -> None: assert "Could not find column: 'a'" in str(exc_info.value) +def test_table_scan_projection_metadata_columns(table_v2: Table) -> None: + projection_schema = table_v2.scan(selected_fields=("_last_updated_sequence_number", "y", "_row_id")).projection() + assert projection_schema == Schema( + NestedField(field_id=2, name="y", field_type=LongType(), required=True, doc="comment"), + LAST_UPDATED_SEQUENCE_NUMBER, + ROW_ID, + identifier_field_ids=[2], + ) + assert projection_schema.schema_id == 1 + + full_projection = table_v2.scan(selected_fields=("*", "_row_id")).projection() + assert full_projection.column_names == ["x", "y", "z", "_row_id"] + + case_insensitive = table_v2.scan(selected_fields=("x", "_ROW_ID"), case_sensitive=False).projection() + assert case_insensitive.column_names == ["x", "_row_id"] + + only_metadata = table_v2.scan().select("_row_id").projection() + assert only_metadata.column_names == ["_row_id"] + + def test_data_scan_plan_files_no_current_snapshot(example_table_metadata_no_snapshot_v1: dict[str, Any]) -> None: table = Table( identifier=("default", "test_no_snapshot"), @@ -1983,6 +2011,169 @@ def test_add_snapshot_update_updates_next_row_id(table_v3: Table) -> None: assert new_metadata.next_row_id == 11 +def test_add_snapshot_update_fails_without_added_rows(table_v3: Table) -> None: + new_snapshot = Snapshot( + snapshot_id=25, + parent_snapshot_id=19, + sequence_number=200, + timestamp_ms=1602638593590, + manifest_list="s3:/a/b/c.avro", + summary=Summary(Operation.APPEND), + schema_id=3, + first_row_id=1, + ) + + with pytest.raises(ValueError, match="Cannot add snapshot without added rows"): + update_table_metadata(table_v3.metadata, (AddSnapshotUpdate(snapshot=new_snapshot),)) + + +def test_add_snapshot_update_rejects_older_sequence_number_v3(table_v3: Table) -> None: + new_snapshot = Snapshot( + snapshot_id=25, + parent_snapshot_id=3055729675574597004, + sequence_number=34, + timestamp_ms=1602638593590, + manifest_list="s3:/a/b/c.avro", + summary=Summary(Operation.APPEND), + schema_id=3, + first_row_id=1, + added_rows=0, + ) + + with pytest.raises(ValueError, match="Cannot add snapshot with sequence number 34 older than last sequence number 34"): + update_table_metadata(table_v3.metadata, (AddSnapshotUpdate(snapshot=new_snapshot),)) + + +@pytest.mark.parametrize("table_fixture", [lf("table_v1"), lf("table_v2")]) +def test_add_snapshot_update_leaves_next_row_id_unset_below_v3(table_fixture: Table) -> None: + new_snapshot = Snapshot( + snapshot_id=25, + parent_snapshot_id=table_fixture.metadata.current_snapshot_id, + sequence_number=200, + timestamp_ms=1602638593590, + manifest_list="s3:/a/b/c.avro", + summary=Summary(Operation.APPEND), + schema_id=0, + ) + + new_metadata = update_table_metadata(table_fixture.metadata, (AddSnapshotUpdate(snapshot=new_snapshot),)) + assert "next-row-id" not in json.loads(new_metadata.model_dump_json()) + + +@pytest.mark.parametrize("table_fixture", [lf("table_v1"), lf("table_v2")]) +def test_upgrade_to_v3(table_fixture: Table) -> None: + transaction = table_fixture.transaction().upgrade_table_version(format_version=3) + assert transaction._updates == (UpgradeFormatVersionUpdate(format_version=3),) + + new_metadata = transaction.table_metadata + assert new_metadata.format_version == 3 + assert new_metadata.next_row_id == 0 + # Existing snapshots are left untouched: their rows have no row ids + assert new_metadata.snapshots == table_fixture.metadata.snapshots + assert all(snapshot.first_row_id is None for snapshot in new_metadata.snapshots) + + validated = TableMetadataUtil.parse_obj(json.loads(new_metadata.model_dump_json())) + assert validated.format_version == 3 + assert validated.next_row_id == 0 + + +def test_upgrade_to_same_version_is_noop(table_v3: Table) -> None: + assert update_table_metadata(table_v3.metadata, (UpgradeFormatVersionUpdate(format_version=3),)) == table_v3.metadata + assert table_v3.transaction().upgrade_table_version(format_version=3)._updates == () + + +def test_upgrade_beyond_supported_version_is_rejected(table_v3: Table) -> None: + with pytest.raises(ValueError, match="Unsupported table format version: 4"): + table_v3.transaction().upgrade_table_version(format_version=4) # type: ignore[arg-type] + with pytest.raises(ValueError, match="Unsupported table format version: 4"): + update_table_metadata(table_v3.metadata, (UpgradeFormatVersionUpdate(format_version=4),)) + + +def test_add_schema_update_rejects_v3_types_on_v2(table_v2: Table) -> None: + schema = Schema(NestedField(field_id=1, name="ts", field_type=TimestampNanoType(), required=False), schema_id=10) + with pytest.raises(ValueError, match="timestamp_ns is only supported in 3 or higher. Current format version is: 2"): + update_table_metadata(table_v2.metadata, (AddSchemaUpdate(schema_=schema),)) + + with pytest.raises(ValueError, match="timestamp_ns is only supported in 3 or higher"): + with table_v2.update_schema() as update: + update.add_column("ts_ns", TimestampNanoType()) + + +def test_add_column_with_default_requires_v3(table_v2: Table, table_v3: Table) -> None: + with pytest.raises(ValueError, match="default values require format version 3 or higher, current format version is 2"): + UpdateSchema(transaction=table_v2.transaction()).add_column("added", IntegerType(), default_value=22) + with pytest.raises(ValueError, match="default values require format version 3 or higher"): + UpdateSchema(transaction=table_v2.transaction()).add_column("added", IntegerType(), required=True, default_value=22) + + allowed = UpdateSchema(transaction=table_v2.transaction(), allow_incompatible_changes=True) + assert allowed.add_column("added", IntegerType(), default_value=22)._apply().find_field("added").initial_default == 22 + + v3_schema = UpdateSchema(transaction=table_v3.transaction()).add_column( + "added", IntegerType(), required=True, default_value=22 + ) + field = v3_schema._apply().find_field("added") + assert (field.required, field.initial_default, field.write_default) == (True, 22, 22) + + +@pytest.mark.parametrize("field_type", [GeometryType(), GeographyType(), UnknownType()]) +def test_add_column_null_only_types_reject_defaults(table_v3: Table, field_type: PrimitiveType) -> None: + with pytest.raises(ValueError, match="must default to null"): + UpdateSchema(transaction=table_v3.transaction()).add_column("added", field_type, default_value=b"\x01") + + field = UpdateSchema(transaction=table_v3.transaction()).add_column("added", field_type)._apply().find_field("added") + assert (field.initial_default, field.write_default) == (None, None) + + +def test_set_default_value_requires_v3(table_v2: Table) -> None: + with pytest.raises(ValueError, match="default values require format version 3 or higher, current format version is 2"): + UpdateSchema(transaction=table_v2.transaction()).set_default_value("x", 22) + + allowed = UpdateSchema(transaction=table_v2.transaction(), allow_incompatible_changes=True) + assert allowed.set_default_value("x", 22)._apply().find_field("x").write_default == 22 + + +def test_set_default_value_unknown_rejects_non_null(table_v3: Table) -> None: + with pytest.raises(ValueError, match="must default to null"): + UpdateSchema(transaction=table_v3.transaction()).set_default_value("u", "value") + + +def test_update_column_unknown_promotion_on_v3(table_v3: Table) -> None: + applied = UpdateSchema(transaction=table_v3.transaction()).update_column("u", StringType())._apply() + assert applied.find_type("u") == StringType() + + +@pytest.mark.parametrize("target", [TimestampType(), TimestampNanoType()]) +def test_update_column_date_to_timestamp_requires_v3(table_v2: Table, table_v3: Table, target: PrimitiveType) -> None: + date_schema = Schema(NestedField(10, "d", DateType(), required=False)) + + with pytest.raises(IcebergValidationError, match="requires format version 3 or higher, current format version is 2"): + UpdateSchema(transaction=table_v2.transaction(), schema=date_schema).update_column("d", target) + + applied = UpdateSchema(transaction=table_v3.transaction(), schema=date_schema).update_column("d", target)._apply() + assert applied.find_type("d") == target + + +def test_update_column_date_to_timestamp_rejects_identity_partition(table_v3: Table) -> None: + # The partition spec of table_v3 is identity(x) on field id 1, which would change value after promotion + date_schema = Schema(NestedField(1, "d", DateType(), required=False)) + with pytest.raises(IcebergValidationError, match="the column is partitioned by identity"): + UpdateSchema(transaction=table_v3.transaction(), schema=date_schema).update_column("d", TimestampType()) + + +def test_union_by_name_uses_table_format_version(table_v3: Table) -> None: + import pyarrow as pa + + new_schema = pa.schema([pa.field("ts_ns", pa.timestamp("ns"), nullable=True)]) + applied = UpdateSchema(transaction=table_v3.transaction()).union_by_name(new_schema)._apply() + assert applied.find_type("ts_ns") == TimestampNanoType() + + +def test_add_schema_update_accepts_v3_types_on_v3(table_v3: Table) -> None: + schema = Schema(NestedField(field_id=10, name="ts", field_type=TimestampNanoType(), required=False), schema_id=10) + new_metadata = update_table_metadata(table_v3.metadata, (AddSchemaUpdate(schema_=schema),)) + assert new_metadata.schema_by_id(10) == schema + + def model_roundtrips(model: BaseModel) -> bool: """Helper assertion that tests if a pydantic model roundtrips successfully. @@ -2067,3 +2258,24 @@ def _spy(*args: Any, **kwargs: Any) -> FileIO: assert seen_locations, "expected at least one load_file_io call" assert all(loc is not None for loc in seen_locations), f"load_file_io called without a location: {seen_locations}" + + +def test_metadata_columns_are_read_but_never_written(catalog: Catalog) -> None: + import pyarrow as pa + + from pyiceberg.io.pyarrow import schema_to_pyarrow + + catalog.create_namespace("default") + schema = Schema(NestedField(1, "id", LongType(), required=False)) + table = catalog.create_table("default.metadata_columns", schema=schema) + table.append(pa.Table.from_pylist([{"id": 1}, {"id": 2}], schema=schema_to_pyarrow(schema))) + + result = table.scan(selected_fields=("id", "_row_id", "_last_updated_sequence_number")).to_arrow() + assert result.column_names == ["id", "_row_id", "_last_updated_sequence_number"] + # Data files added without a first_row_id do not carry row lineage + assert result.column("_row_id").to_pylist() == [None, None] + assert result.column("_last_updated_sequence_number").to_pylist() == [None, None] + + with pytest.raises(ValueError, match="PyArrow table contains more columns: _last_updated_sequence_number, _row_id"): + table.append(result) + assert len(table.scan().to_arrow()) == 2 diff --git a/tests/table/test_inspect.py b/tests/table/test_inspect.py index 2404eadc74..52f4c2b690 100644 --- a/tests/table/test_inspect.py +++ b/tests/table/test_inspect.py @@ -15,6 +15,7 @@ # specific language governing permissions and limitations # under the License. +import json from pathlib import PosixPath from typing import Any @@ -25,11 +26,13 @@ from pyiceberg.manifest import DataFile, DataFileContent from pyiceberg.partitioning import PartitionField, PartitionSpec from pyiceberg.schema import Schema +from pyiceberg.table import Table from pyiceberg.table.inspect import InspectTable, _readable_bound +from pyiceberg.table.metadata import TableMetadataV3 from pyiceberg.table.snapshots import Snapshot from pyiceberg.transforms import IdentityTransform from pyiceberg.typedef import Record -from pyiceberg.types import NestedField, StringType +from pyiceberg.types import BinaryType, GeographyType, GeometryType, NestedField, StringType, UnknownType from tests.catalog.test_base import InMemoryCatalog @@ -76,6 +79,27 @@ def test_inspect_entries_and_files_render_null_bound(catalog: InMemoryCatalog) - assert files_metrics["upper_bound"] is None +def test_inspect_files_and_entries_expose_v3_content_file_fields(catalog: InMemoryCatalog) -> None: + schema = Schema(NestedField(1, "s", StringType(), required=False)) + tbl = catalog.create_table("default.v3_fields", schema) + tbl.append(pa.table({"s": ["a"]}, schema=pa.schema([pa.field("s", pa.large_string(), nullable=True)]))) + + v3_fields = ["first_row_id", "referenced_data_file", "content_offset", "content_size_in_bytes"] + for files in (tbl.inspect.files(), tbl.inspect.data_files(), tbl.inspect.delete_files(), tbl.inspect.all_files()): + names = files.column_names + # Matches the column order of Spark's files metadata tables + assert names[names.index("sort_order_id") + 1 : names.index("readable_metrics")] == v3_fields + assert files.schema.field("first_row_id").type == pa.int64() + assert files.schema.field("referenced_data_file").type == pa.string() + + row = tbl.inspect.files().to_pylist()[0] + assert {field: row[field] for field in v3_fields} == dict.fromkeys(v3_fields) + + data_file_type = tbl.inspect.entries().schema.field("data_file").type + data_file_names = [data_file_type.field(i).name for i in range(data_file_type.num_fields)] + assert data_file_names[data_file_names.index("sort_order_id") + 1 :] == v3_fields + + @pytest.mark.parametrize("newest_first", [False, True]) def test_partitions_last_updated_uses_latest_snapshot_regardless_of_order(newest_first: bool) -> None: # Manifest entries are visited in manifest order, which is not chronological, so the @@ -106,3 +130,46 @@ def test_inspect_manifests_preserves_empty_string_bounds(catalog: InMemoryCatalo partition_summary = tbl.inspect.manifests().to_pydict()["partition_summaries"][0][0] assert partition_summary["lower_bound"] == "" assert partition_summary["upper_bound"] == "" + + +def test_readable_bound_for_unknown_type() -> None: + assert _readable_bound(UnknownType(), b"\x00") is None + + +def test_inspect_entries_and_files_with_geo_columns(catalog: InMemoryCatalog) -> None: + wkb_point = bytes.fromhex("0101000000000000000000f03f0000000000000040") + binary_schema = Schema( + NestedField(1, "geom", BinaryType(), required=False), + NestedField(2, "geog", BinaryType(), required=False), + ) + # v3 tables cannot be written yet, so write WKB as binary and read the files back through v3 metadata that + # declares the columns as geometry/geography, like a table written by another engine + tbl = catalog.create_table("default.geo", binary_schema) + tbl.append(pa.table({"geom": [wkb_point, None], "geog": [wkb_point, None]}, schema=binary_schema.as_arrow())) + geo_schema = Schema( + NestedField(1, "geom", GeometryType("srid:3857"), required=False), + NestedField(2, "geog", GeographyType("srid:4326", "vincenty"), required=False), + ) + metadata = json.loads(tbl.metadata.model_dump_json()) + metadata.update({"format-version": 3, "next-row-id": 2, "schemas": [json.loads(geo_schema.model_dump_json())]}) + geo_tbl = Table( + identifier=("default", "geo"), + metadata=TableMetadataV3.model_validate(metadata), + metadata_location=tbl.metadata_location, + io=tbl.io, + catalog=catalog, + ) + + for inspect_table in (geo_tbl.inspect.files(), geo_tbl.inspect.entries(), geo_tbl.inspect.all_files()): + readable_metrics = inspect_table.to_pydict()["readable_metrics"][0] + for column in ("geom", "geog"): + metrics_type = inspect_table.schema.field("readable_metrics").type.field(column).type + assert metrics_type.field("lower_bound").type == pa.binary() + assert metrics_type.field("upper_bound").type == pa.binary() + assert readable_metrics[column]["value_count"] == 2 + assert readable_metrics[column]["null_value_count"] == 1 + + geom = geo_tbl.scan().to_arrow().column("geom") + if isinstance(geom.type, pa.ExtensionType): + geom = geom.cast(geom.type.storage_type) + assert geom.to_pylist() == [wkb_point, None] diff --git a/tests/table/test_metadata.py b/tests/table/test_metadata.py index fb20f9726a..1038519838 100644 --- a/tests/table/test_metadata.py +++ b/tests/table/test_metadata.py @@ -186,12 +186,30 @@ def test_serialize_v2(example_table_metadata_v2: dict[str, Any]) -> None: def test_serialize_v3(example_table_metadata_v3: dict[str, Any]) -> None: - # Writing will be part of https://github.com/apache/iceberg-python/issues/1551 + table_metadata = TableMetadataV3(**example_table_metadata_v3) + serialized = json.loads(table_metadata.model_dump_json()) - with pytest.raises(NotImplementedError) as exc_info: - _ = TableMetadataV3(**example_table_metadata_v3).model_dump_json() + assert serialized["format-version"] == 3 + assert serialized["next-row-id"] == 1 + # Empty encryption keys are left out + assert "encryption-keys" not in serialized + assert TableMetadataV3(**serialized) == table_metadata - assert "Writing V3 is not yet supported, see: https://github.com/apache/iceberg-python/issues/1551" in str(exc_info.value) + +def test_serialize_v3_with_encryption_keys(example_table_metadata_v3: dict[str, Any]) -> None: + encryption_keys = [{"key-id": "key-1", "encrypted-key-metadata": "c2VjcmV0", "encrypted-by-id": "kms-key", "properties": {}}] + table_metadata = TableMetadataV3(**{**example_table_metadata_v3, "encryption-keys": encryption_keys}) + serialized = json.loads(table_metadata.model_dump_json()) + + assert serialized["encryption-keys"] == encryption_keys + assert TableMetadataV3(**serialized) == table_metadata + + +def test_parse_v3_without_next_row_id(example_table_metadata_v3: dict[str, Any]) -> None: + # Defaulting to 0 could hand out row IDs that are already assigned, so the spec-required key must be present + metadata = {key: value for key, value in example_table_metadata_v3.items() if key != "next-row-id"} + with pytest.raises((ValidationError, PydanticValidationError), match="next-row-id"): + TableMetadataV3(**metadata) def test_migrate_v1_schemas(example_table_metadata_v1: dict[str, Any]) -> None: @@ -839,6 +857,7 @@ def test_new_table_metadata_with_v3_schema() -> None: default_sort_order_id=1, refs={}, format_version=3, + next_row_id=0, ) assert actual.model_dump() == expected.model_dump() @@ -847,6 +866,20 @@ def test_new_table_metadata_with_v3_schema() -> None: assert actual.sort_orders == [expected_sort_order] +def test_new_table_metadata_v3_initializes_next_row_id() -> None: + actual = new_table_metadata( + schema=Schema(NestedField(field_id=1, name="foo", field_type=StringType(), required=False)), + partition_spec=PartitionSpec(), + sort_order=SortOrder(), + location="s3://some_v3_location/", + properties={"format-version": "3"}, + ) + + assert actual.format_version == 3 + assert actual.next_row_id == 0 + assert json.loads(actual.model_dump_json())["next-row-id"] == 0 + + @pytest.mark.parametrize( "field_type", [ diff --git a/tests/table/test_partitioning.py b/tests/table/test_partitioning.py index 494addaef1..7abe9d3b4a 100644 --- a/tests/table/test_partitioning.py +++ b/tests/table/test_partitioning.py @@ -31,6 +31,7 @@ IdentityTransform, MonthTransform, TruncateTransform, + VoidTransform, YearTransform, ) from pyiceberg.typedef import Record @@ -45,11 +46,14 @@ PrimitiveType, StringType, StructType, + TimestampNanoType, TimestampType, + TimestamptzNanoType, TimestamptzType, TimeType, UnknownType, UUIDType, + VariantType, ) @@ -105,6 +109,31 @@ def test_partition_compatible_with() -> None: assert not lhs.compatible_with(rhs) +def test_partition_compatible_with_compares_every_source_id() -> None: + # multi-argument fields differ in their later source ids, which the first one cannot show + lhs = PartitionSpec.model_validate( + {"spec-id": 0, "fields": [{"field-id": 1000, "name": "p", "source-ids": [1, 2], "transform": "zorder(a,b)"}]} + ) + rhs = PartitionSpec.model_validate( + {"spec-id": 0, "fields": [{"field-id": 1000, "name": "p", "source-ids": [1, 3], "transform": "zorder(a,c)"}]} + ) + + assert not lhs.compatible_with(rhs) + assert lhs.compatible_with(lhs) + + +def test_partition_compatible_with_compares_unknown_transform_names() -> None: + # same source ids, different transform: the names are all that separates them + lhs = PartitionSpec.model_validate( + {"spec-id": 0, "fields": [{"field-id": 1000, "name": "p", "source-ids": [1, 2], "transform": "zorder(a,b)"}]} + ) + rhs = PartitionSpec.model_validate( + {"spec-id": 0, "fields": [{"field-id": 1000, "name": "p", "source-ids": [1, 2], "transform": "somethingelse(a,b)"}]} + ) + + assert not lhs.compatible_with(rhs) + + def test_unpartitioned() -> None: assert len(UNPARTITIONED_PARTITION_SPEC.fields) == 0 assert UNPARTITIONED_PARTITION_SPEC.is_unpartitioned() @@ -172,6 +201,26 @@ def test_partition_spec_to_path() -> None: assert spec.partition_to_path(record, schema) == "my%23str%25bucket=my%2Bstr/other+str%2Bbucket=%28+%29/my%21int%3Abucket=10" +@pytest.mark.parametrize( + "field_type, expected_path", + [ + (TimestampType(), "ts=2025-02-23T20%3A21%3A44.375612"), + (TimestamptzType(), "ts=2025-02-23T20%3A21%3A44.375612%2B00%3A00"), + (TimestampNanoType(), "ts=2025-02-23T20%3A21%3A44.375612001"), + (TimestamptzNanoType(), "ts=2025-02-23T20%3A21%3A44.375612001%2B00%3A00"), + ], +) +def test_partition_spec_to_path_timestamp(field_type: PrimitiveType, expected_path: str) -> None: + schema = Schema(NestedField(field_id=1, name="ts", field_type=field_type, required=False)) + spec = PartitionSpec(PartitionField(source_id=1, field_id=1000, transform=IdentityTransform(), name="ts")) + + # Nanosecond timestamps carry three extra sub-microsecond digits; the microsecond + # value 1740342104375612 is the same instant as 1740342104375612001 nanoseconds. + value = 1740342104375612001 if isinstance(field_type, (TimestampNanoType, TimestamptzNanoType)) else 1740342104375612 + + assert spec.partition_to_path(Record(value), schema) == expected_path + + def test_partition_spec_to_path_dropped_source_id() -> None: schema = Schema( NestedField(field_id=1, name="str", field_type=StringType(), required=False), @@ -279,6 +328,52 @@ def test_deserialize_partition_field_source_id_and_source_ids_rejected() -> None PartitionField.model_validate_json(json_partition_spec) +def test_deserialize_partition_field_multi_arg() -> None: + import json as json_lib + + from pyiceberg.transforms import UnknownTransform + + json_partition_spec = """{"source-ids": [1, 2], "field-id": 1000, "transform": "bucket[4]", "name": "multi_bucket"}""" + field = PartitionField.model_validate_json(json_partition_spec) + + # v3 readers must read tables with multi-argument transforms, treating them as unknown + assert isinstance(field.transform, UnknownTransform) + assert field.source_id == 1 + assert field.source_ids == [1, 2] + + # the field must round-trip: source-ids only, with the original transform name + serialized = json_lib.loads(field.model_dump_json()) + assert serialized["source-ids"] == [1, 2] + assert "source-id" not in serialized + assert serialized["transform"] == "bucket[4]" + + assert str(field) == "1000: multi_bucket: bucket[4](1, 2)" + + +def test_serialize_partition_field_single_source_id_only() -> None: + import json as json_lib + + json_partition_spec = """{"source-ids": [1], "field-id": 1000, "transform": "truncate[19]", "name": "str_truncate"}""" + field = PartitionField.model_validate_json(json_partition_spec) + serialized = json_lib.loads(field.model_dump_json()) + assert serialized["source-id"] == 1 + assert "source-ids" not in serialized + # a single-element source-ids is normalized onto source-id + assert field.source_ids is None + + +def test_partition_type_with_multi_arg_field() -> None: + from pyiceberg.types import StringType + + schema = Schema(NestedField(1, "a", IntegerType()), NestedField(2, "b", IntegerType())) + field = PartitionField.model_validate_json( + """{"source-ids": [1, 2], "field-id": 1000, "transform": "bucket[4]", "name": "m"}""" + ) + spec = PartitionSpec(field) + struct = spec.partition_type(schema) + assert struct.fields[0].field_type == StringType() + + def test_incompatible_source_column_not_found() -> None: schema = Schema(NestedField(1, "foo", IntegerType()), NestedField(2, "bar", IntegerType())) @@ -310,3 +405,85 @@ def test_incompatible_transform_source_type() -> None: spec.check_compatible(schema) assert "Invalid source field foo with type int for transform: year" in str(exc.value) + + +def test_deserialize_partition_field_multi_arg_requires_transform() -> None: + json_partition_spec = """{"source-ids": [1, 2], "field-id": 1000, "name": "m"}""" + with pytest.raises(Exception, match="Transform is required for a multi-argument field"): + PartitionField.model_validate_json(json_partition_spec) + + +def test_partition_field_transform_arguments() -> None: + single = PartitionField(source_id=1, field_id=1000, transform=TruncateTransform(width=19), name="str_truncate") + assert single.transform_arguments == [1] + assert single.is_multi_argument is False + + multi = PartitionField.model_validate_json( + """{"source-ids": [1, 2], "field-id": 1000, "transform": "bucket[4]", "name": "m"}""" + ) + assert multi.transform_arguments == [1, 2] + assert multi.is_multi_argument is True + + +@pytest.mark.parametrize("transform", [IdentityTransform(), BucketTransform(4), TruncateTransform(4), VoidTransform()]) +def test_variant_cannot_be_partition_source(transform: Any) -> None: + schema = Schema(NestedField(1, "id", IntegerType()), NestedField(2, "v", VariantType())) + assert not transform.can_transform(VariantType()) + if not isinstance(transform, VoidTransform): + spec = PartitionSpec(PartitionField(source_id=2, field_id=1000, transform=transform, name="v_part")) + with pytest.raises(ValidationError, match="Invalid source field v with type variant"): + spec.check_compatible(schema) + + +def test_multi_arg_partition_spec_in_table_metadata() -> None: + from pyiceberg.expressions import EqualTo + from pyiceberg.expressions.visitors import inclusive_projection + from pyiceberg.table.metadata import TableMetadataUtil + from pyiceberg.transforms import UnknownTransform + + metadata = TableMetadataUtil.parse_obj( + { + "format-version": 3, + "table-uuid": "9c12d441-03fe-4693-9a96-a0705ddf69c1", + "location": "s3://bucket/table", + "last-sequence-number": 0, + "last-updated-ms": 1602638573590, + "last-column-id": 2, + "current-schema-id": 0, + "schemas": [ + { + "type": "struct", + "schema-id": 0, + "fields": [ + {"id": 1, "name": "a", "required": False, "type": "int"}, + {"id": 2, "name": "b", "required": False, "type": "int"}, + ], + } + ], + "default-spec-id": 0, + "partition-specs": [ + {"spec-id": 0, "fields": [{"name": "ab", "transform": "zorder", "source-ids": [1, 2], "field-id": 1000}]} + ], + "last-partition-id": 1000, + "default-sort-order-id": 1, + "sort-orders": [ + { + "order-id": 1, + "fields": [{"transform": "zorder", "source-ids": [1, 2], "direction": "asc", "null-order": "nulls-first"}], + } + ], + "next-row-id": 0, + } + ) + + spec = metadata.spec() + assert isinstance(spec.fields[0].transform, UnknownTransform) + assert spec.fields[0].transform_arguments == [1, 2] + assert metadata.sort_order_by_id(1).fields[0].transform_arguments == [1, 2] # type: ignore[union-attr] + # unknown transforms produce no partition predicate, so nothing is pruned + projected = inclusive_projection(metadata.schema(), spec)(EqualTo("a", 1)) + assert str(projected) == "AlwaysTrue()" + # the metadata round-trips with source-ids + serialized = metadata.model_dump_json() + assert '"source-ids":[1,2]' in serialized + assert TableMetadataUtil.parse_raw(serialized).spec() == spec diff --git a/tests/table/test_puffin.py b/tests/table/test_puffin.py index 93b16158bb..c692022fdb 100644 --- a/tests/table/test_puffin.py +++ b/tests/table/test_puffin.py @@ -14,9 +14,28 @@ # KIND, either express or implied. See the License for the # specific language governing permissions and limitations # under the License. +import json +import struct +import zlib from os import path +from pathlib import Path -from pyiceberg.table.puffin import PuffinFile +import pyarrow as pa +import pytest +from pyroaring import BitMap + +from pyiceberg import __version__ +from pyiceberg.io.pyarrow import PyArrowFileIO +from pyiceberg.manifest import DataFile, DataFileContent, FileFormat +from pyiceberg.table.deletion_vector import ( + _DV_BLOB_MAGIC_NUMBER, + DELETION_VECTOR_V1_BLOB_TYPE, + PROPERTY_REFERENCED_DATA_FILE, + DeletionVector, + deletion_vectors_from_puffin_file, + read_deletion_vectors, +) +from pyiceberg.table.puffin import MAGIC_BYTES, PuffinFile, PuffinWriter def _open_file(file: str) -> bytes: @@ -79,3 +98,204 @@ def test_read_two_blobs_uncompressed() -> None: assert pf.get_blob_payload(blob2) == ( b"some blob \x00 binary data \xf0\x9f\xa4\xaf that is not very very very very very very long, is it?" ) + + +def _dv_blob(positions: list[int]) -> bytes: + bitmap_data = ( + struct.pack("I", len(bitmap_data)) + bitmap_data + struct.pack(">I", zlib.crc32(bitmap_data) & 0xFFFFFFFF) + + +def _puffin_with_blobs(blobs: list[tuple[str, str, bytes]]) -> tuple[bytes, list[tuple[int, int]]]: + """Build a Puffin file from (type, referenced data file, payload) tuples, returning the (offset, length) of each blob.""" + body = MAGIC_BYTES + ranges = [] + blob_metadata = [] + for blob_type, referenced_data_file, payload in blobs: + ranges.append((len(body), len(payload))) + blob_metadata.append( + { + "type": blob_type, + "fields": [2147483545], + "snapshot-id": 1, + "sequence-number": 1, + "offset": len(body), + "length": len(payload), + "properties": {PROPERTY_REFERENCED_DATA_FILE: referenced_data_file}, + } + ) + body += payload + footer_payload = json.dumps({"blobs": blob_metadata, "properties": {}}).encode() + footer = ( + MAGIC_BYTES + footer_payload + len(footer_payload).to_bytes(4, byteorder="little") + b"\x00\x00\x00\x00" + MAGIC_BYTES + ) + return body + footer, ranges + + +def _dv_data_file(delete_file_path: str, referenced_data_file: str, offset: int, length: int, cardinality: int) -> DataFile: + return DataFile.from_args( + _table_format_version=3, + content=DataFileContent.POSITION_DELETES, + file_path=delete_file_path, + file_format=FileFormat.PUFFIN, + record_count=cardinality, + referenced_data_file=referenced_data_file, + content_offset=offset, + content_size_in_bytes=length, + ) + + +def test_read_two_deletion_vectors_by_range(tmp_path: Path) -> None: + puffin_bytes, ranges = _puffin_with_blobs( + [ + (DELETION_VECTOR_V1_BLOB_TYPE, "s3://bucket/a.parquet", _dv_blob([1, 3])), + (DELETION_VECTOR_V1_BLOB_TYPE, "s3://bucket/b.parquet", _dv_blob([0, 2, 4])), + ] + ) + delete_file_path = str(tmp_path / "deletes.puffin") + with open(delete_file_path, "wb") as f: + f.write(puffin_bytes) + + io = PyArrowFileIO() + (first,) = read_deletion_vectors(io, _dv_data_file(delete_file_path, "s3://bucket/a.parquet", *ranges[0], cardinality=2)) + (second,) = read_deletion_vectors(io, _dv_data_file(delete_file_path, "s3://bucket/b.parquet", *ranges[1], cardinality=3)) + + assert first.referenced_data_file == "s3://bucket/a.parquet" + assert first.to_vector() == pa.chunked_array([[1, 3]]) + assert second.referenced_data_file == "s3://bucket/b.parquet" + assert second.to_vector() == pa.chunked_array([[0, 2, 4]]) + + # Reading the whole file yields both vectors + assert {dv.referenced_data_file: dv.to_vector() for dv in deletion_vectors_from_puffin_file(PuffinFile(puffin_bytes))} == { + "s3://bucket/a.parquet": pa.chunked_array([[1, 3]]), + "s3://bucket/b.parquet": pa.chunked_array([[0, 2, 4]]), + } + + +def test_deletion_vectors_from_puffin_file_skips_other_blob_types() -> None: + puffin_bytes, _ = _puffin_with_blobs( + [ + ("apache-datasketches-theta-v1", "s3://bucket/a.parquet", b"not-a-deletion-vector"), + (DELETION_VECTOR_V1_BLOB_TYPE, "s3://bucket/b.parquet", _dv_blob([7])), + ] + ) + + deletion_vectors = deletion_vectors_from_puffin_file(PuffinFile(puffin_bytes)) + + assert [dv.referenced_data_file for dv in deletion_vectors] == ["s3://bucket/b.parquet"] + + +def _write(tmp_path: Path, *deletion_vectors: DeletionVector, created_by: str | None = None) -> Path: + puffin_path = tmp_path / "test.puffin" + with PuffinWriter(PyArrowFileIO().new_output(str(puffin_path)), created_by=created_by) as writer: + for dv in deletion_vectors: + writer.add_blob(dv.to_blob()) + return puffin_path + + +def test_puffin_writer_round_trips_single_blob(tmp_path: Path) -> None: + positions = [0, 1, 5, (1 << 32) + 7] + puffin_path = _write(tmp_path, DeletionVector.from_positions("file.parquet", positions)) + + dvs = deletion_vectors_from_puffin_file(PuffinFile(puffin_path.read_bytes())) + + assert len(dvs) == 1 + assert dvs[0].referenced_data_file == "file.parquet" + assert dvs[0].to_vector().to_pylist() == positions + + +def test_puffin_writer_round_trips_multiple_blobs(tmp_path: Path) -> None: + puffin_path = _write( + tmp_path, + DeletionVector.from_positions("file1.parquet", [1, 2, 3]), + DeletionVector.from_positions("file2.parquet", [4, 5, 6]), + ) + + dvs = deletion_vectors_from_puffin_file(PuffinFile(puffin_path.read_bytes())) + + assert {dv.referenced_data_file: dv.to_vector().to_pylist() for dv in dvs} == { + "file1.parquet": [1, 2, 3], + "file2.parquet": [4, 5, 6], + } + + +def test_puffin_writer_layout(tmp_path: Path) -> None: + puffin_path = tmp_path / "test.puffin" + first = DeletionVector.from_positions("file1.parquet", [1, 2, 3]).to_blob() + second = DeletionVector.from_positions("file2.parquet", [4]).to_blob() + with PuffinWriter(PyArrowFileIO().new_output(str(puffin_path))) as writer: + first_metadata = writer.add_blob(first) + second_metadata = writer.add_blob(second) + puffin_bytes = puffin_path.read_bytes() + + # The offsets and lengths are assigned when a blob is added + assert (first_metadata.offset, first_metadata.length) == (4, len(first.payload)) + assert (second_metadata.offset, second_metadata.length) == (4 + len(first.payload), len(second.payload)) + assert puffin_bytes[first_metadata.offset : first_metadata.offset + first_metadata.length] == first.payload + assert puffin_bytes[second_metadata.offset : second_metadata.offset + second_metadata.length] == second.payload + + # Magic | blobs | magic | footer payload | footer payload size | flags | magic + footer_start = second_metadata.offset + second_metadata.length + footer_size = int.from_bytes(puffin_bytes[-12:-8], "little") + assert puffin_bytes[:4] == MAGIC_BYTES + assert puffin_bytes[footer_start : footer_start + 4] == MAGIC_BYTES + assert len(puffin_bytes) == footer_start + 4 + footer_size + 12 + assert puffin_bytes[-8:-4] == b"\x00\x00\x00\x00" + assert puffin_bytes[-4:] == MAGIC_BYTES + assert writer.file_size == len(puffin_bytes) + + footer = json.loads(puffin_bytes[footer_start + 4 : footer_start + 4 + footer_size]) + assert footer["blobs"][0] == { + "type": "deletion-vector-v1", + "fields": [2147483645], + "snapshot-id": -1, + "sequence-number": -1, + "offset": 4, + "length": len(first.payload), + "properties": {"referenced-data-file": "file1.parquet", "cardinality": "3"}, + } + assert PuffinFile(puffin_bytes).footer.blobs == [first_metadata, second_metadata] + + +def test_puffin_writer_created_by(tmp_path: Path) -> None: + default = PuffinFile(_write(tmp_path, DeletionVector.from_positions("file.parquet", [1])).read_bytes()) + assert default.footer.properties == {"created-by": f"PyIceberg version {__version__}"} + + custom = PuffinFile(_write(tmp_path, created_by="my-test-app").read_bytes()) + assert custom.footer.properties == {"created-by": "my-test-app"} + + +def test_puffin_writer_empty(tmp_path: Path) -> None: + reader = PuffinFile(_write(tmp_path).read_bytes()) + + assert reader.footer.blobs == [] + assert deletion_vectors_from_puffin_file(reader) == [] + + +def test_puffin_writer_does_not_write_on_exception(tmp_path: Path) -> None: + output_file = PyArrowFileIO().new_output(str(tmp_path / "test.puffin")) + + def _fail() -> None: + raise ValueError("boom") + + with pytest.raises(ValueError, match="boom"): + with PuffinWriter(output_file) as writer: + writer.add_blob(DeletionVector.from_positions("file.parquet", [1]).to_blob()) + _fail() + + assert writer.closed + assert not output_file.exists() + + +def test_puffin_writer_rejects_use_after_close(tmp_path: Path) -> None: + writer = PuffinWriter(PyArrowFileIO().new_output(str(tmp_path / "test.puffin"))) + writer.close() + + with pytest.raises(RuntimeError, match="Cannot add blob to closed Puffin writer"): + writer.add_blob(DeletionVector.from_positions("file.parquet", [2]).to_blob()) + with pytest.raises(RuntimeError, match="Puffin writer is already closed"): + writer.close() diff --git a/tests/table/test_snapshots.py b/tests/table/test_snapshots.py index 39aeb7b349..acdadabde5 100644 --- a/tests/table/test_snapshots.py +++ b/tests/table/test_snapshots.py @@ -28,7 +28,7 @@ from pyiceberg.catalog import Catalog from pyiceberg.exceptions import ValidationException from pyiceberg.io.pyarrow import _dataframe_to_data_files -from pyiceberg.manifest import DataFile, DataFileContent, ManifestContent, ManifestFile +from pyiceberg.manifest import DataFile, DataFileContent, FileFormat, ManifestContent, ManifestEntryStatus, ManifestFile from pyiceberg.partitioning import PartitionField, PartitionSpec from pyiceberg.schema import Schema from pyiceberg.table import Table @@ -821,3 +821,224 @@ def test_overwrite_rejects_explicit_delete_without_parent_snapshot( with empty.transaction() as tx: with tx.update_snapshot().overwrite() as overwrite: overwrite.delete_data_file(stale_file) + + +def test_snapshot_summary_collector_deletion_vectors(table_schema_simple: Schema) -> None: + def _delete_file(file_format: FileFormat, record_count: int) -> DataFile: + return DataFile.from_args( + _table_format_version=3, + content=DataFileContent.POSITION_DELETES, + file_format=file_format, + record_count=record_count, + file_size_in_bytes=1000, + content_size_in_bytes=40 if file_format == FileFormat.PUFFIN else None, + partition=Record(), + ) + + ssc = SnapshotSummaryCollector() + ssc.add_file(_delete_file(FileFormat.PUFFIN, 3), schema=table_schema_simple) + ssc.remove_file(_delete_file(FileFormat.PUFFIN, 1), schema=table_schema_simple) + ssc.remove_file(_delete_file(FileFormat.PARQUET, 2), schema=table_schema_simple) + + # As in Java, a DV is not a position delete file, and its size is the size of its blob + assert ssc.build() == { + "added-delete-files": "1", + "added-dvs": "1", + "added-files-size": "40", + "added-position-deletes": "3", + "removed-delete-files": "2", + "removed-dvs": "1", + "removed-files-size": "1040", + "removed-position-delete-files": "1", + "removed-position-deletes": "3", + } + + +def _merge_on_read_table(catalog: Catalog, format_version: int = 3, partition_spec: PartitionSpec | None = None) -> Table: + catalog.create_namespace("default") + schema = Schema( + NestedField(1, "id", IntegerType(), required=False), + NestedField(2, "category", StringType(), required=False), + ) + table = catalog.create_table( + "default.merge_on_read", + schema=schema, + partition_spec=partition_spec or PartitionSpec(), + properties={"format-version": str(format_version), "write.delete.mode": "merge-on-read"}, + ) + table.append(pa.table({"id": pa.array(range(10), pa.int32()), "category": ["a"] * 5 + ["b"] * 5})) + return table + + +def _live_delete_files(table: Table) -> list[DataFile]: + snapshot = table.current_snapshot() + assert snapshot is not None + return [ + entry.data_file + for manifest in snapshot.manifests(table.io) + if manifest.content == ManifestContent.DELETES + for entry in manifest.fetch_manifest_entry(table.io, discard_deleted=True) + ] + + +def _row_ids(table: Table) -> dict[int, int]: + rows = table.scan(selected_fields=("id", "_row_id")).to_arrow() + return dict(zip(rows["id"].to_pylist(), rows["_row_id"].to_pylist(), strict=True)) + + +def test_merge_on_read_delete_writes_deletion_vector(catalog: Catalog) -> None: + table = _merge_on_read_table(catalog) + data_file = next(iter(table.scan().plan_files())).file + row_ids_before = _row_ids(table) + + table.delete("id == 2 or id == 7") + + snapshot = table.current_snapshot() + assert snapshot is not None and snapshot.summary is not None + assert snapshot.summary.operation == Operation.DELETE + assert snapshot.summary["added-dvs"] == "1" + assert snapshot.summary["added-delete-files"] == "1" + assert snapshot.summary["added-position-deletes"] == "2" + assert snapshot.summary["total-delete-files"] == "1" + assert snapshot.summary["total-position-deletes"] == "2" + assert snapshot.summary["total-data-files"] == "1" + assert snapshot.summary["total-records"] == "10" + # No rows are added, so no row ids are assigned + assert snapshot.added_rows == 0 + + (dv,) = _live_delete_files(table) + assert dv.content == DataFileContent.POSITION_DELETES + assert dv.file_format == FileFormat.PUFFIN + assert dv.file_path.endswith("-deletes.puffin") + assert dv.referenced_data_file == data_file.file_path + assert dv.content_offset == 4 + assert dv.content_size_in_bytes is not None and dv.content_size_in_bytes > 0 + assert dv.record_count == 2 + + assert sorted(table.scan().to_arrow()["id"].to_pylist()) == [0, 1, 3, 4, 5, 6, 8, 9] + assert table.scan().count() == 8 + # The data file is not rewritten, so the row ids of the remaining rows are unchanged + row_ids_after = _row_ids(table) + assert row_ids_after == {id_: row_id for id_, row_id in row_ids_before.items() if id_ not in (2, 7)} + + +def test_merge_on_read_delete_replaces_deletion_vector(catalog: Catalog) -> None: + table = _merge_on_read_table(catalog) + table.delete("id == 2") + (previous_dv,) = _live_delete_files(table) + + table.delete("id in (2, 5, 6)") + + snapshot = table.current_snapshot() + assert snapshot is not None and snapshot.summary is not None + assert snapshot.summary.operation == Operation.DELETE + assert snapshot.summary["added-dvs"] == "1" + assert snapshot.summary["removed-dvs"] == "1" + assert snapshot.summary["removed-delete-files"] == "1" + assert snapshot.summary["added-position-deletes"] == "3" + assert snapshot.summary["removed-position-deletes"] == "1" + assert snapshot.summary["total-delete-files"] == "1" + assert snapshot.summary["total-position-deletes"] == "3" + + # One DV per data file, holding the union of the deleted positions + (dv,) = _live_delete_files(table) + assert dv.file_path != previous_dv.file_path + assert dv.record_count == 3 + assert sorted(table.scan().to_arrow()["id"].to_pylist()) == [0, 1, 3, 4, 7, 8, 9] + + # The previous DV is recorded as deleted in this snapshot + deleted = [ + entry + for manifest in snapshot.manifests(table.io) + if manifest.content == ManifestContent.DELETES + for entry in manifest.fetch_manifest_entry(table.io, discard_deleted=False) + if entry.status == ManifestEntryStatus.DELETED + ] + assert [(entry.data_file.file_path, entry.snapshot_id) for entry in deleted] == [ + (previous_dv.file_path, snapshot.snapshot_id) + ] + + +def test_merge_on_read_delete_of_deleted_rows_is_a_no_op(catalog: Catalog) -> None: + table = _merge_on_read_table(catalog) + table.delete("id == 2") + snapshot_count = len(table.snapshots()) + + table.delete("id == 2") + + assert len(table.snapshots()) == snapshot_count + + +def test_merge_on_read_delete_drops_fully_matching_files(catalog: Catalog) -> None: + table = _merge_on_read_table( + catalog, + partition_spec=PartitionSpec(PartitionField(source_id=2, field_id=1000, transform=IdentityTransform(), name="category")), + ) + + table.delete("category == 'a'") + + snapshot = table.current_snapshot() + assert snapshot is not None and snapshot.summary is not None + assert snapshot.summary["deleted-data-files"] == "1" + assert "added-dvs" not in snapshot.summary.additional_properties + assert _live_delete_files(table) == [] + assert sorted(table.scan().to_arrow()["id"].to_pylist()) == [5, 6, 7, 8, 9] + + +def test_merge_on_read_delete_partitioned(catalog: Catalog) -> None: + table = _merge_on_read_table( + catalog, + partition_spec=PartitionSpec(PartitionField(source_id=2, field_id=1000, transform=IdentityTransform(), name="category")), + ) + data_files = {task.file.file_path: task.file for task in table.scan().plan_files()} + + table.delete("id in (1, 8)") + + dvs = _live_delete_files(table) + assert len(dvs) == 2 + # All DVs of a commit are written to one Puffin file, as one blob per data file + assert len({dv.file_path for dv in dvs}) == 1 + assert len({dv.content_offset for dv in dvs}) == 2 + for dv in dvs: + assert dv.referenced_data_file is not None + data_file = data_files[dv.referenced_data_file] + assert dv.partition == data_file.partition + assert dv.spec_id == data_file.spec_id + assert dv.record_count == 1 + assert sorted(table.scan().to_arrow()["id"].to_pylist()) == [0, 2, 3, 4, 5, 6, 7, 9] + assert table.scan(row_filter="category == 'b'").count() == 4 + + +def test_merge_on_read_delete_requires_format_version_3(catalog: Catalog) -> None: + table = _merge_on_read_table(catalog, format_version=2) + + with pytest.warns(UserWarning, match=re.escape("Merge-on-read deletes require format version 3 (deletion vectors)")): + table.delete("id == 2") + + snapshot = table.current_snapshot() + assert snapshot is not None and snapshot.summary is not None + assert snapshot.summary.operation == Operation.OVERWRITE + assert _live_delete_files(table) == [] + assert table.scan().count() == 9 + + +def test_overwrite_keeps_copy_on_write_with_merge_on_read(catalog: Catalog) -> None: + table = _merge_on_read_table(catalog) + + table.overwrite(pa.table({"id": pa.array([100], pa.int32()), "category": ["c"]}), overwrite_filter="id == 2") + + assert _live_delete_files(table) == [] + assert sorted(table.scan().to_arrow()["id"].to_pylist()) == [0, 1, 3, 4, 5, 6, 7, 8, 9, 100] + + +def test_failed_merge_on_read_delete_removes_puffin_file(catalog: Catalog) -> None: + table = _merge_on_read_table(catalog) + concurrent = catalog.load_table(table.name()) + concurrent.delete("id == 3") + + with pytest.raises(ValidationException): + table.delete("id == 2") + + data_path = Path(urlparse(table.location()).path) / "data" + assert len(list(data_path.glob("*-deletes.puffin"))) == 1 + assert sorted(catalog.load_table(table.name()).scan().to_arrow()["id"].to_pylist()) == [0, 1, 2, 4, 5, 6, 7, 8, 9] diff --git a/tests/table/test_sorting.py b/tests/table/test_sorting.py index f8329091b4..a047578430 100644 --- a/tests/table/test_sorting.py +++ b/tests/table/test_sorting.py @@ -181,3 +181,50 @@ def test_incompatible_transform_source_type() -> None: sort_order.check_compatible(schema) assert "Invalid source field foo with type int for transform: year" in str(exc.value) + + +def test_variant_cannot_be_sort_source() -> None: + from pyiceberg.types import VariantType + + schema = Schema(NestedField(1, "v", VariantType())) + sort_order = SortOrder(SortField(source_id=1, transform=IdentityTransform(), null_order=NullOrder.NULLS_FIRST)) + + with pytest.raises(ValidationError, match="Invalid source field v with type variant for transform: identity"): + sort_order.check_compatible(schema) + + +def test_deserialize_sort_field_multi_arg() -> None: + from pyiceberg.transforms import UnknownTransform + + payload = '{"source-ids":[19,20],"transform":"bucket[4]","direction":"asc","null-order":"nulls-first"}' + field = SortField.model_validate_json(payload) + + # v3 readers must read tables with multi-argument transforms, treating them as unknown + assert isinstance(field.transform, UnknownTransform) + assert field.source_id == 19 + assert field.source_ids == [19, 20] + + serialized = json.loads(field.model_dump_json()) + assert serialized["source-ids"] == [19, 20] + assert "source-id" not in serialized + assert serialized["transform"] == "bucket[4]" + + assert str(field) == "bucket[4](19, 20) ASC NULLS FIRST" + + +def test_deserialize_sort_field_multi_arg_requires_transform() -> None: + payload = '{"source-ids":[19,20],"direction":"asc","null-order":"nulls-first"}' + with pytest.raises(Exception, match="Transform is required for a multi-argument field"): + SortField.model_validate_json(payload) + + +def test_sort_field_transform_arguments() -> None: + single = SortField(source_id=19, transform=BucketTransform(num_buckets=4), null_order=NullOrder.NULLS_FIRST) + assert single.transform_arguments == [19] + assert single.is_multi_argument is False + + multi = SortField.model_validate_json( + '{"source-ids":[19,20],"transform":"bucket[4]","direction":"asc","null-order":"nulls-first"}' + ) + assert multi.transform_arguments == [19, 20] + assert multi.is_multi_argument is True diff --git a/tests/test_conversions.py b/tests/test_conversions.py index e786ae0683..4da884927f 100644 --- a/tests/test_conversions.py +++ b/tests/test_conversions.py @@ -100,6 +100,8 @@ DoubleType, FixedType, FloatType, + GeographyType, + GeometryType, IntegerType, LongType, PrimitiveType, @@ -111,6 +113,7 @@ TimeType, UnknownType, UUIDType, + VariantType, ) @@ -609,6 +612,51 @@ def test_json_serialize_roundtrip(primitive_type: PrimitiveType, value: Any) -> assert value == conversions.from_json(primitive_type, conversions.to_json(primitive_type, value)) +@pytest.mark.parametrize( + "primitive_type, value, expected", + [ + (TimestampNanoType(), 1510871468123456789, "2017-11-16T22:31:08.123456789"), + (TimestampNanoType(), -1, "1969-12-31T23:59:59.999999999"), + (TimestampNanoType(), datetime(2017, 11, 16, 22, 31, 8, 123456), "2017-11-16T22:31:08.123456000"), + (TimestamptzNanoType(), 1510871468123456789, "2017-11-16T22:31:08.123456789+00:00"), + ( + TimestamptzNanoType(), + datetime(2017, 11, 16, 22, 31, 8, 123456, tzinfo=timezone.utc), + "2017-11-16T22:31:08.123456000+00:00", + ), + ], +) +def test_json_timestamp_nano_serialization(primitive_type: PrimitiveType, value: Any, expected: str) -> None: + assert conversions.to_json(primitive_type, value) == expected + + +@pytest.mark.parametrize( + "primitive_type, json_value, expected", + [ + (TimestampNanoType(), "2017-11-16T22:31:08.123456789", 1510871468123456789), + (TimestampNanoType(), "2017-11-16T22:31:08.1", 1510871468100000000), + (TimestampNanoType(), "2017-11-16T22:31:08", 1510871468000000000), + (TimestamptzNanoType(), "2017-11-16T22:31:08.123456789+00:00", 1510871468123456789), + (TimestamptzNanoType(), "2017-11-16T14:31:08.123456789-08:00", 1510871468123456789), + ], +) +def test_json_timestamp_nano_deserialization(primitive_type: PrimitiveType, json_value: str, expected: int) -> None: + assert conversions.from_json(primitive_type, json_value) == expected + + +@pytest.mark.parametrize("primitive_type", [TimestampNanoType(), TimestamptzNanoType()]) +def test_json_timestamp_nano_roundtrip(primitive_type: PrimitiveType) -> None: + value = 1510871468123456789 + assert conversions.from_json(primitive_type, conversions.to_json(primitive_type, value)) == value + + +def test_json_timestamp_nano_zone_errors() -> None: + with pytest.raises(ValueError, match="Zone offset provided, but not expected"): + conversions.from_json(TimestampNanoType(), "2017-11-16T22:31:08.123456789+00:00") + with pytest.raises(ValueError, match="Missing zone offset"): + conversions.from_json(TimestamptzNanoType(), "2017-11-16T22:31:08.123456789") + + def test_unknown_type_conversions() -> None: """Unknown values are always null, so they have no binary or partition representation.""" unknown = UnknownType() @@ -619,8 +667,27 @@ def test_unknown_type_conversions() -> None: # The spec defines no JSON single-value representation for unknown, since a # non-null initial-default or write-default is invalid for the type. - with pytest.raises(TypeError): + assert conversions.to_json(unknown, None) is None + assert conversions.from_json(unknown, None) is None + with pytest.raises(ValueError, match="must default to null"): conversions.to_json(unknown, "iceberg") - with pytest.raises(TypeError): + with pytest.raises(ValueError, match="must default to null"): conversions.from_json(unknown, "iceberg") + + +@pytest.mark.parametrize("geo_type", [GeometryType(), GeographyType()]) +def test_geo_json_null_default(geo_type: PrimitiveType) -> None: + """Geometry and geography columns must default to null, so null deserializes to None.""" + assert conversions.from_json(geo_type, None) is None + + +def test_variant_type_conversions() -> None: + assert conversions.to_json(VariantType(), None) is None + assert conversions.from_json(VariantType(), None) is None + with pytest.raises(ValueError, match="must default to null"): + conversions.from_json(VariantType(), {"a": 1}) + with pytest.raises(ValueError, match="no single-value binary serialization"): + conversions.to_bytes(VariantType(), b"\x00") + with pytest.raises(ValueError, match="no single-value binary serialization"): + conversions.from_bytes(VariantType(), b"\x00") diff --git a/tests/test_schema.py b/tests/test_schema.py index 94931277b4..68846d1a32 100644 --- a/tests/test_schema.py +++ b/tests/test_schema.py @@ -26,6 +26,7 @@ Accessor, Schema, _check_schema_compatible, + assign_fresh_schema_ids, build_position_accessors, index_by_id, index_by_name, @@ -44,6 +45,8 @@ DoubleType, FixedType, FloatType, + GeographyType, + GeometryType, IcebergType, IntegerType, ListType, @@ -53,11 +56,13 @@ PrimitiveType, StringType, StructType, + TimestampNanoType, TimestampType, TimestamptzType, TimeType, UnknownType, UUIDType, + VariantType, ) TEST_PRIMITIVE_TYPES = [ @@ -897,6 +902,8 @@ def should_promote(file_type: IcebergType, read_type: IcebergType) -> bool: return can_promote_decimal(file_type, read_type) if isinstance(file_type, FixedType) and isinstance(read_type, UUIDType) and len(file_type) == 16: return True + if isinstance(file_type, DateType) and isinstance(read_type, (TimestampType, TimestampNanoType)): + return True return False @@ -974,6 +981,23 @@ def test_decimal_promotion() -> None: promote(DecimalType(18, 2), DecimalType(9, 2)) +@pytest.mark.parametrize("read_type", [GeometryType(), GeometryType("srid:3857"), GeographyType("srid:4326", "vincenty")]) +def test_binary_resolves_to_geo_types(read_type: IcebergType) -> None: + # Parquet GEOMETRY/GEOGRAPHY columns are reported as binary by PyArrow + assert promote(BinaryType(), read_type) == read_type + with pytest.raises(ResolveError): + promote(StringType(), read_type) + with pytest.raises(ResolveError): + promote(FixedType(16), read_type) + + +def test_update_column_binary_to_geometry_is_rejected(table_v2: Table) -> None: + current_schema = Schema(NestedField(field_id=1, name="aCol", field_type=BinaryType(), required=False)) + update = UpdateSchema(transaction=Transaction(table_v2), schema=current_schema) + with pytest.raises(ValidationError, match="Cannot change column type: aCol: binary -> geometry"): + update.update_column("aCol", GeometryType()) + + def test_unknown_type_promotion_to_primitive() -> None: """Test that UnknownType can be promoted to primitive types (V3+ behavior)""" unknown_type = UnknownType() @@ -984,6 +1008,21 @@ def test_unknown_type_promotion_to_primitive() -> None: assert promote(unknown_type, FloatType()) == FloatType() +def test_variant_group_struct_promotion_to_variant() -> None: + variant_group = StructType( + NestedField(3, "metadata", BinaryType(), required=True), NestedField(4, "value", BinaryType(), required=True) + ) + assert promote(variant_group, VariantType()) == VariantType() + + with pytest.raises(ResolveError, match="Cannot promote"): + promote(StructType(NestedField(3, "metadata", BinaryType(), required=True)), VariantType()) + with pytest.raises(ResolveError, match="Cannot promote"): + promote( + StructType(NestedField(3, "metadata", BinaryType()), NestedField(4, "value", StringType())), + VariantType(), + ) + + def test_unknown_type_promotion_to_non_primitive_raises_resolve_error() -> None: """Test that UnknownType cannot be promoted to non-primitive types and raises ResolveError""" unknown_type = UnknownType() @@ -1844,3 +1883,14 @@ def test_check_schema_compatible_optional_map_field_present() -> None: ) # Should not raise - schemas match _check_schema_compatible(requested_schema, provided_schema) + + +def test_assign_fresh_schema_ids_keeps_defaults() -> None: + schema = Schema( + NestedField(10, "c", StringType(), required=True, initial_default="init", write_default="wd"), + NestedField(11, "s", StructType(NestedField(12, "n", IntegerType(), required=False, initial_default=1, write_default=2))), + ) + fresh = assign_fresh_schema_ids(schema) + assert fresh.find_field("c").field_id == 1 + assert (fresh.find_field("c").initial_default, fresh.find_field("c").write_default) == ("init", "wd") + assert (fresh.find_field("s.n").initial_default, fresh.find_field("s.n").write_default) == (1, 2) diff --git a/tests/test_transforms.py b/tests/test_transforms.py index 998d8a224d..fb2e48a290 100644 --- a/tests/test_transforms.py +++ b/tests/test_transforms.py @@ -17,7 +17,7 @@ # under the License. # pylint: disable=eval-used,protected-access,redefined-outer-name from collections.abc import Callable -from datetime import date, datetime +from datetime import date, datetime, timezone from decimal import Decimal from typing import Annotated, Any from uuid import UUID @@ -72,10 +72,11 @@ DateLiteral, DecimalLiteral, TimestampLiteral, + TimestampNanoLiteral, literal, ) from pyiceberg.partitioning import _to_partition_representation -from pyiceberg.schema import Accessor +from pyiceberg.schema import Accessor, Schema from pyiceberg.transforms import ( BucketTransform, DayTransform, @@ -417,6 +418,8 @@ def test_satisfies_order_of_method(transform: TimeTransform[Any], other_transfor (TimeType(), 36775038194, "10:12:55.038194"), (TimestamptzType(), 1512151975038194, "2017-12-01T18:12:55.038194+00:00"), (TimestampType(), 1512151975038194, "2017-12-01T18:12:55.038194"), + (TimestampNanoType(), 1740342104375612001, "2025-02-23T20:21:44.375612001"), + (TimestamptzNanoType(), 1740342104375612001, "2025-02-23T20:21:44.375612001+00:00"), (LongType(), -1234567890000, "-1234567890000"), (StringType(), "a/b/c=d", "a/b/c=d"), (DecimalType(9, 2), Decimal("-1.50"), "-1.50"), @@ -575,6 +578,50 @@ def test_unknown_transform_str() -> None: assert str(UnknownTransform("unknown")) == "unknown" +def test_unknown_transform_str_preserves_original_name() -> None: + # serializing metadata with an unknown transform must not rewrite its name + assert str(UnknownTransform("zorder")) == "zorder" + assert str(UnknownTransform("bucketv2[4]")) == "bucketv2[4]" + + +def test_parse_transform_unknown_with_known_prefix() -> None: + # unknown transforms that share a prefix with known ones must not fail parsing + from pyiceberg.transforms import parse_transform + + for name in ("bucketv2[4]", "truncatev2[8]", "bucket", "truncate"): + transform = parse_transform(name) + assert isinstance(transform, UnknownTransform), name + assert str(transform) == name + + +def test_parse_transform_rejects_malformed_arguments() -> None: + # a known transform with an unparsable argument is an error, not an unknown transform + from pyiceberg.exceptions import ValidationError + from pyiceberg.transforms import parse_transform + + for name in ("bucket[abc]", "truncate[]", "bucket[]"): + with pytest.raises(ValidationError): + parse_transform(name) + + +def test_parse_transform_rejects_trailing_text_after_brackets() -> None: + # a name that merely starts with a known transform is a different transform, so it has to be + # preserved rather than silently rewritten to the known one + from pyiceberg.transforms import parse_transform + + for name in ("bucket[4]garbage", "truncate[8]v2"): + transform = parse_transform(name) + assert isinstance(transform, UnknownTransform), name + assert str(transform) == name + + +def test_unknown_transforms_compare_by_name() -> None: + # every unknown transform shares the same root value, so equality has to use the preserved name + assert UnknownTransform("zorder(x,y)") == UnknownTransform("zorder(x,y)") + assert UnknownTransform("zorder(x,y)") != UnknownTransform("bucketv2[4]") + assert len({UnknownTransform("zorder(x,y)"), UnknownTransform("bucketv2[4]")}) == 2 + + def test_unknown_transform_repr() -> None: assert repr(UnknownTransform("unknown")) == "UnknownTransform(transform='unknown')" @@ -1734,3 +1781,52 @@ def test_calling_pyarrow_transform_without_pyiceberg_core_installed_correctly_ra with pytest.raises(NotInstalledError): transform.pyarrow_transform(StringType()) + + +@pytest.mark.parametrize( + "transform, expected", + [ + (YearTransform(), 47), + (MonthTransform(), 574), + (DayTransform(), 17486), + (HourTransform(), 419686), + (BucketTransform(100), 7), + ], +) +@pytest.mark.parametrize("source_type", [TimestampNanoType(), TimestamptzNanoType()]) +def test_nano_transforms_accept_datetimes(transform: Transform[Any, int], expected: int, source_type: PrimitiveType) -> None: + value = datetime(2017, 11, 16, 22, 31, 8) + if isinstance(source_type, TimestamptzNanoType): + value = value.replace(tzinfo=timezone.utc) + nanos = 1510871468000000000 + assert transform.transform(source_type)(value) == expected + assert transform.transform(source_type)(nanos) == expected + + +NS_SCHEMA = Schema(NestedField(1, "ts_ns", TimestampNanoType(), required=False)) + + +@pytest.mark.parametrize( + "transform, expected", + [ + (YearTransform(), GreaterThanOrEqual(term="part", literal=54)), + (MonthTransform(), GreaterThanOrEqual(term="part", literal=649)), + (DayTransform(), GreaterThanOrEqual(term="part", literal=19754)), + (HourTransform(), GreaterThanOrEqual(term="part", literal=474096)), + (IdentityTransform(), GreaterThanOrEqual(term="part", literal=TimestampNanoLiteral(1706745600000000001))), + ], +) +def test_projection_nano_timestamp(transform: Transform[Any, Any], expected: UnboundPredicate) -> None: + bound = GreaterThanOrEqual("ts_ns", "2024-02-01T00:00:00.000000001").bind(NS_SCHEMA) + assert transform.project("part", bound) == expected + + +def test_projection_nano_timestamp_bucket_and_strict() -> None: + bound_eq = EqualTo("ts_ns", "2024-02-01T00:00:00.000000001").bind(NS_SCHEMA) + # Bucketing uses microseconds, so the sub-microsecond digits do not change the bucket + expected_bucket = BucketTransform(4).transform(TimestampType())(1706745600000000) + assert BucketTransform(4).project("part", bound_eq) == EqualTo(term="part", literal=expected_bucket) + + bound_lt = LessThan("ts_ns", "2024-02-01T00:00:00").bind(NS_SCHEMA) + assert DayTransform().project("part", bound_lt) == LessThanOrEqual(term="part", literal=19753) + assert DayTransform().strict_project("part", bound_lt) == LessThan(term="part", literal=19754) diff --git a/tests/test_types.py b/tests/test_types.py index eb8ae2ea52..3ae7f75938 100644 --- a/tests/test_types.py +++ b/tests/test_types.py @@ -46,7 +46,11 @@ TimestampType, TimestamptzType, TimeType, + UnknownType, UUIDType, + VariantType, + _parse_geography_type, + _parse_geometry_type, strtobool, transform_dict_value_to_str, ) @@ -64,6 +68,7 @@ (10, StringType), (11, UUIDType), (12, BinaryType), + (13, VariantType), ] primitive_types = { @@ -79,6 +84,7 @@ "string": StringType, "uuid": UUIDType, "binary": BinaryType, + "variant": VariantType, } @@ -248,6 +254,28 @@ def test_nested_field() -> None: _ = (NestedField(1, "field", StringType(), required=True, write_default=(1, "a", True)),) # type: ignore +def test_nested_field_pickle_keeps_defaults() -> None: + field_var = NestedField(1, "color", StringType(), required=True, initial_default="blue", write_default="green") + unpickled = pickle.loads(pickle.dumps(field_var)) + assert unpickled == field_var + assert (unpickled.initial_default, unpickled.write_default) == ("blue", "green") + assert field_var.__getnewargs__() == (1, "color", StringType(), True, None, "blue", "green") + + +def test_nested_field_unknown_must_be_optional() -> None: + assert NestedField(1, "u", UnknownType(), required=False).optional + with pytest.raises(ValueError, match="Columns of type unknown must be optional"): + NestedField(1, "u", UnknownType(), required=True) + + +@pytest.mark.parametrize("field_type", [UnknownType(), GeometryType(), GeographyType()]) +def test_nested_field_null_only_types_reject_defaults(field_type: PrimitiveType) -> None: + with pytest.raises(ValueError, match="must default to null"): + NestedField(1, "f", field_type, required=False, initial_default=b"\x01") + with pytest.raises(ValueError, match="must default to null"): + NestedField(1, "f", field_type, required=False, write_default=b"\x01") + + def test_nested_field_complex_type_as_str_unsupported() -> None: unsupported_types = ["list", "map", "struct"] for type_str in unsupported_types: @@ -752,7 +780,7 @@ def test_geometry_type_custom_crs() -> None: """Test GeometryType with custom CRS.""" type_var = GeometryType("EPSG:4326") assert type_var.crs == "EPSG:4326" - assert str(type_var) == "geometry('EPSG:4326')" + assert str(type_var) == "geometry(EPSG:4326)" assert repr(type_var) == "GeometryType(crs='EPSG:4326')" @@ -774,20 +802,21 @@ def test_geometry_type_pickle() -> None: def test_geometry_type_serialization() -> None: """Test GeometryType JSON serialization.""" assert GeometryType().model_dump_json() == '"geometry"' - assert GeometryType("EPSG:4326").model_dump_json() == "\"geometry('EPSG:4326')\"" + assert GeometryType("EPSG:4326").model_dump_json() == '"geometry(EPSG:4326)"' def test_geometry_type_deserialization() -> None: """Test GeometryType JSON deserialization.""" assert GeometryType.model_validate_json('"geometry"') == GeometryType() + assert GeometryType.model_validate_json('"geometry(EPSG:4326)"') == GeometryType("EPSG:4326") assert GeometryType.model_validate_json("\"geometry('EPSG:4326')\"") == GeometryType("EPSG:4326") def test_geometry_type_deserialization_failure() -> None: """Test GeometryType deserialization with invalid input.""" with pytest.raises(ValidationError) as exc_info: - GeometryType.model_validate_json('"geometry(invalid)"') - assert "Could not parse geometry(invalid) into a GeometryType" in str(exc_info.value) + GeometryType.model_validate_json('"geometry(invalid)x"') + assert "Could not parse geometry(invalid)x into a GeometryType" in str(exc_info.value) def test_geometry_type_singleton() -> None: @@ -816,61 +845,62 @@ def test_geography_type_custom_crs() -> None: type_var = GeographyType("EPSG:4326") assert type_var.crs == "EPSG:4326" assert type_var.algorithm == DEFAULT_GEOGRAPHY_ALGORITHM - assert str(type_var) == "geography('EPSG:4326')" + assert str(type_var) == "geography(EPSG:4326)" assert repr(type_var) == "GeographyType(crs='EPSG:4326')" def test_geography_type_custom_crs_and_algorithm() -> None: """Test GeographyType with custom CRS and algorithm.""" - type_var = GeographyType("EPSG:4326", "planar") + type_var = GeographyType("EPSG:4326", "vincenty") assert type_var.crs == "EPSG:4326" - assert type_var.algorithm == "planar" - assert str(type_var) == "geography('EPSG:4326', 'planar')" - assert repr(type_var) == "GeographyType(crs='EPSG:4326', algorithm='planar')" + assert type_var.algorithm == "vincenty" + assert str(type_var) == "geography(EPSG:4326, vincenty)" + assert repr(type_var) == "GeographyType(crs='EPSG:4326', algorithm='vincenty')" def test_geography_type_equality() -> None: """Test GeographyType equality and hashing.""" assert GeographyType() == GeographyType() assert GeographyType("EPSG:4326") == GeographyType("EPSG:4326") - assert GeographyType("EPSG:4326", "planar") == GeographyType("EPSG:4326", "planar") + assert GeographyType("EPSG:4326", "vincenty") == GeographyType("EPSG:4326", "vincenty") assert GeographyType() != GeographyType("EPSG:4326") - assert GeographyType("EPSG:4326") != GeographyType("EPSG:4326", "planar") + assert GeographyType("EPSG:4326") != GeographyType("EPSG:4326", "vincenty") assert hash(GeographyType()) == hash(GeographyType()) def test_geography_type_pickle() -> None: """Test GeographyType pickle round-trip.""" assert GeographyType() == pickle.loads(pickle.dumps(GeographyType())) - assert GeographyType("EPSG:4326", "planar") == pickle.loads(pickle.dumps(GeographyType("EPSG:4326", "planar"))) + assert GeographyType("EPSG:4326", "vincenty") == pickle.loads(pickle.dumps(GeographyType("EPSG:4326", "vincenty"))) def test_geography_type_serialization() -> None: """Test GeographyType JSON serialization.""" assert GeographyType().model_dump_json() == '"geography"' - assert GeographyType("EPSG:4326").model_dump_json() == "\"geography('EPSG:4326')\"" - assert GeographyType("EPSG:4326", "planar").model_dump_json() == "\"geography('EPSG:4326', 'planar')\"" + assert GeographyType("EPSG:4326").model_dump_json() == '"geography(EPSG:4326)"' + assert GeographyType("EPSG:4326", "vincenty").model_dump_json() == '"geography(EPSG:4326, vincenty)"' def test_geography_type_deserialization() -> None: """Test GeographyType JSON deserialization.""" assert GeographyType.model_validate_json('"geography"') == GeographyType() - assert GeographyType.model_validate_json("\"geography('EPSG:4326')\"") == GeographyType("EPSG:4326") - assert GeographyType.model_validate_json("\"geography('EPSG:4326', 'planar')\"") == GeographyType("EPSG:4326", "planar") + assert GeographyType.model_validate_json('"geography(EPSG:4326)"') == GeographyType("EPSG:4326") + assert GeographyType.model_validate_json('"geography(EPSG:4326, vincenty)"') == GeographyType("EPSG:4326", "vincenty") + assert GeographyType.model_validate_json("\"geography('EPSG:4326', 'vincenty')\"") == GeographyType("EPSG:4326", "vincenty") def test_geography_type_deserialization_failure() -> None: """Test GeographyType deserialization with invalid input.""" with pytest.raises(ValidationError) as exc_info: - GeographyType.model_validate_json('"geography(invalid)"') - assert "Could not parse geography(invalid) into a GeographyType" in str(exc_info.value) + GeographyType.model_validate_json('"geography(a, b, c)"') + assert "Could not parse geography(a, b, c) into a GeographyType" in str(exc_info.value) def test_geography_type_singleton() -> None: """Test that GeographyType uses singleton pattern for same parameters.""" assert id(GeographyType()) == id(GeographyType()) assert id(GeographyType("EPSG:4326")) == id(GeographyType("EPSG:4326")) - assert id(GeographyType("EPSG:4326", "planar")) == id(GeographyType("EPSG:4326", "planar")) + assert id(GeographyType("EPSG:4326", "vincenty")) == id(GeographyType("EPSG:4326", "vincenty")) assert id(GeographyType()) != id(GeographyType("EPSG:4326")) @@ -886,10 +916,10 @@ def test_nested_field_with_geometry() -> None: def test_nested_field_with_geography() -> None: """Test NestedField with GeographyType.""" - field = NestedField(1, "location", GeographyType("EPSG:4326", "planar"), required=True) + field = NestedField(1, "location", GeographyType("EPSG:4326", "vincenty"), required=True) assert isinstance(field.field_type, GeographyType) assert field.field_type.crs == "EPSG:4326" - assert field.field_type.algorithm == "planar" + assert field.field_type.algorithm == "vincenty" def test_nested_field_geometry_as_string() -> None: @@ -909,17 +939,66 @@ def test_nested_field_geography_as_string() -> None: def test_nested_field_geometry_with_params_as_string() -> None: """Test NestedField with parameterized geometry type as string.""" - field = NestedField(1, "location", "geometry('EPSG:4326')", required=True) + field = NestedField(1, "location", "geometry(EPSG:4326)", required=True) assert isinstance(field.field_type, GeometryType) assert field.field_type.crs == "EPSG:4326" def test_nested_field_geography_with_params_as_string() -> None: """Test NestedField with parameterized geography type as string.""" - field = NestedField(1, "location", "geography('EPSG:4326', 'planar')", required=True) + field = NestedField(1, "location", "geography(EPSG:4326, vincenty)", required=True) assert isinstance(field.field_type, GeographyType) assert field.field_type.crs == "EPSG:4326" - assert field.field_type.algorithm == "planar" + assert field.field_type.algorithm == "vincenty" + + +@pytest.mark.parametrize( + "type_string, expected, serialized", + [ + # Strings emitted by Java's Types.GeometryType / Types.GeographyType + ("geometry", GeometryType(), "geometry"), + ("geometry(srid:3857)", GeometryType("srid:3857"), "geometry(srid:3857)"), + ("geometry(EPSG:4326)", GeometryType("EPSG:4326"), "geometry(EPSG:4326)"), + ("geography", GeographyType(), "geography"), + ("geography(srid:4326)", GeographyType("srid:4326"), "geography(srid:4326)"), + ("geography(srid:4326, vincenty)", GeographyType("srid:4326", "vincenty"), "geography(srid:4326, vincenty)"), + ("geography(OGC:CRS84, karney)", GeographyType("OGC:CRS84", "karney"), "geography(OGC:CRS84, karney)"), + # Case-insensitive, whitespace-tolerant, like Java + ("GEOMETRY", GeometryType(), "geometry"), + ("Geometry ( srid:3857 )", GeometryType("srid:3857"), "geometry(srid:3857)"), + ("GEOGRAPHY(srid:4326, VINCENTY)", GeographyType("srid:4326", "vincenty"), "geography(srid:4326, vincenty)"), + ("geography(srid:4326,thomas)", GeographyType("srid:4326", "thomas"), "geography(srid:4326, thomas)"), + # Default values are collapsed to the bare form + ("geometry(OGC:CRS84)", GeometryType(), "geometry"), + ("geography(OGC:CRS84, spherical)", GeographyType(), "geography"), + # Quoted values written by older PyIceberg versions + ("geometry('srid:3857')", GeometryType("srid:3857"), "geometry(srid:3857)"), + ('geometry("srid:3857")', GeometryType("srid:3857"), "geometry(srid:3857)"), + ("geography('srid:4326')", GeographyType("srid:4326"), "geography(srid:4326)"), + ("geography('srid:4326', 'vincenty')", GeographyType("srid:4326", "vincenty"), "geography(srid:4326, vincenty)"), + ('geography("srid:4326", "andoyer")', GeographyType("srid:4326", "andoyer"), "geography(srid:4326, andoyer)"), + ], +) +def test_geo_type_string_round_trip(type_string: str, expected: IcebergType, serialized: str) -> None: + """Test that geo type strings are parsed like Java and serialized in the unquoted spec form.""" + field = NestedField(1, "location", type_string, required=False) + assert field.field_type == expected + assert str(field.field_type) == serialized + assert field.field_type.model_dump_json() == f'"{serialized}"' + assert NestedField.model_validate_json(field.model_dump_json()) == field + + +def test_geography_type_string_invalid_algorithm() -> None: + """Test that an unknown geography edge algorithm is rejected.""" + with pytest.raises(ValidationError, match="Invalid geography edge algorithm: planar"): + NestedField(1, "location", "geography(srid:4326, planar)", required=False) + + +def test_geography_type_dict_input() -> None: + """Test that the dict form of geo types is still accepted.""" + assert _parse_geometry_type({"crs": "srid:3857"}) == "srid:3857" + assert _parse_geography_type({"crs": "srid:4326", "algorithm": "Vincenty"}) == ("srid:4326", "vincenty") + assert _parse_geography_type({}) == (DEFAULT_GEOGRAPHY_CRS, DEFAULT_GEOGRAPHY_ALGORITHM) def test_decimal_precision_validation() -> None: @@ -947,3 +1026,23 @@ def test_decimal_scale_validation() -> None: with pytest.raises(ValidationError, match="Decimal scale must be between 0 and the precision"): DecimalType(5, 10) + + +def test_variant_type() -> None: + assert VariantType() is VariantType() + assert str(VariantType()) == "variant" + assert repr(VariantType()) == "VariantType()" + assert VariantType().is_primitive + assert VariantType().minimum_format_version() == 3 + field = NestedField.model_validate_json('{"id": 1, "name": "v", "type": "variant", "required": false}') + assert field.field_type == VariantType() + assert '"type":"variant"' in field.model_dump_json() + + +def test_variant_type_requires_format_version_3() -> None: + from pyiceberg.schema import Schema + + schema = Schema(NestedField(1, "v", VariantType(), required=False)) + schema.check_format_version_compatibility(3) + with pytest.raises(ValueError, match="variant is only supported in 3 or higher"): + schema.check_format_version_compatibility(2) diff --git a/tests/test_variant.py b/tests/test_variant.py new file mode 100644 index 0000000000..fd88e3931f --- /dev/null +++ b/tests/test_variant.py @@ -0,0 +1,181 @@ +# Licensed to the Apache Software Foundation (ASF) under one +# or more contributor license agreements. See the NOTICE file +# distributed with this work for additional information +# regarding copyright ownership. The ASF licenses this file +# to you under the Apache License, Version 2.0 (the +# "License"); you may not use this file except in compliance +# with the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, +# software distributed under the License is distributed on an +# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +# KIND, either express or implied. See the License for the +# specific language governing permissions and limitations +# under the License. +import json +import struct +import uuid +from datetime import date, datetime, time, timezone +from decimal import Decimal +from typing import Any + +import pytest + +from pyiceberg.variant import to_json, to_python + +EMPTY_METADATA = b"\x01\x00\x00" + + +def _metadata(*keys: str) -> bytes: + """Encode a metadata dictionary with 1-byte offsets.""" + encoded = [key.encode() for key in keys] + offsets = [0] + for key in encoded: + offsets.append(offsets[-1] + len(key)) + return bytes([0x01, len(keys), *offsets]) + b"".join(encoded) + + +def _primitive(type_id: int, payload: bytes = b"") -> bytes: + return bytes([type_id << 2]) + payload + + +def _short_string(value: str) -> bytes: + data = value.encode() + return bytes([(len(data) << 2) | 1]) + data + + +def _array(*elements: bytes) -> bytes: + offsets = [0] + for element in elements: + offsets.append(offsets[-1] + len(element)) + return bytes([0x03, len(elements), *offsets]) + b"".join(elements) + + +def _object(fields: list[tuple[int, bytes]]) -> bytes: + offsets = [] + position = 0 + for _, element in fields: + offsets.append(position) + position += len(element) + offsets.append(position) + return bytes([0x02, len(fields), *[field_id for field_id, _ in fields], *offsets]) + b"".join( + element for _, element in fields + ) + + +# Written by Spark 4 / Iceberg 1.11 for parse_json('...') +SPARK_SAMPLES = [ + (b"\x01\x01\x00\x01a", b"\x02\x01\x00\x00\x02\x0c\x01", {"a": 1}, '{"a":1}'), + ( + EMPTY_METADATA, + b"\x03\x04\x00\x02\x06\x07\r\x0c\x01\rtwo\x00 \x01#\x00\x00\x00", + [1, "two", None, Decimal("3.5")], + '[1,"two",null,3.5]', + ), + (EMPTY_METADATA, b"\x35just a string", "just a string", '"just a string"'), + (EMPTY_METADATA, b"\x00", None, "null"), +] + + +@pytest.mark.parametrize("metadata,value,expected_python,expected_json", SPARK_SAMPLES) +def test_decode_spark_written_variants(metadata: bytes, value: bytes, expected_python: Any, expected_json: str) -> None: + assert to_python(metadata, value) == expected_python + assert to_json(metadata, value) == expected_json + + +@pytest.mark.parametrize( + "value,expected", + [ + (_primitive(0), None), + (_primitive(1), True), + (_primitive(2), False), + (_primitive(3, struct.pack(" None: + assert to_python(EMPTY_METADATA, value) == expected + + +def test_to_json_primitives() -> None: + assert to_json(EMPTY_METADATA, _primitive(1)) == "true" + assert to_json(EMPTY_METADATA, _primitive(10, bytes([2]) + (10**30).to_bytes(16, "little", signed=True))) == ( + "10000000000000000000000000000.00" + ) + assert to_json(EMPTY_METADATA, _primitive(11, struct.pack(" None: + metadata = _metadata("id", "tags", "nested") + value = _object( + [ + (0, _primitive(3, b"\x07")), + (1, _array(_short_string("x"), _primitive(0), _array())), + (2, _object([(0, _short_string("inner"))])), + ] + ) + + assert to_python(metadata, value) == {"id": 7, "tags": ["x", None, []], "nested": {"id": "inner"}} + assert json.loads(to_json(metadata, value)) == {"id": 7, "tags": ["x", None, []], "nested": {"id": "inner"}} + + +def test_object_with_out_of_order_offsets() -> None: + # field ids are sorted by key, but values may be laid out in any order + metadata = _metadata("a", "b") + values = _short_string("B") + _short_string("A") + value = bytes([0x02, 2, 0, 1, 2, 0, 4]) + values + assert to_python(metadata, value) == {"a": "A", "b": "B"} + + +def test_large_array_and_wide_offsets() -> None: + elements = [_primitive(3, bytes([index])) for index in range(3)] + # is_large (4-byte element count) with 2-byte offsets + header = (0b101 << 2) | 0x03 + value = ( + bytes([header]) + struct.pack(" None: + with pytest.raises(ValueError, match="Unsupported variant metadata version"): + to_python(b"\x02\x00\x00", _primitive(0)) + with pytest.raises(ValueError, match="unexpected end of buffer"): + to_python(EMPTY_METADATA, b"\x35short") + with pytest.raises(ValueError, match="Unsupported variant primitive type id"): + to_python(EMPTY_METADATA, _primitive(63)) + with pytest.raises(ValueError, match="not in the metadata dictionary"): + to_python(EMPTY_METADATA, _object([(0, _primitive(0))])) diff --git a/tests/utils/test_datetime.py b/tests/utils/test_datetime.py index 54fd3eefcc..c125e7ba15 100644 --- a/tests/utils/test_datetime.py +++ b/tests/utils/test_datetime.py @@ -30,6 +30,8 @@ time_to_nanos, timestamp_to_nanos, timestamptz_to_nanos, + to_human_timestamp_ns, + to_human_timestamptz_ns, ) timezones = [ @@ -169,3 +171,35 @@ def test_nanos_to_micros(nanos: int, micros: int) -> None: ) def test_nanos_to_hours(nanos: int, hours: int) -> None: assert hours == nanos_to_hours(nanos) + + +@pytest.mark.parametrize( + "nanos, human_str", + [ + (0, "1970-01-01T00:00:00.000000000"), + (1, "1970-01-01T00:00:00.000000001"), + (999, "1970-01-01T00:00:00.000000999"), + (1000, "1970-01-01T00:00:00.000001000"), + (1740342104375612001, "2025-02-23T20:21:44.375612001"), + (1510871468000001001, "2017-11-16T22:31:08.000001001"), + (1740342104000000001, "2025-02-23T20:21:44.000000001"), + ], +) +def test_to_human_timestamp_ns(nanos: int, human_str: str) -> None: + assert human_str == to_human_timestamp_ns(nanos) + + +@pytest.mark.parametrize( + "nanos, human_str", + [ + (0, "1970-01-01T00:00:00.000000000+00:00"), + (1, "1970-01-01T00:00:00.000000001+00:00"), + (999, "1970-01-01T00:00:00.000000999+00:00"), + (1000, "1970-01-01T00:00:00.000001000+00:00"), + (1740342104375612001, "2025-02-23T20:21:44.375612001+00:00"), + (1510871468000001001, "2017-11-16T22:31:08.000001001+00:00"), + (1740342104000000001, "2025-02-23T20:21:44.000000001+00:00"), + ], +) +def test_to_human_timestamptz_ns(nanos: int, human_str: str) -> None: + assert human_str == to_human_timestamptz_ns(nanos) diff --git a/tests/utils/test_geo.py b/tests/utils/test_geo.py new file mode 100644 index 0000000000..6b26dbf4e4 --- /dev/null +++ b/tests/utils/test_geo.py @@ -0,0 +1,59 @@ +# Licensed to the Apache Software Foundation (ASF) under one +# or more contributor license agreements. See the NOTICE file +# distributed with this work for additional information +# regarding copyright ownership. The ASF licenses this file +# to you under the Apache License, Version 2.0 (the +# "License"); you may not use this file except in compliance +# with the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, +# software distributed under the License is distributed on an +# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +# KIND, either express or implied. See the License for the +# specific language governing permissions and limitations +# under the License. +import math +import struct + +import pytest + +from pyiceberg.utils.geo import GeospatialBound, geo_bounds_from_bbox + + +@pytest.mark.parametrize( + "bound, expected", + [ + (GeospatialBound(1.0, 2.0), struct.pack("<2d", 1.0, 2.0)), + (GeospatialBound(1.0, 2.0, 3.0), struct.pack("<3d", 1.0, 2.0, 3.0)), + (GeospatialBound(1.0, 2.0, 3.0, 4.0), struct.pack("<4d", 1.0, 2.0, 3.0, 4.0)), + ], +) +def test_geospatial_bound_round_trip(bound: GeospatialBound, expected: bytes) -> None: + assert bound.to_bytes() == expected + assert GeospatialBound.from_bytes(expected) == bound + + +def test_geospatial_bound_xym_uses_nan_z() -> None: + serialized = GeospatialBound(1.0, 2.0, m=4.0).to_bytes() + x, y, z, m = struct.unpack("<4d", serialized) + assert (x, y, m) == (1.0, 2.0, 4.0) + assert math.isnan(z) + assert GeospatialBound.from_bytes(serialized) == GeospatialBound(1.0, 2.0, m=4.0) + + +def test_geospatial_bound_invalid_length() -> None: + with pytest.raises(ValueError, match="Invalid geospatial bound of 21 bytes"): + GeospatialBound.from_bytes(b"\x00" * 21) + + +def test_geo_bounds_from_bbox() -> None: + lower, upper = geo_bounds_from_bbox(-3.0, 2.0, 1.0, 5.0) + assert GeospatialBound.from_bytes(lower) == GeospatialBound(-3.0, 2.0) + assert GeospatialBound.from_bytes(upper) == GeospatialBound(1.0, 5.0) + + lower, upper = geo_bounds_from_bbox(-3.0, 2.0, 1.0, 5.0, zmin=0.0, zmax=9.0, mmin=None, mmax=7.0) + # M is only included when both its minimum and maximum are known + assert GeospatialBound.from_bytes(lower) == GeospatialBound(-3.0, 2.0, 0.0) + assert GeospatialBound.from_bytes(upper) == GeospatialBound(1.0, 5.0, 9.0) diff --git a/tests/utils/test_manifest.py b/tests/utils/test_manifest.py index 331146346e..454e7b950a 100644 --- a/tests/utils/test_manifest.py +++ b/tests/utils/test_manifest.py @@ -26,6 +26,7 @@ import pyiceberg.manifest as manifest_module from pyiceberg.avro.codecs import AvroCompressionCodec from pyiceberg.avro.file import AvroFile, AvroOutputFile +from pyiceberg.exceptions import ValidationError from pyiceberg.io import load_file_io from pyiceberg.io.pyarrow import PyArrowFileIO from pyiceberg.manifest import ( @@ -228,6 +229,67 @@ def test_fetch_manifest_entry_with_filter(generated_manifest_entry_file: str) -> assert len(no_match) == 0 +def test_fetch_manifest_entry_inherits_first_row_id(tmp_path: Path) -> None: + """Data files inherit first row IDs from data manifests in file order, and have them cleared otherwise.""" + io = PyArrowFileIO() + manifest_path = str(tmp_path / "manifest.avro") + + def entry(status: ManifestEntryStatus, record_count: int, first_row_id: int | None = None) -> ManifestEntry: + return ManifestEntry.from_args( + _table_format_version=3, + status=status, + snapshot_id=25, + sequence_number=1, + file_sequence_number=1, + data_file=DataFile.from_args( + _table_format_version=3, + content=DataFileContent.DATA, + file_path=f"s3://bucket/data-{record_count}.parquet", + file_format=FileFormat.PARQUET, + partition=Record(), + record_count=record_count, + file_size_in_bytes=1024, + first_row_id=first_row_id, + ), + ) + + with AvroOutputFile[ManifestEntry]( + output_file=io.new_output(manifest_path), + file_schema=MANIFEST_ENTRY_SCHEMAS[3], + record_schema=MANIFEST_ENTRY_SCHEMAS[3], + schema_name="manifest_entry", + metadata={"format-version": "3"}, + ) as writer: + writer.write_block( + [ + entry(ManifestEntryStatus.ADDED, 10), + entry(ManifestEntryStatus.EXISTING, 5, first_row_id=500), + entry(ManifestEntryStatus.DELETED, 7), + entry(ManifestEntryStatus.ADDED, 3), + ] + ) + + def first_row_ids( + manifest_first_row_id: int | None, discard_deleted: bool, content: ManifestContent = ManifestContent.DATA + ) -> list[int | None]: + manifest = ManifestFile.from_args( + manifest_path=manifest_path, + manifest_length=0, + partition_spec_id=0, + content=content, + added_snapshot_id=25, + sequence_number=1, + min_sequence_number=1, + first_row_id=manifest_first_row_id, + ) + return [e.data_file.first_row_id for e in manifest.fetch_manifest_entry(io, discard_deleted=discard_deleted)] + + assert first_row_ids(1000, discard_deleted=False) == [1000, 500, None, 1010] + assert first_row_ids(1000, discard_deleted=True) == [1000, 500, 1010] + assert first_row_ids(None, discard_deleted=False) == [None, None, None, None] + assert first_row_ids(1000, discard_deleted=False, content=ManifestContent.DELETES) == [None, None, None, None] + + def test_read_manifest_entry_v3_fields(tmp_path: Path) -> None: io = PyArrowFileIO() @@ -257,6 +319,7 @@ def write_and_read(file_name: str, data_file: DataFile) -> DataFile: added_snapshot_id=25, sequence_number=1, min_sequence_number=1, + first_row_id=0, ) return manifest.fetch_manifest_entry(io)[0].data_file @@ -554,7 +617,7 @@ def test_write_empty_manifest() -> None: pass -@pytest.mark.parametrize("format_version", [1, 2]) +@pytest.mark.parametrize("format_version", [1, 2, 3]) @pytest.mark.parametrize("compression", ["null", "deflate", "zstd"]) def test_write_manifest( generated_manifest_file_file_v1: str, @@ -726,7 +789,7 @@ def test_write_manifest( assert data_file.sort_order_id == 0 -@pytest.mark.parametrize("format_version", [1, 2]) +@pytest.mark.parametrize("format_version", [1, 2, 3]) @pytest.mark.parametrize("parent_snapshot_id", [19, None]) @pytest.mark.parametrize("compression", ["null", "deflate"]) def test_write_manifest_list( @@ -758,6 +821,7 @@ def test_write_manifest_list( parent_snapshot_id=parent_snapshot_id, sequence_number=0, avro_compression=compression, + first_row_id=0 if format_version >= 3 else None, ) as writer: writer.add_manifests(demo_manifest_list) new_manifest_list = list(read_manifest_list(io.new_input(path))) @@ -767,8 +831,10 @@ def test_write_manifest_list( else: expected_metadata = {"snapshot-id": "25", "parent-snapshot-id": "null", "format-version": str(format_version)} - if format_version == 2: + if format_version >= 2: expected_metadata["sequence-number"] = "0" + if format_version >= 3: + expected_metadata["first-row-id"] = "0" _verify_metadata_with_fastavro(path, expected_metadata) manifest_file = new_manifest_list[0] @@ -1091,7 +1157,7 @@ def test_manifest_cache_efficiency_with_many_overlapping_lists() -> None: assert ref is references[0], f"All references to manifest {i} should be the same object instance" -@pytest.mark.parametrize("format_version", [1, 2]) +@pytest.mark.parametrize("format_version", [1, 2, 3]) def test_manifest_writer_tell(format_version: TableVersion) -> None: io = load_file_io() test_schema = Schema(NestedField(1, "foo", IntegerType(), False)) @@ -1126,7 +1192,7 @@ def test_manifest_writer_tell(format_version: TableVersion) -> None: assert after_entry_bytes > initial_bytes, "Bytes should increase after adding entry" -@pytest.mark.parametrize("format_version", [1, 2]) +@pytest.mark.parametrize("format_version", [1, 2, 3]) def test_write_manifest_min_sequence_number_zero(format_version: TableVersion) -> None: # A data sequence number of 0 is a legitimate min for a live file (e.g. files from a # v1 table or the initial commit of a v2 table). It must be preserved in the manifest, @@ -1272,6 +1338,19 @@ def test_clear_manifest_cache() -> None: assert len(manifest_module._manifest_cache) == 0, "Cache should be empty after clear" +def test_manifest_cache_refreshes_assigned_first_row_id() -> None: + """A manifest written before a V3 upgrade gets its first row id from a later manifest list.""" + unassigned = ManifestFile.from_args(**_manifest_args("/m1.avro")) + assigned = ManifestFile.from_args(first_row_id=10, **_manifest_args("/m1.avro")) + + assert manifest_module._manifest_cache.get_or_cache(unassigned) is unassigned + assert manifest_module._manifest_cache.get_or_cache(assigned) is assigned + assert ( + manifest_module._manifest_cache.get_or_cache(ManifestFile.from_args(first_row_id=10, **_manifest_args("/m1.avro"))) + is assigned + ) + + def test_manifest_cache_can_be_disabled_with_size_zero(monkeypatch: pytest.MonkeyPatch) -> None: """Test that manifest-cache-size=0 disables caching.""" monkeypatch.setenv("PYICEBERG_MANIFEST_CACHE_SIZE", "0") @@ -1376,3 +1455,192 @@ def test_negative_manifest_cache_size_raises_value_error(monkeypatch: pytest.Mon finally: monkeypatch.delenv("PYICEBERG_MANIFEST_CACHE_SIZE", raising=False) importlib.reload(manifest_module) + + +@pytest.mark.parametrize("content", [ManifestContent.DATA, ManifestContent.DELETES]) +def test_write_manifest_v3(tmp_path: Path, content: ManifestContent) -> None: + io = load_file_io() + test_schema = Schema(NestedField(1, "foo", IntegerType(), False)) + file_content = DataFileContent.DATA if content == ManifestContent.DATA else DataFileContent.POSITION_DELETES + + v3_data_file = DataFile.from_args( + _table_format_version=3, + content=file_content, + file_path="/data/file-v3.puffin", + file_format=FileFormat.PUFFIN, + partition=Record(), + record_count=100, + file_size_in_bytes=1024, + key_metadata=b"\x01", + first_row_id=1000, + referenced_data_file="/data/referenced.parquet", + content_offset=4, + content_size_in_bytes=42, + ) + # Records bound to an older layout are rebound to the V3 layout + v2_data_file = DataFile.from_args( + _table_format_version=2, + content=file_content, + file_path="/data/file-v2.parquet", + file_format=FileFormat.PARQUET, + partition=Record(), + record_count=50, + file_size_in_bytes=512, + ) + v1_data_file = DataFile.from_args( + _table_format_version=1, + file_path="/data/file-v1.parquet", + file_format=FileFormat.PARQUET, + partition=Record(), + record_count=10, + file_size_in_bytes=128, + ) + + path = str(tmp_path / "manifest-v3.avro") + with write_manifest( + format_version=3, + spec=UNPARTITIONED_PARTITION_SPEC, + schema=test_schema, + output_file=io.new_output(path), + snapshot_id=25, + avro_compression="null", + content=content, + ) as writer: + for data_file in (v3_data_file, v2_data_file, v1_data_file): + writer.add(ManifestEntry.from_args(status=ManifestEntryStatus.ADDED, snapshot_id=25, data_file=data_file)) + manifest_file = writer.to_manifest_file() + + expected_content = "data" if content == ManifestContent.DATA else "deletes" + _verify_metadata_with_fastavro(path, {"format-version": "3", "content": expected_content}) + assert manifest_file.content == content + assert manifest_file.first_row_id is None + + with open(path, "rb") as f: + records = [r["data_file"] for r in fastavro.reader(f)] + assert [r["file_path"] for r in records] == ["/data/file-v3.puffin", "/data/file-v2.parquet", "/data/file-v1.parquet"] + assert [r["record_count"] for r in records] == [100, 50, 10] + # The first row id of added files stays unset unless it was already assigned + assert [r["first_row_id"] for r in records] == [1000, None, None] + assert records[0]["key_metadata"] == b"\x01" + assert records[0]["referenced_data_file"] == "/data/referenced.parquet" + assert records[0]["content_offset"] == 4 + assert records[0]["content_size_in_bytes"] == 42 + + read_back = manifest_file.fetch_manifest_entry(io)[0].data_file + assert read_back.referenced_data_file == "/data/referenced.parquet" + assert read_back.content_offset == 4 + assert read_back.content_size_in_bytes == 42 + + +def test_write_delete_manifest_v1_is_rejected(tmp_path: Path) -> None: + with pytest.raises(ValidationError, match="Cannot write delete manifests in a v1 table"): + write_manifest( + format_version=1, + spec=UNPARTITIONED_PARTITION_SPEC, + schema=Schema(NestedField(1, "foo", IntegerType(), False)), + output_file=load_file_io().new_output(str(tmp_path / "manifest.avro")), + snapshot_id=25, + avro_compression="null", + content=ManifestContent.DELETES, + ) + + +def _manifest_args( + path: str, content: ManifestContent = ManifestContent.DATA, added: int = 0, existing: int = 0 +) -> dict[str, Any]: + return { + "manifest_path": path, + "manifest_length": 100, + "partition_spec_id": 0, + "content": content, + "sequence_number": 1, + "min_sequence_number": 1, + "added_snapshot_id": 25, + "added_files_count": 1, + "existing_files_count": 1, + "deleted_files_count": 0, + "added_rows_count": added, + "existing_rows_count": existing, + "deleted_rows_count": 0, + } + + +@pytest.mark.parametrize("compression", ["null", "deflate"]) +def test_write_manifest_list_v3_assigns_first_row_id(tmp_path: Path, compression: AvroCompressionCodec) -> None: + io = load_file_io() + manifests = [ + # Bound to the V2 layout to exercise rebinding in the writer + ManifestFile.from_args(_table_format_version=2, **_manifest_args("/m1.avro", added=100, existing=25)), + ManifestFile.from_args(first_row_id=77, **_manifest_args("/m2.avro", added=10)), + ManifestFile.from_args(**_manifest_args("/m3.avro", content=ManifestContent.DELETES, added=10)), + ManifestFile.from_args(**_manifest_args("/m4.avro", added=5)), + ] + + path = str(tmp_path / "manifest-list-v3.avro") + with write_manifest_list( + format_version=3, + output_file=io.new_output(path), + snapshot_id=25, + parent_snapshot_id=19, + sequence_number=2, + avro_compression=compression, + first_row_id=1000, + ) as writer: + writer.add_manifests(manifests) + + # 1000 + (100 + 25) + 5: only unassigned data manifests advance the row id + assert writer.next_row_id == 1130 # type: ignore[attr-defined] + # The first row ids are assigned on copies, the input records are left untouched + assert [m.first_row_id for m in manifests[1:]] == [77, None, None] + + _verify_metadata_with_fastavro( + path, + {"snapshot-id": "25", "parent-snapshot-id": "19", "sequence-number": "2", "first-row-id": "1000", "format-version": "3"}, + ) + with open(path, "rb") as f: + assert [r["first_row_id"] for r in fastavro.reader(f)] == [1000, 77, None, 1125] + + read_back = list(read_manifest_list(io.new_input(path))) + assert [m.first_row_id for m in read_back] == [1000, 77, None, 1125] + + # Carrying the manifests of an existing V3 list into a new one preserves their first row ids + rewritten_path = str(tmp_path / "manifest-list-v3-rewritten.avro") + with write_manifest_list( + format_version=3, + output_file=io.new_output(rewritten_path), + snapshot_id=26, + parent_snapshot_id=25, + sequence_number=3, + avro_compression=compression, + first_row_id=1130, + ) as rewriter: + rewriter.add_manifests(read_back) + assert rewriter.next_row_id == 1130 # type: ignore[attr-defined] + assert [m.first_row_id for m in read_manifest_list(io.new_input(rewritten_path))] == [1000, 77, None, 1125] + + +def test_write_manifest_list_v3_requires_first_row_id(tmp_path: Path) -> None: + with pytest.raises(ValueError, match="First-row-id is required for V3 tables"): + write_manifest_list( + format_version=3, + output_file=load_file_io().new_output(str(tmp_path / "manifest-list.avro")), + snapshot_id=25, + parent_snapshot_id=19, + sequence_number=2, + avro_compression="null", + ) + + +def test_write_manifest_list_v3_rejects_unknown_row_counts(tmp_path: Path) -> None: + manifest = ManifestFile.from_args(**{**_manifest_args("/m1.avro"), "added_rows_count": None, "existing_rows_count": None}) + with pytest.raises(ValueError, match="unknown row counts"): + with write_manifest_list( + format_version=3, + output_file=load_file_io().new_output(str(tmp_path / "manifest-list.avro")), + snapshot_id=25, + parent_snapshot_id=19, + sequence_number=2, + avro_compression="null", + first_row_id=1000, + ) as writer: + writer.add_manifests([manifest]) diff --git a/tests/utils/test_schema_conversion.py b/tests/utils/test_schema_conversion.py index b110f92bd2..864fe56457 100644 --- a/tests/utils/test_schema_conversion.py +++ b/tests/utils/test_schema_conversion.py @@ -33,7 +33,9 @@ NestedField, StringType, StructType, + TimestampNanoType, TimestampType, + TimestamptzNanoType, UnknownType, UUIDType, ) @@ -347,6 +349,46 @@ def test_convert_timestamp_micros_type() -> None: assert actual == TimestampType() +@pytest.mark.parametrize( + "avro_logical_type, expected", + [ + ({"type": "long", "logicalType": "timestamp-nanos"}, TimestampNanoType()), + ({"type": "long", "logicalType": "timestamp-nanos", "adjust-to-utc": False}, TimestampNanoType()), + ({"type": "long", "logicalType": "timestamp-nanos", "adjust-to-utc": True}, TimestamptzNanoType()), + ], +) +def test_convert_timestamp_nanos_type(avro_logical_type: dict[str, Any], expected: Any) -> None: + assert AvroSchemaConversion()._convert_logical_type(avro_logical_type) == expected + + +def test_timestamp_nanos_avro_roundtrip() -> None: + schema = Schema( + NestedField(1, "ts_ns", TimestampNanoType(), required=True), + NestedField(2, "tstz_ns", TimestamptzNanoType(), required=True), + ) + avro_schema = AvroSchemaConversion().iceberg_to_avro(schema) + assert isinstance(avro_schema, dict) + assert AvroSchemaConversion().avro_to_iceberg(avro_schema) == schema + + +def test_iceberg_to_avro_defaults() -> None: + schema = Schema( + NestedField(1, "d", DateType(), required=True, write_default=19052), + NestedField(2, "b", BinaryType(), required=False, write_default=b"\xff"), + NestedField(3, "u", UnknownType(), required=False), + ) + avro_schema = AvroSchemaConversion().iceberg_to_avro(schema) + assert isinstance(avro_schema, dict) + fields = {field["name"]: field for field in avro_schema["fields"]} + # Required fields carry the default in the Avro encoding of the type + assert fields["d"]["default"] == 19052 + # The default of a nullable union must match its first branch, null + assert (fields["b"]["type"], fields["b"]["default"]) == (["null", "bytes"], None) + # A column of type unknown is written as a plain null instead of an invalid union of nulls + assert (fields["u"]["type"], fields["u"]["default"]) == ("null", None) + assert AvroSchemaConversion().avro_to_iceberg(avro_schema).find_field("u") == NestedField(3, "u", UnknownType()) + + def test_unknown_logical_type() -> None: """Test raising a ValueError when converting an unknown logical type as part of an Avro schema conversion""" avro_logical_type = {"type": "bytes", "logicalType": "date"} @@ -389,3 +431,21 @@ def test_iceberg_to_avro_manifest(avro_schema_manifest_entry: dict[str, Any]) -> iceberg_schema = AvroSchemaConversion().avro_to_iceberg(avro_schema_manifest_entry) avro_result = AvroSchemaConversion().iceberg_to_avro(iceberg_schema, schema_name="manifest_entry") assert avro_schema_manifest_entry == avro_result + + +def test_iceberg_variant_to_avro() -> None: + from pyiceberg.types import NestedField, VariantType + + schema = Schema(NestedField(1, "v", VariantType(), required=False)) + avro_schema = AvroSchemaConversion().iceberg_to_avro(schema, schema_name="table") + + assert isinstance(avro_schema, dict) + assert avro_schema["fields"][0]["type"] == [ + "null", + { + "type": "record", + "logicalType": "variant", + "fields": [{"name": "metadata", "type": "bytes"}, {"name": "value", "type": "bytes"}], + "name": "r1", + }, + ]