diff --git a/parquet/pqarrow/file_reader.go b/parquet/pqarrow/file_reader.go index 37b847fc..ae659d22 100644 --- a/parquet/pqarrow/file_reader.go +++ b/parquet/pqarrow/file_reader.go @@ -219,6 +219,13 @@ func (fr *FileReader) allRowGroupFactory() itrFactory { // // IncludedLeaves and RowGroups are used to specify precisely which leaf indexes and row groups to read a subset of. func (fr *FileReader) GetFieldReader(ctx context.Context, i int, includedLeaves map[int]bool, rowGroups []int) (*ColumnReader, error) { + if i < 0 || i >= len(fr.Manifest.Fields) { + return nil, fmt.Errorf("%w: invalid field index chosen %d, there are only %d fields", arrow.ErrIndex, i, len(fr.Manifest.Fields)) + } + if err := fr.checkRowGroups(rowGroups); err != nil { + return nil, err + } + ctx = context.WithValue(ctx, rdrCtxKey{}, readerCtx{ rdr: fr.rdr, mem: fr.mem, @@ -287,6 +294,10 @@ func (fr *FileReader) RowGroup(idx int) RowGroupReader { // ReadColumn reads data to create a chunked array only from the requested row groups. func (fr *FileReader) ReadColumn(rowGroups []int, rdr *ColumnReader) (*arrow.Chunked, error) { + if err := fr.checkRowGroups(rowGroups); err != nil { + return nil, err + } + recs := int64(0) for _, rg := range rowGroups { recs += fr.rdr.MetaData().RowGroups[rg].GetNumRows() @@ -312,7 +323,7 @@ func (fr *FileReader) ReadTable(ctx context.Context) (arrow.Table, error) { func (fr *FileReader) checkCols(indices []int) (err error) { for _, col := range indices { if col < 0 || col >= fr.rdr.MetaData().Schema.NumColumns() { - err = fmt.Errorf("invalid column index specified %d out of %d", col, fr.rdr.MetaData().Schema.NumColumns()) + err = fmt.Errorf("%w: invalid column index specified %d out of %d", arrow.ErrIndex, col, fr.rdr.MetaData().Schema.NumColumns()) break } } @@ -322,7 +333,7 @@ func (fr *FileReader) checkCols(indices []int) (err error) { func (fr *FileReader) checkRowGroups(indices []int) (err error) { for _, rg := range indices { if rg < 0 || rg >= fr.rdr.NumRowGroups() { - err = fmt.Errorf("invalid row group specified: %d, file only has %d row groups", rg, fr.rdr.NumRowGroups()) + err = fmt.Errorf("%w: invalid row group specified: %d, file only has %d row groups", arrow.ErrIndex, rg, fr.rdr.NumRowGroups()) break } } @@ -452,7 +463,7 @@ func (fr *FileReader) ReadRowGroups(ctx context.Context, indices, rowGroups []in func (fr *FileReader) getColumnReader(ctx context.Context, i int, colFactory itrFactory) (*ColumnReader, error) { if i < 0 || i >= len(fr.Manifest.Fields) { - return nil, fmt.Errorf("invalid column index chosen %d, there are only %d columns", i, len(fr.Manifest.Fields)) + return nil, fmt.Errorf("%w: invalid column index chosen %d, there are only %d columns", arrow.ErrIndex, i, len(fr.Manifest.Fields)) } ctx = context.WithValue(ctx, rdrCtxKey{}, readerCtx{ diff --git a/parquet/pqarrow/file_reader_test.go b/parquet/pqarrow/file_reader_test.go index 45e0a4f3..16c0c954 100644 --- a/parquet/pqarrow/file_reader_test.go +++ b/parquet/pqarrow/file_reader_test.go @@ -606,6 +606,51 @@ func TestFileReaderColumnChunkBoundsErrors(t *testing.T) { } } +func TestFileReaderIndexValidation(t *testing.T) { + schema := arrow.NewSchema([]arrow.Field{{Name: "value", Type: arrow.PrimitiveTypes.Int32}}, nil) + record, _, err := array.RecordFromJSON(memory.DefaultAllocator, schema, + strings.NewReader(`[{"value": 1}]`)) + require.NoError(t, err) + defer record.Release() + + var buf bytes.Buffer + writer, err := pqarrow.NewFileWriter(schema, &buf, nil, pqarrow.DefaultWriterProps()) + require.NoError(t, err) + require.NoError(t, writer.Write(record)) + require.NoError(t, writer.Close()) + + fileReader, err := file.NewParquetReader(bytes.NewReader(buf.Bytes())) + require.NoError(t, err) + defer fileReader.Close() + + arrowReader, err := pqarrow.NewFileReader(fileReader, pqarrow.ArrowReadProperties{}, memory.DefaultAllocator) + require.NoError(t, err) + + _, err = arrowReader.GetFieldReader(context.Background(), -1, nil, []int{0}) + require.ErrorIs(t, err, arrow.ErrIndex) + _, err = arrowReader.GetFieldReader(context.Background(), 1, nil, []int{0}) + require.ErrorIs(t, err, arrow.ErrIndex) + _, err = arrowReader.GetFieldReader(context.Background(), 0, nil, []int{1}) + require.ErrorIs(t, err, arrow.ErrIndex) + _, err = arrowReader.GetFieldReader(context.Background(), 0, nil, []int{-1}) + require.ErrorIs(t, err, arrow.ErrIndex) + + fieldReader, err := arrowReader.GetFieldReader(context.Background(), 0, map[int]bool{0: true}, []int{0}) + require.NoError(t, err) + fieldReader.Release() + + _, err = arrowReader.GetColumn(context.Background(), -1) + require.ErrorIs(t, err, arrow.ErrIndex) + columnReader, err := arrowReader.GetColumn(context.Background(), 0) + require.NoError(t, err) + defer columnReader.Release() + _, err = arrowReader.ReadColumn([]int{1}, columnReader) + require.ErrorIs(t, err, arrow.ErrIndex) + chunked, err := arrowReader.ReadColumn([]int{0}, columnReader) + require.NoError(t, err) + chunked.Release() +} + func TestReadParquetFile(t *testing.T) { dir := os.Getenv("PARQUET_TEST_BAD_DATA") if dir == "" {