Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
17 changes: 14 additions & 3 deletions parquet/pqarrow/file_reader.go
Original file line number Diff line number Diff line change
Expand Up @@ -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))
Comment thread
zeroshade marked this conversation as resolved.
}
if err := fr.checkRowGroups(rowGroups); err != nil {
return nil, err
}

ctx = context.WithValue(ctx, rdrCtxKey{}, readerCtx{
rdr: fr.rdr,
mem: fr.mem,
Expand Down Expand Up @@ -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()
Expand All @@ -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
}
}
Expand All @@ -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
}
}
Expand Down Expand Up @@ -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{
Expand Down
45 changes: 45 additions & 0 deletions parquet/pqarrow/file_reader_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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 == "" {
Expand Down