This is an automated email from the ASF dual-hosted git repository.

laskoviymishka pushed a commit to branch main
in repository https://gitbox.apache.org/repos/asf/iceberg-go.git


The following commit(s) were added to refs/heads/main by this push:
     new 8f1a87eb2 fix(table): canonicalize partition values used as map keys 
(#1581)
8f1a87eb2 is described below

commit 8f1a87eb2184e0f01dcaf1a8c0db753c733d9c2e
Author: Minh Vu <[email protected]>
AuthorDate: Fri Jul 31 16:21:49 2026 +0200

    fix(table): canonicalize partition values used as map keys (#1581)
    
    ## What changed
    
    Canonicalize partition values used internally for fanout and
    clustered-writer lookup. Binary and fixed values use a distinct
    string-backed key, and float NaNs use stable typed sentinels.
    
    Original partition values remain unchanged for paths and metadata. Byte
    values are cloned when a partition is first recorded so metadata does
    not retain Arrow-owned storage. Fixed-size Arrow binary values are also
    converted to Iceberg literals before grouping.
    
    ## Why
    
    Binary and fixed partition values are byte slices, which panic when used
    as Go map keys. Clustered partition equality could panic for the same
    reason, while repeated NaNs did not compare equal and could split one
    logical partition.
    
    ## Testing
    
    Coverage exercises binary, fixed, and NaN grouping; source-buffer
    mutation; clustered equality; and closed-partition tracking.
    
    - `go test ./table`
    - `go vet ./table`
    
    ---------
    
    Signed-off-by: Minh Vu <[email protected]>
---
 table/clustered_writer.go                      |   9 +-
 table/clustered_writer_test.go                 |  21 ++++
 table/partitioned_fanout_writer.go             |  76 ++++++++++++--
 table/partitioned_fanout_writer_review_test.go | 133 +++++++++++++++++++++++++
 table/partitioned_fanout_writer_test.go        |  69 ++++++++++++-
 5 files changed, 297 insertions(+), 11 deletions(-)

diff --git a/table/clustered_writer.go b/table/clustered_writer.go
index 1620112a0..54296017e 100644
--- a/table/clustered_writer.go
+++ b/table/clustered_writer.go
@@ -113,7 +113,7 @@ func clusteredPartitionedWrite(
                        }
                        defer subBatch.Release()
 
-                       if currentWriter == nil || !slices.Equal(currentRec, 
part.partitionRec) {
+                       if currentWriter == nil || 
!partitionRecordsEqual(currentRec, part.partitionRec) {
                                if err := closeCurrentWriter(); err != nil {
                                        return err
                                }
@@ -213,10 +213,11 @@ type closedPartitionSet map[any]closedPartitionSet
 func (s closedPartitionSet) add(rec partitionRecord) {
        node := s
        for _, part := range rec {
-               next, ok := node[part]
+               key := comparablePartitionKey(part)
+               next, ok := node[key]
                if !ok {
                        next = make(closedPartitionSet)
-                       node[part] = next
+                       node[key] = next
                }
                node = next
        }
@@ -225,7 +226,7 @@ func (s closedPartitionSet) add(rec partitionRecord) {
 func (s closedPartitionSet) contains(rec partitionRecord) bool {
        node := s
        for _, part := range rec {
-               next, ok := node[part]
+               next, ok := node[comparablePartitionKey(part)]
                if !ok {
                        return false
                }
diff --git a/table/clustered_writer_test.go b/table/clustered_writer_test.go
index 25fa7216c..c816bf285 100644
--- a/table/clustered_writer_test.go
+++ b/table/clustered_writer_test.go
@@ -20,6 +20,7 @@ package table
 import (
        "context"
        "fmt"
+       "math"
        "path/filepath"
        "runtime"
        "testing"
@@ -38,6 +39,26 @@ import (
 
 // -- Functional tests --
 
+func TestClusteredPartitionTrackingUsesComparableKeys(t *testing.T) {
+       t.Parallel()
+
+       records := []partitionRecord{
+               {[]byte{1, 2, 3}},
+               {math.NaN()},
+       }
+       completed := make(closedPartitionSet)
+       for _, record := range records {
+               completed.add(record)
+               require.True(t, completed.contains(record))
+               require.True(t, partitionRecordsEqual(record, record))
+       }
+
+       require.True(t, completed.contains(partitionRecord{[]byte{1, 2, 3}}))
+       require.True(t, completed.contains(partitionRecord{math.NaN()}))
+       require.True(t, partitionRecordsEqual(partitionRecord{[]byte{1, 2, 3}}, 
partitionRecord{[]byte{1, 2, 3}}))
+       require.True(t, partitionRecordsEqual(partitionRecord{math.NaN()}, 
partitionRecord{math.NaN()}))
+}
+
 type ClusteredWriterTestSuite struct {
        suite.Suite
 
diff --git a/table/partitioned_fanout_writer.go 
b/table/partitioned_fanout_writer.go
index 39b716d02..48d0f23f4 100644
--- a/table/partitioned_fanout_writer.go
+++ b/table/partitioned_fanout_writer.go
@@ -23,6 +23,7 @@ import (
        "fmt"
        "iter"
        "math"
+       "slices"
 
        "github.com/apache/arrow-go/v18/arrow"
        "github.com/apache/arrow-go/v18/arrow/array"
@@ -54,10 +55,55 @@ type partitionInfo struct {
 
 type partitionFieldInfo struct {
        sourceField iceberg.PartitionField
+       sourceName  string
        fieldID     int
        sourceType  iceberg.Type
 }
 
+type binaryPartitionKey string
+
+type nanPartitionKey struct {
+       bits int
+}
+
+func comparablePartitionKey(value any) any {
+       switch value := value.(type) {
+       case []byte:
+               return binaryPartitionKey(value)
+       case float32:
+               if math.IsNaN(float64(value)) {
+                       return nanPartitionKey{bits: 32}
+               }
+       case float64:
+               if math.IsNaN(value) {
+                       return nanPartitionKey{bits: 64}
+               }
+       }
+
+       return value
+}
+
+func partitionRecordsEqual(left, right partitionRecord) bool {
+       if len(left) != len(right) {
+               return false
+       }
+       for i := range left {
+               if comparablePartitionKey(left[i]) != 
comparablePartitionKey(right[i]) {
+                       return false
+               }
+       }
+
+       return true
+}
+
+func clonePartitionValue(value any) any {
+       if bytes, ok := value.([]byte); ok {
+               return slices.Clone(bytes)
+       }
+
+       return value
+}
+
 // NewPartitionedFanoutWriter creates a new PartitionedFanoutWriter with the 
specified
 // partition specification, schema, record iterator, and writerFactory.
 func newPartitionedFanoutWriter(partitionSpec iceberg.PartitionSpec, schema 
*iceberg.Schema, itr iter.Seq2[arrow.RecordBatch, error], writerFactory 
*writerFactory) *partitionedFanoutWriter {
@@ -285,6 +331,7 @@ func getRecordPartitions(spec iceberg.PartitionSpec, schema 
*iceberg.Schema, rec
                partitionColumns[i] = record.Column(colIndices[0])
                partitionFieldsInfo[i] = partitionFieldInfo{
                        sourceField: sourceField,
+                       sourceName:  colName,
                        fieldID:     sourceField.FieldID,
                        sourceType:  sourceType,
                }
@@ -298,7 +345,14 @@ func getRecordPartitions(spec iceberg.PartitionSpec, 
schema *iceberg.Schema, rec
                                sourceField := fieldInfo.sourceField
                                val, err := getArrowValueAsIcebergLiteral(col, 
int(row), fieldInfo.sourceType)
                                if err != nil {
-                                       return nil, fmt.Errorf("failed to get 
arrow values as iceberg literal: %w", err)
+                                       return nil, fmt.Errorf(
+                                               "failed to convert source 
column %q (field ID %d) from Arrow type %s to Iceberg type %s: %w",
+                                               fieldInfo.sourceName,
+                                               sourceField.SourceID(),
+                                               col.DataType(),
+                                               fieldInfo.sourceType,
+                                               err,
+                                       )
                                }
 
                                transformedLiteral := 
sourceField.Transform.Apply(iceberg.Optional[iceberg.Literal]{Valid: true, Val: 
val})
@@ -343,10 +397,11 @@ func (n *partitionMapNode) getOrCreate(partitionRec 
partitionRecord, fieldInfo [
        // Navigate through all but the last partition field
        node := n
        for _, part := range partitionRec[:len(partitionRec)-1] {
-               val, ok := node.children[part]
+               key := comparablePartitionKey(part)
+               val, ok := node.children[key]
                if !ok {
                        newNode := newPartitionMapNode()
-                       node.children[part] = newNode
+                       node.children[key] = newNode
                        node = newNode
                } else {
                        node = val.(*partitionMapNode)
@@ -354,7 +409,7 @@ func (n *partitionMapNode) getOrCreate(partitionRec 
partitionRecord, fieldInfo [
        }
 
        // Last level stores the actual partitionInfo
-       lastKey := partitionRec[len(partitionRec)-1]
+       lastKey := comparablePartitionKey(partitionRec[len(partitionRec)-1])
        partVal, ok := node.children[lastKey].(*partitionInfo)
        if ok {
                return partVal
@@ -366,8 +421,9 @@ func (n *partitionMapNode) getOrCreate(partitionRec 
partitionRecord, fieldInfo [
        // Copy partitionRec values so they don't get overwritten
        partRecCopy := make(partitionRecord, len(partitionRec))
        for i := range partitionRec {
-               partitionValues[fieldInfo[i].fieldID] = partitionRec[i]
-               partRecCopy[i] = partitionRec[i]
+               value := clonePartitionValue(partitionRec[i])
+               partitionValues[fieldInfo[i].fieldID] = value
+               partRecCopy[i] = value
        }
 
        partVal = &partitionInfo{
@@ -518,6 +574,14 @@ func getArrowValueAsIcebergLiteral(column arrow.Array, row 
int, sourceType icebe
        case *array.LargeBinary:
 
                return iceberg.NewLiteral(arr.Value(row)), nil
+       case *array.FixedSizeBinary:
+               switch sourceType.(type) {
+               case iceberg.BinaryType, iceberg.FixedType, iceberg.UUIDType:
+               default:
+                       return nil, fmt.Errorf("%w: cannot convert Arrow %s to 
Iceberg type %v", iceberg.ErrInvalidSchema, arr.DataType(), sourceType)
+               }
+
+               return iceberg.NewLiteral(arr.Value(row)).To(sourceType)
 
        default:
                val := column.GetOneForMarshal(row)
diff --git a/table/partitioned_fanout_writer_review_test.go 
b/table/partitioned_fanout_writer_review_test.go
new file mode 100644
index 000000000..14e30bb08
--- /dev/null
+++ b/table/partitioned_fanout_writer_review_test.go
@@ -0,0 +1,133 @@
+// 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.
+
+package table
+
+import (
+       "unsafe"
+
+       "github.com/apache/arrow-go/v18/arrow"
+       "github.com/apache/arrow-go/v18/arrow/array"
+       "github.com/apache/iceberg-go"
+)
+
+func (s *FanoutWriterTestSuite) 
TestBinaryPartitionValuesDoNotAliasArrowStorage() {
+       tests := []struct {
+               name        string
+               arrowType   arrow.DataType
+               icebergType iceberg.Type
+       }{
+               {name: "binary", arrowType: arrow.BinaryTypes.Binary, 
icebergType: iceberg.PrimitiveTypes.Binary},
+               {name: "fixed", arrowType: 
&arrow.FixedSizeBinaryType{ByteWidth: 4}, icebergType: iceberg.FixedTypeOf(4)},
+       }
+
+       for _, test := range tests {
+               s.Run(test.name, func() {
+                       arrowSchema := arrow.NewSchema([]arrow.Field{{Name: 
"part", Type: test.arrowType}}, nil)
+                       record := s.createCustomTestRecord(arrowSchema, [][]any{
+                               {[]byte{1, 2, 3, 4}},
+                               {[]byte{1, 2, 3, 4}},
+                               {[]byte{5, 6, 7, 8}},
+                       })
+                       defer record.Release()
+
+                       icebergSchema := iceberg.NewSchema(1, 
iceberg.NestedField{ID: 1, Name: "part", Type: test.icebergType})
+                       spec := iceberg.NewPartitionSpec(iceberg.PartitionField{
+                               SourceIDs: []int{1}, FieldID: 1000, Name: 
"part", Transform: iceberg.IdentityTransform{},
+                       })
+
+                       partitions, err := getRecordPartitions(spec, 
icebergSchema, record)
+                       s.Require().NoError(err)
+                       s.Require().Len(partitions, 2)
+
+                       var arrowValue []byte
+                       switch values := record.Column(0).(type) {
+                       case *array.Binary:
+                               arrowValue = values.Value(0)
+                       case *array.FixedSizeBinary:
+                               arrowValue = values.Value(0)
+                       default:
+                               s.FailNow("unsupported byte-slice column", 
"%T", record.Column(0))
+                       }
+
+                       var storedValue []byte
+                       for _, partition := range partitions {
+                               value, ok := partition.partitionRec[0].([]byte)
+                               s.Require().True(ok)
+                               if string(value) == string([]byte{1, 2, 3, 4}) {
+                                       storedValue = value
+
+                                       break
+                               }
+                       }
+                       s.Require().NotNil(storedValue)
+                       s.NotEqual(
+                               
uintptr(unsafe.Pointer(unsafe.SliceData(arrowValue))),
+                               
uintptr(unsafe.Pointer(unsafe.SliceData(storedValue))),
+                               "stored partition value must not alias 
Arrow-owned storage",
+                       )
+
+                       arrowValue[0] = 9
+                       s.Equal([]byte{1, 2, 3, 4}, storedValue)
+               })
+       }
+}
+
+func (s *FanoutWriterTestSuite) 
TestFixedSizeBinaryPartitionReportsWidthMismatch() {
+       arrowSchema := arrow.NewSchema([]arrow.Field{{
+               Name: "part",
+               Type: &arrow.FixedSizeBinaryType{ByteWidth: 4},
+       }}, nil)
+       record := s.createCustomTestRecord(arrowSchema, [][]any{{[]byte{1, 2, 
3, 4}}})
+       defer record.Release()
+
+       icebergSchema := iceberg.NewSchema(1, iceberg.NestedField{
+               ID: 1, Name: "part", Type: iceberg.FixedTypeOf(3),
+       })
+       spec := iceberg.NewPartitionSpec(iceberg.PartitionField{
+               SourceIDs: []int{1}, FieldID: 1000, Name: "part", Transform: 
iceberg.IdentityTransform{},
+       })
+
+       _, err := getRecordPartitions(spec, icebergSchema, record)
+       s.Require().Error(err)
+       s.ErrorContains(err, `source column "part"`)
+       s.ErrorContains(err, "field ID 1")
+       s.ErrorContains(err, "fixed[3]")
+}
+
+func (s *FanoutWriterTestSuite) 
TestFixedSizeBinaryPartitionRejectsUnsupportedIcebergType() {
+       arrowSchema := arrow.NewSchema([]arrow.Field{{
+               Name: "part",
+               Type: &arrow.FixedSizeBinaryType{ByteWidth: 4},
+       }}, nil)
+       record := s.createCustomTestRecord(arrowSchema, [][]any{{[]byte{1, 2, 
3, 4}}})
+       defer record.Release()
+
+       icebergSchema := iceberg.NewSchema(1, iceberg.NestedField{
+               ID: 1, Name: "part", Type: iceberg.PrimitiveTypes.String,
+       })
+       spec := iceberg.NewPartitionSpec(iceberg.PartitionField{
+               SourceIDs: []int{1}, FieldID: 1000, Name: "part", Transform: 
iceberg.IdentityTransform{},
+       })
+
+       _, err := getRecordPartitions(spec, icebergSchema, record)
+       s.Require().ErrorIs(err, iceberg.ErrInvalidSchema)
+       s.ErrorContains(err, `source column "part"`)
+       s.ErrorContains(err, "field ID 1")
+       s.ErrorContains(err, "fixed_size_binary[4]")
+       s.ErrorContains(err, "string")
+}
diff --git a/table/partitioned_fanout_writer_test.go 
b/table/partitioned_fanout_writer_test.go
index 929e44be1..16a74d203 100644
--- a/table/partitioned_fanout_writer_test.go
+++ b/table/partitioned_fanout_writer_test.go
@@ -83,7 +83,14 @@ func (s *FanoutWriterTestSuite) 
createCustomTestRecord(arrSchema *arrow.Schema,
                        case uuid.UUID:
                                field.(*extensions.UUIDBuilder).Append(t)
                        case []byte:
-                               field.(*array.BinaryBuilder).Append(t)
+                               switch builder := field.(type) {
+                               case *array.BinaryBuilder:
+                                       builder.Append(t)
+                               case *array.FixedSizeBinaryBuilder:
+                                       builder.Append(t)
+                               default:
+                                       s.FailNow("unsupported byte-slice 
builder", "%T", field)
+                               }
                        default:
                                appendMethod.Call([]reflect.Value{v})
                        }
@@ -282,6 +289,66 @@ func (s *FanoutWriterTestSuite) TestIdentityTransform() {
        s.testTransformPartition(iceberg.IdentityTransform{}, "large_name", 
"identity_large_string", testRecord, 5)
 }
 
+func (s *FanoutWriterTestSuite) TestBinaryPartitionValuesUseComparableKeys() {
+       tests := []struct {
+               name        string
+               arrowType   arrow.DataType
+               icebergType iceberg.Type
+       }{
+               {name: "binary", arrowType: arrow.BinaryTypes.Binary, 
icebergType: iceberg.PrimitiveTypes.Binary},
+               {name: "fixed", arrowType: 
&arrow.FixedSizeBinaryType{ByteWidth: 4}, icebergType: iceberg.FixedTypeOf(4)},
+       }
+
+       for _, test := range tests {
+               s.Run(test.name, func() {
+                       arrowSchema := arrow.NewSchema([]arrow.Field{{Name: 
"part", Type: test.arrowType}}, nil)
+                       record := s.createCustomTestRecord(arrowSchema, 
[][]any{{[]byte{1, 2, 3, 4}}, {[]byte{1, 2, 3, 4}}, {[]byte{5, 6, 7, 8}}})
+                       defer record.Release()
+
+                       icebergSchema := iceberg.NewSchema(1, 
iceberg.NestedField{ID: 1, Name: "part", Type: test.icebergType})
+                       spec := iceberg.NewPartitionSpec(iceberg.PartitionField{
+                               SourceIDs: []int{1}, FieldID: 1000, Name: 
"part", Transform: iceberg.IdentityTransform{},
+                       })
+
+                       partitions, err := getRecordPartitions(spec, 
icebergSchema, record)
+                       s.Require().NoError(err)
+                       s.Require().Len(partitions, 2)
+                       switch values := record.Column(0).(type) {
+                       case *array.Binary:
+                               values.Value(0)[0] = 9
+                       case *array.FixedSizeBinary:
+                               values.Value(0)[0] = 9
+                       }
+
+                       rowsByValue := make(map[string]int)
+                       for _, partition := range partitions {
+                               value, ok := partition.partitionRec[0].([]byte)
+                               s.Require().True(ok)
+                               rowsByValue[string(value)] = len(partition.rows)
+                       }
+                       s.Equal(2, rowsByValue[string([]byte{1, 2, 3, 4})])
+                       s.Equal(1, rowsByValue[string([]byte{5, 6, 7, 8})])
+               })
+       }
+}
+
+func (s *FanoutWriterTestSuite) TestNaNPartitionValuesUseStableKeys() {
+       arrowSchema := arrow.NewSchema([]arrow.Field{{Name: "part", Type: 
arrow.PrimitiveTypes.Float64}}, nil)
+       record := s.createCustomTestRecord(arrowSchema, [][]any{{math.NaN()}, 
{math.NaN()}})
+       defer record.Release()
+
+       icebergSchema := iceberg.NewSchema(1, iceberg.NestedField{ID: 1, Name: 
"part", Type: iceberg.PrimitiveTypes.Float64})
+       spec := iceberg.NewPartitionSpec(iceberg.PartitionField{
+               SourceIDs: []int{1}, FieldID: 1000, Name: "part", Transform: 
iceberg.IdentityTransform{},
+       })
+
+       partitions, err := getRecordPartitions(spec, icebergSchema, record)
+       s.Require().NoError(err)
+       s.Require().Len(partitions, 1)
+       s.Len(partitions[0].rows, 2)
+       s.True(math.IsNaN(partitions[0].partitionRec[0].(float64)))
+}
+
 func (s *FanoutWriterTestSuite) TestBucketTransform() {
        arrSchema := arrow.NewSchema([]arrow.Field{
                {Name: "id", Type: arrow.PrimitiveTypes.Int32, Nullable: true},

Reply via email to