diff --git a/vector/src/main/java/org/apache/arrow/vector/dictionary/DictionaryEncoder.java b/vector/src/main/java/org/apache/arrow/vector/dictionary/DictionaryEncoder.java index 4af1a8693f..a6d490c805 100644 --- a/vector/src/main/java/org/apache/arrow/vector/dictionary/DictionaryEncoder.java +++ b/vector/src/main/java/org/apache/arrow/vector/dictionary/DictionaryEncoder.java @@ -164,12 +164,12 @@ static void retrieveIndexVector( BaseIntVector indices, TransferPair transfer, int dictionaryCount, int start, int end) { for (int i = start; i < end; i++) { if (!indices.isNull(i)) { - int indexAsInt = (int) indices.getValueAsLong(i); - if (indexAsInt > dictionaryCount) { + long index = indices.getValueAsLong(i); + if (index < 0 || index >= dictionaryCount) { throw new IllegalArgumentException( - "Provided dictionary does not contain value for index " + indexAsInt); + "Provided dictionary does not contain value for index " + index); } - transfer.copyValueSafe(indexAsInt, i); + transfer.copyValueSafe((int) index, i); } } } diff --git a/vector/src/main/java/org/apache/arrow/vector/dictionary/StructSubfieldEncoder.java b/vector/src/main/java/org/apache/arrow/vector/dictionary/StructSubfieldEncoder.java index 8ff152fb1c..b183bc84c4 100644 --- a/vector/src/main/java/org/apache/arrow/vector/dictionary/StructSubfieldEncoder.java +++ b/vector/src/main/java/org/apache/arrow/vector/dictionary/StructSubfieldEncoder.java @@ -207,7 +207,8 @@ public static StructVector decode( TransferPair transfer = dictionary.getVector().makeTransferPair(decodedChildVector); BaseIntVector indices = (BaseIntVector) childVector; - DictionaryEncoder.retrieveIndexVector(indices, transfer, valueCount, 0, valueCount); + DictionaryEncoder.retrieveIndexVector( + indices, transfer, dictionary.getVector().getValueCount(), 0, valueCount); } } diff --git a/vector/src/test/java/org/apache/arrow/vector/TestDictionaryVector.java b/vector/src/test/java/org/apache/arrow/vector/TestDictionaryVector.java index 0945919b91..65c04433e5 100644 --- a/vector/src/test/java/org/apache/arrow/vector/TestDictionaryVector.java +++ b/vector/src/test/java/org/apache/arrow/vector/TestDictionaryVector.java @@ -21,6 +21,7 @@ import static org.junit.jupiter.api.Assertions.assertArrayEquals; import static org.junit.jupiter.api.Assertions.assertEquals; import static org.junit.jupiter.api.Assertions.assertFalse; +import static org.junit.jupiter.api.Assertions.assertThrows; import static org.junit.jupiter.api.Assertions.assertTrue; import static org.junit.jupiter.api.Assertions.fail; @@ -942,6 +943,61 @@ public void testNoMemoryLeak() { assertEquals(0, allocator.getAllocatedMemory(), "decode memory leak"); } + @Test + public void testDecodeRejectsDictionaryIndicesOutsideBounds() { + try (final IntVector indices = newVector(IntVector.class, "", Types.MinorType.INT, allocator); + final VarCharVector dictionaryVector = newVarCharVector("dict", allocator)) { + setVector(dictionaryVector, zero, one); + Dictionary dictionary = + new Dictionary(dictionaryVector, new DictionaryEncoding(1L, false, null)); + + setVector(indices, dictionaryVector.getValueCount()); + IllegalArgumentException upperBoundException = + assertThrows( + IllegalArgumentException.class, + () -> DictionaryEncoder.decode(indices, dictionary, allocator)); + assertEquals( + "Provided dictionary does not contain value for index 2", + upperBoundException.getMessage()); + + setVector(indices, -1); + IllegalArgumentException negativeException = + assertThrows( + IllegalArgumentException.class, + () -> DictionaryEncoder.decode(indices, dictionary, allocator)); + assertEquals( + "Provided dictionary does not contain value for index -1", + negativeException.getMessage()); + } + assertEquals(0, allocator.getAllocatedMemory(), "decode memory leak"); + } + + @Test + public void testDecodeRejectsBigIntDictionaryIndexOutsideBounds() { + try (final BigIntVector indices = new BigIntVector("indices", allocator); + final VarCharVector dictionaryVector = newVarCharVector("dict", allocator)) { + setVector(dictionaryVector, zero, one); + Dictionary dictionary = + new Dictionary(dictionaryVector, new DictionaryEncoding(1L, false, null)); + + setVector(indices, 1L); + try (ValueVector decoded = DictionaryEncoder.decode(indices, dictionary, allocator)) { + assertEquals(new Text("bar"), decoded.getObject(0)); + } + + long largeIndex = 1L << 32; + setVector(indices, largeIndex); + IllegalArgumentException exception = + assertThrows( + IllegalArgumentException.class, + () -> DictionaryEncoder.decode(indices, dictionary, allocator)); + assertEquals( + "Provided dictionary does not contain value for index " + largeIndex, + exception.getMessage()); + } + assertEquals(0, allocator.getAllocatedMemory(), "decode memory leak"); + } + @Test public void testListNoMemoryLeak() { // Create a new value vector @@ -1053,7 +1109,7 @@ public void testStructNoMemoryLeak() { NullableStructWriter writer = indices.getWriter(); writer.allocate(); writer.start(); - writer.integer("f0").writeInt(1); + writer.integer("f0").writeInt(0); writer.integer("f1").writeInt(3); writer.end(); writer.setValueCount(1); @@ -1067,6 +1123,59 @@ public void testStructNoMemoryLeak() { assertEquals(0, allocator.getAllocatedMemory(), "struct decode memory leak"); } + @Test + public void testStructDecodeUsesDictionaryValueCount() { + try (final StructVector validIndices = StructVector.empty("valid", allocator); + final StructVector outOfRangeIndices = StructVector.empty("outOfRange", allocator); + final VarCharVector dictionaryVector = new VarCharVector("f0", allocator)) { + + setVector( + dictionaryVector, + "aa".getBytes(StandardCharsets.UTF_8), + "bb".getBytes(StandardCharsets.UTF_8)); + + DictionaryProvider.MapDictionaryProvider provider = + new DictionaryProvider.MapDictionaryProvider(); + Dictionary dictionary = + new Dictionary(dictionaryVector, new DictionaryEncoding(1L, false, null)); + provider.put(dictionary); + + ArrowType int32 = new ArrowType.Int(32, true); + FieldType indexFieldType = new FieldType(true, int32, dictionary.getEncoding()); + validIndices.addOrGet("f0", indexFieldType, IntVector.class); + outOfRangeIndices.addOrGet("f0", indexFieldType, IntVector.class); + + NullableStructWriter validWriter = validIndices.getWriter(); + validWriter.allocate(); + validWriter.start(); + validWriter.integer("f0").writeInt(1); + validWriter.end(); + validIndices.setValueCount(1); + + try (StructVector decoded = StructSubfieldEncoder.decode(validIndices, provider, allocator)) { + assertArrayEquals( + new Object[] {new Text("bb")}, convertMapValuesToArray(decoded.getObject(0))); + } + + NullableStructWriter outOfRangeWriter = outOfRangeIndices.getWriter(); + outOfRangeWriter.allocate(); + for (int i = 0; i < 5; i++) { + outOfRangeWriter.start(); + outOfRangeWriter.integer("f0").writeInt(i == 0 ? 2 : 0); + outOfRangeWriter.end(); + } + outOfRangeIndices.setValueCount(5); + + IllegalArgumentException exception = + assertThrows( + IllegalArgumentException.class, + () -> StructSubfieldEncoder.decode(outOfRangeIndices, provider, allocator)); + assertEquals( + "Provided dictionary does not contain value for index 2", exception.getMessage()); + } + assertEquals(0, allocator.getAllocatedMemory(), "struct decode memory leak"); + } + private void testDictionary( Dictionary dictionary, ToIntBiFunction valGetter) { try (VarCharVector vector = new VarCharVector("vector", allocator)) {