Skip to content

Commit c2e889f

Browse files
dfa1claude
andcommitted
refactor(reader): remove footgun ArraySegments.of(Array) overload
The no-arena of(Array) threw on lazy arrays, and its throw was being abused as an "is it segment-backed?" probe in ScanIterator. Replace it: - of(Array) body becomes private primarySegment; the surviving of(Array, SegmentAllocator) overload uses it for materialized/segment-backed inputs and materializes lazy ones. - Add public trySegment(Array) -> Optional<MemorySegment> for the zero-alloc segment-backed probe (ScanIterator dict-codes capacity check). - Production decoders that had an allocator in scope (RunEnd, Rle) move to the arena overload. - Tests/bench route through of(arr, Arena.ofAuto()) (zero-copy for materialized, correct for now-lazy scan output); PatchedEncodingDecoderTest uses typed getters. Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com>
1 parent a7dcb6a commit c2e889f

20 files changed

Lines changed: 102 additions & 88 deletions

integration/src/test/java/io/github/dfa1/vortex/integration/RustWritesJavaReadsIntegrationTest.java

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -34,6 +34,7 @@
3434
import org.junit.jupiter.api.io.TempDir;
3535

3636
import java.io.IOException;
37+
import java.lang.foreign.Arena;
3738
import java.lang.foreign.ValueLayout;
3839
import java.net.URI;
3940
import java.nio.ByteOrder;
@@ -129,7 +130,7 @@ private static List<JavaChunk> scanAll(VortexReader vf,
129130
/// into a heap primitive array — long[]/int[]/double[]/float[]/short[]/byte[].
130131
private static Object snapshotArray(Array arr) {
131132
var ptype = ((DType.Primitive) arr.dtype()).ptype();
132-
var seg = ArraySegments.of(arr);
133+
var seg = ArraySegments.of(arr, Arena.ofAuto());
133134
return switch (ptype) {
134135
case I64, U64 -> seg.toArray(ValueLayout.JAVA_LONG_UNALIGNED.withOrder(ByteOrder.LITTLE_ENDIAN));
135136
case I32, U32 -> seg.toArray(ValueLayout.JAVA_INT_UNALIGNED.withOrder(ByteOrder.LITTLE_ENDIAN));

performance/src/main/java/io/github/dfa1/vortex/performance/RustWritesJavaReadsBigFileBenchmark.java

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -37,6 +37,7 @@
3737
import org.openjdk.jmh.annotations.Warmup;
3838

3939
import java.io.IOException;
40+
import java.lang.foreign.Arena;
4041
import java.lang.foreign.MemorySegment;
4142
import java.lang.foreign.ValueLayout;
4243
import java.nio.ByteOrder;
@@ -180,7 +181,7 @@ private long scanJava() throws IOException {
180181
while (iter.hasNext()) {
181182
try (Chunk c = iter.next()) {
182183
Array arr = c.columns().get("c0");
183-
MemorySegment buf = ArraySegments.of(arr);
184+
MemorySegment buf = ArraySegments.of(arr, Arena.ofAuto());
184185
long count = buf.byteSize() / Long.BYTES;
185186
for (long i = 0; i < count; i++) {
186187
sum += buf.getAtIndex(LE_LONG, i);

reader/src/main/java/io/github/dfa1/vortex/reader/ScanIterator.java

Lines changed: 4 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -47,6 +47,7 @@
4747
import java.util.List;
4848
import java.util.Map;
4949
import java.util.NoSuchElementException;
50+
import java.util.Optional;
5051
import java.util.function.Consumer;
5152

5253
/// Iterates over decoded chunks from a [io.github.dfa1.vortex.reader.VortexReader].
@@ -641,13 +642,11 @@ private Array decodeDictLayout(Layout dictLayout, DType dtype, SegmentAllocator
641642
/// @param codesPType code ptype reported by the dict layout metadata
642643
/// @param n claimed dict row count
643644
private static void validateDictCodesCapacity(Array codes, PType codesPType, long n) {
644-
MemorySegment seg;
645-
try {
646-
seg = ArraySegments.of(codes);
647-
} catch (VortexException e) {
645+
Optional<MemorySegment> maybeSeg = ArraySegments.trySegment(codes);
646+
if (maybeSeg.isEmpty()) {
648647
return;
649648
}
650-
long bufferCodes = seg.byteSize() / (long) codesPType.byteSize();
649+
long bufferCodes = maybeSeg.get().byteSize() / (long) codesPType.byteSize();
651650
if (bufferCodes < n) {
652651
throw new VortexException(EncodingId.VORTEX_DICT,
653652
"dict codes: layout row_count=" + n + " exceeds buffer capacity=" + bufferCodes);

reader/src/main/java/io/github/dfa1/vortex/reader/array/ArraySegments.java

Lines changed: 29 additions & 20 deletions
Original file line numberDiff line numberDiff line change
@@ -5,6 +5,7 @@
55

66
import java.lang.foreign.MemorySegment;
77
import java.lang.foreign.SegmentAllocator;
8+
import java.util.Optional;
89

910
/// Utility for extracting the primary {@link MemorySegment} from any {@link Array}.
1011
///
@@ -25,31 +26,39 @@ public final class ArraySegments {
2526
private ArraySegments() {
2627
}
2728

28-
/// Returns the primary backing segment of `arr`.
29+
/// Returns the primary backing segment of `arr` if it is segment-backed, otherwise empty.
30+
///
31+
/// Non-throwing probe for callers that want to operate on the raw buffer only when one
32+
/// exists (e.g. zone-map / capacity validation) and skip lazy variants without allocating.
33+
/// To force a segment for a lazy array, use [#of(Array, SegmentAllocator)].
2934
///
3035
/// @param arr the array whose segment is needed
31-
/// @return the primary {@link MemorySegment}
32-
/// @throws VortexException if the array type has no primary segment (e.g. lazy variants — use
33-
/// {@link #of(Array, SegmentAllocator)} instead)
34-
public static MemorySegment of(Array arr) {
36+
/// @return the primary [MemorySegment], or empty if `arr` has no segment backing
37+
public static Optional<MemorySegment> trySegment(Array arr) {
3538
Array data = arr instanceof MaskedArray m ? m.inner() : arr;
3639
return switch (data) {
37-
case MaterializedIntArray a -> a.buffer();
38-
case MaterializedLongArray a -> a.buffer();
39-
case MaterializedDoubleArray a -> a.buffer();
40-
case MaterializedFloatArray a -> a.buffer();
41-
case MaterializedShortArray a -> a.buffer();
42-
case MaterializedByteArray a -> a.buffer();
43-
case MaterializedBoolArray a -> a.buffer();
44-
case MaterializedFloat16Array a -> a.buffer();
45-
case VarBinArray a -> a.bytesSegment();
46-
case GenericArray a -> a.buffer(0);
47-
case LazyDecimalArray a -> a.buf();
48-
case DecimalArray a -> throw new VortexException(a.getClass().getSimpleName() + " has no primary segment — use of(arr, arena)");
49-
default -> throw new VortexException(data.getClass().getSimpleName() + " has no primary segment");
40+
case MaterializedIntArray a -> Optional.of(a.buffer());
41+
case MaterializedLongArray a -> Optional.of(a.buffer());
42+
case MaterializedDoubleArray a -> Optional.of(a.buffer());
43+
case MaterializedFloatArray a -> Optional.of(a.buffer());
44+
case MaterializedShortArray a -> Optional.of(a.buffer());
45+
case MaterializedByteArray a -> Optional.of(a.buffer());
46+
case MaterializedBoolArray a -> Optional.of(a.buffer());
47+
case MaterializedFloat16Array a -> Optional.of(a.buffer());
48+
case VarBinArray a -> Optional.of(a.bytesSegment());
49+
case GenericArray a -> Optional.of(a.buffer(0));
50+
case LazyDecimalArray a -> Optional.of(a.buf());
51+
default -> Optional.empty();
5052
};
5153
}
5254

55+
private static MemorySegment primarySegment(Array arr) {
56+
return trySegment(arr).orElseThrow(() -> {
57+
Array data = arr instanceof MaskedArray m ? m.inner() : arr;
58+
return new VortexException(data.getClass().getSimpleName() + " has no primary segment — use of(arr, arena)");
59+
});
60+
}
61+
5362
/// Returns the primary backing segment of `arr`, materialising lazy variants into a
5463
/// fresh segment allocated from `arena`.
5564
///
@@ -90,8 +99,8 @@ public static MemorySegment of(Array arr, SegmentAllocator arena) {
9099
case ShortArray a -> materialiseShort(a, arena);
91100
case ByteArray a -> materialiseByte(a, arena);
92101
case LazyConstantDecimalArray a -> materialiseConstantDecimal(a, arena);
93-
case DecimalArray _ -> of(arr);
94-
default -> of(arr);
102+
case DecimalArray _ -> primarySegment(arr);
103+
default -> primarySegment(arr);
95104
};
96105
}
97106

reader/src/main/java/io/github/dfa1/vortex/reader/decode/ConstantEncodingDecoder.java

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -85,7 +85,7 @@ private static Array arrayFromScalar(DecodeContext ctx, ScalarValue scalar, DTyp
8585
Array storage = arrayFromScalar(ctx, scalar, ext.storageDType(), n);
8686
// GenericArray needs a backing buffer; the recursive call returns a metadata-only
8787
// LazyConstantXxxArray. Materialise once into the chunk arena so downstream
88-
// extension consumers that read via ArraySegments.of(arr) still find a segment.
88+
// extension consumers that read via ArraySegments.of(arr, arena) still find a segment.
8989
// Extension-on-constant is rare enough that the small alloc doesn't matter — the
9090
// bare primitive path stays buffer-free.
9191
return new GenericArray(dtype, n, ArraySegments.of(storage, ctx.arena()));

reader/src/main/java/io/github/dfa1/vortex/reader/decode/RleEncodingDecoder.java

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -85,7 +85,7 @@ public Array decode(DecodeContext ctx) {
8585
indicesValidity = masked.validity();
8686
}
8787

88-
int[] indices = readIndices(ArraySegments.of(indicesArr), (int) indicesLen, indicesPtype);
88+
int[] indices = readIndices(ArraySegments.of(indicesArr, ctx.arena()), (int) indicesLen, indicesPtype);
8989
long[] valuesIdxOffsets = readUnsignedLongs(
9090
ctx.decodeChildSegment(2, offsetsDtype, offsetsLen), (int) offsetsLen, offsetsPtype);
9191
long firstOffset = valuesLen > 0 && valuesIdxOffsets.length > 0 ? valuesIdxOffsets[0] : 0L;

reader/src/main/java/io/github/dfa1/vortex/reader/decode/RunEndEncodingDecoder.java

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -104,7 +104,7 @@ private static Array expandStrings(
104104
PType endsPtype, long numRuns, long offset, long n,
105105
DType dtype, SegmentAllocator arena
106106
) {
107-
MemorySegment endsSeg = ArraySegments.of(endsArr);
107+
MemorySegment endsSeg = ArraySegments.of(endsArr, arena);
108108
long endsCap = SegmentBroadcast.capacity(endsSeg, endsPtype.byteSize());
109109
MemorySegment valBytes = valuesArr.bytesSegment();
110110
MemorySegment valOffsets = valuesArr.offsetsSegment();

reader/src/test/java/io/github/dfa1/vortex/reader/decode/PatchedEncodingDecoderTest.java

Lines changed: 18 additions & 20 deletions
Original file line numberDiff line numberDiff line change
@@ -5,11 +5,9 @@
55
import io.github.dfa1.vortex.core.DType;
66
import io.github.dfa1.vortex.core.PType;
77
import io.github.dfa1.vortex.reader.array.Array;
8-
import io.github.dfa1.vortex.reader.array.ArraySegments;
98
import io.github.dfa1.vortex.reader.array.IntArray;
109
import io.github.dfa1.vortex.reader.array.LongArray;
1110
import io.github.dfa1.vortex.encoding.EncodingId;
12-
import io.github.dfa1.vortex.encoding.PTypeIO;
1311
import io.github.dfa1.vortex.proto.PatchedMetadata;
1412
import org.junit.jupiter.api.Test;
1513
import org.junit.jupiter.params.ParameterizedTest;
@@ -102,9 +100,9 @@ void decode_noPatches_returnsInnerUnchanged() {
102100

103101
// Then
104102
assertThat(result).isInstanceOf(IntArray.class);
105-
MemorySegment seg = ArraySegments.of(result);
103+
IntArray ints = (IntArray) result;
106104
for (int i = 0; i < n; i++) {
107-
assertThat(seg.getAtIndex(PTypeIO.LE_INT, i)).as("index %d", i).isEqualTo(inner[i]);
105+
assertThat(ints.getInt(i)).as("index %d", i).isEqualTo(inner[i]);
108106
}
109107
}
110108

@@ -114,11 +112,11 @@ void decode_singlePatch_overwrites() {
114112
Array result = decode(4, new int[]{10, 20, 30, 40}, new int[]{0, 1}, new short[]{2}, new int[]{99});
115113

116114
// Then
117-
MemorySegment seg = ArraySegments.of(result);
118-
assertThat(seg.getAtIndex(PTypeIO.LE_INT, 0)).isEqualTo(10);
119-
assertThat(seg.getAtIndex(PTypeIO.LE_INT, 1)).isEqualTo(20);
120-
assertThat(seg.getAtIndex(PTypeIO.LE_INT, 2)).isEqualTo(99);
121-
assertThat(seg.getAtIndex(PTypeIO.LE_INT, 3)).isEqualTo(40);
115+
IntArray ints = (IntArray) result;
116+
assertThat(ints.getInt(0)).isEqualTo(10);
117+
assertThat(ints.getInt(1)).isEqualTo(20);
118+
assertThat(ints.getInt(2)).isEqualTo(99);
119+
assertThat(ints.getInt(3)).isEqualTo(40);
122120
}
123121

124122
@Test
@@ -127,11 +125,11 @@ void decode_multiplePatches_allApplied() {
127125
Array result = decode(4, new int[]{0, 0, 0, 0}, new int[]{0, 2}, new short[]{0, 3}, new int[]{1, 7});
128126

129127
// Then
130-
MemorySegment seg = ArraySegments.of(result);
131-
assertThat(seg.getAtIndex(PTypeIO.LE_INT, 0)).isEqualTo(1);
132-
assertThat(seg.getAtIndex(PTypeIO.LE_INT, 1)).isEqualTo(0);
133-
assertThat(seg.getAtIndex(PTypeIO.LE_INT, 2)).isEqualTo(0);
134-
assertThat(seg.getAtIndex(PTypeIO.LE_INT, 3)).isEqualTo(7);
128+
IntArray ints = (IntArray) result;
129+
assertThat(ints.getInt(0)).isEqualTo(1);
130+
assertThat(ints.getInt(1)).isEqualTo(0);
131+
assertThat(ints.getInt(2)).isEqualTo(0);
132+
assertThat(ints.getInt(3)).isEqualTo(7);
135133
}
136134

137135
@ParameterizedTest
@@ -144,9 +142,9 @@ void decode_variousLengths_noPatches(int n) {
144142
Array result = decode(n, inner, new int[]{0, 0}, new short[]{}, new int[]{});
145143

146144
// Then
147-
MemorySegment seg = ArraySegments.of(result);
145+
IntArray ints = (IntArray) result;
148146
for (int i = 0; i < n; i++) {
149-
assertThat(seg.getAtIndex(PTypeIO.LE_INT, i)).as("index %d", i).isZero();
147+
assertThat(ints.getInt(i)).as("index %d", i).isZero();
150148
}
151149
}
152150

@@ -161,10 +159,10 @@ void decode_i64_singlePatch() {
161159

162160
// Then
163161
assertThat(result).isInstanceOf(LongArray.class);
164-
MemorySegment seg = ArraySegments.of(result);
165-
assertThat(seg.getAtIndex(PTypeIO.LE_LONG, 0)).isEqualTo(100L);
166-
assertThat(seg.getAtIndex(PTypeIO.LE_LONG, 1)).isEqualTo(999L);
167-
assertThat(seg.getAtIndex(PTypeIO.LE_LONG, 2)).isEqualTo(300L);
162+
LongArray longs = (LongArray) result;
163+
assertThat(longs.getLong(0)).isEqualTo(100L);
164+
assertThat(longs.getLong(1)).isEqualTo(999L);
165+
assertThat(longs.getLong(2)).isEqualTo(300L);
168166
}
169167

170168
@Test

writer/src/test/java/io/github/dfa1/vortex/writer/DictEncodingTest.java

Lines changed: 8 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -13,6 +13,7 @@
1313
import org.junit.jupiter.api.io.TempDir;
1414

1515
import java.io.IOException;
16+
import java.lang.foreign.Arena;
1617
import java.lang.foreign.ValueLayout;
1718
import java.nio.ByteOrder;
1819
import java.nio.channels.FileChannel;
@@ -101,19 +102,19 @@ void roundTrip_multipleChunks(@TempDir Path tmp) throws IOException {
101102
try (Chunk c1 = iter.next()) {
102103
Array a1 = c1.columns().get("category");
103104
assertThat(a1.length()).isEqualTo(3L);
104-
assertThat(ArraySegments.of(a1).get(LE_INT, 0)).isEqualTo(10);
105-
assertThat(ArraySegments.of(a1).get(LE_INT, 4)).isEqualTo(20);
106-
assertThat(ArraySegments.of(a1).get(LE_INT, 8)).isEqualTo(10);
105+
assertThat(ArraySegments.of(a1, Arena.ofAuto()).get(LE_INT, 0)).isEqualTo(10);
106+
assertThat(ArraySegments.of(a1, Arena.ofAuto()).get(LE_INT, 4)).isEqualTo(20);
107+
assertThat(ArraySegments.of(a1, Arena.ofAuto()).get(LE_INT, 8)).isEqualTo(10);
107108
}
108109

109110
assertThat(iter.hasNext()).isTrue();
110111
try (Chunk c2 = iter.next()) {
111112
Array a2 = c2.columns().get("category");
112113
assertThat(a2.length()).isEqualTo(4L);
113-
assertThat(ArraySegments.of(a2).get(LE_INT, 0)).isEqualTo(30);
114-
assertThat(ArraySegments.of(a2).get(LE_INT, 4)).isEqualTo(10);
115-
assertThat(ArraySegments.of(a2).get(LE_INT, 8)).isEqualTo(20);
116-
assertThat(ArraySegments.of(a2).get(LE_INT, 12)).isEqualTo(30);
114+
assertThat(ArraySegments.of(a2, Arena.ofAuto()).get(LE_INT, 0)).isEqualTo(30);
115+
assertThat(ArraySegments.of(a2, Arena.ofAuto()).get(LE_INT, 4)).isEqualTo(10);
116+
assertThat(ArraySegments.of(a2, Arena.ofAuto()).get(LE_INT, 8)).isEqualTo(20);
117+
assertThat(ArraySegments.of(a2, Arena.ofAuto()).get(LE_INT, 12)).isEqualTo(30);
117118
}
118119

119120
assertThat(iter.hasNext()).isFalse();

writer/src/test/java/io/github/dfa1/vortex/writer/encode/BitpackedConstantPatchesBroadcastTest.java

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -73,7 +73,7 @@ void bitpackedDecode_withConstantPatchesValues_broadcastsValueAcrossPatches() {
7373

7474
// Then
7575
assertThat(result.length()).isEqualTo(n);
76-
MemorySegment data = ArraySegments.of(result);
76+
MemorySegment data = ArraySegments.of(result, Arena.ofAuto());
7777
assertThat(data.getAtIndex(PTypeIO.LE_LONG, 2)).isEqualTo(constantPatchValue);
7878
for (long i = 0; i < n; i++) {
7979
if (i == 2) {

0 commit comments

Comments
 (0)