@@ -343,6 +343,55 @@ TEST_F(AvroFileBatchReaderTest, TestGetPreviousBatchFirstRowNumber) {
343343 ASSERT_TRUE (BatchReader::IsEofBatch (batch5));
344344}
345345
346+ TEST_F (AvroFileBatchReaderTest, TestSetReadSchemaResetsReaderToFirstRow) {
347+ std::string file_path = PathUtil::JoinPath (dir_->Str (), " file.avro" );
348+
349+ arrow::FieldVector fields = {
350+ arrow::field (" f0" , arrow::int32 ()),
351+ arrow::field (" f1" , arrow::int32 ()),
352+ };
353+ auto file_data_type = arrow::struct_ (fields);
354+ auto src_array = arrow::ipc::internal::json::ArrayFromJSON (file_data_type, R"( [
355+ [1, 10],
356+ [2, 20],
357+ [3, 30],
358+ [4, 40]
359+ ])" )
360+ .ValueOrDie ();
361+ WriteData (src_array, file_path, /* compression=*/ " null" );
362+
363+ ASSERT_OK_AND_ASSIGN (auto reader_builder, file_format_->CreateReaderBuilder (/* batch_size=*/ 2 ));
364+ ASSERT_OK_AND_ASSIGN (std::shared_ptr<InputStream> in, fs_->Open (file_path));
365+ ASSERT_OK_AND_ASSIGN (auto reader, reader_builder->Build (in));
366+
367+ ASSERT_OK_AND_ASSIGN (auto first_batch, reader->NextBatch ());
368+ ASSERT_EQ (0 , reader->GetPreviousBatchFirstRowNumber ().value ());
369+ auto first_array =
370+ arrow::ImportArray (first_batch.first .get (), first_batch.second .get ()).ValueOrDie ();
371+ ASSERT_TRUE (first_array->Equals (src_array->Slice (0 , 2 ))) << first_array->ToString ();
372+
373+ auto read_schema = arrow::schema ({arrow::field (" f1" , arrow::int32 ())});
374+ std::unique_ptr<ArrowSchema> c_schema = std::make_unique<ArrowSchema>();
375+ ASSERT_TRUE (arrow::ExportSchema (*read_schema, c_schema.get ()).ok ());
376+ ASSERT_OK (reader->SetReadSchema (c_schema.get (), /* predicate=*/ nullptr ,
377+ /* selection_bitmap=*/ std::nullopt ));
378+ ASSERT_EQ (std::numeric_limits<uint64_t >::max (),
379+ reader->GetPreviousBatchFirstRowNumber ().value ());
380+
381+ ASSERT_OK_AND_ASSIGN (auto projected_batch, reader->NextBatch ());
382+ ASSERT_EQ (0 , reader->GetPreviousBatchFirstRowNumber ().value ());
383+ auto projected_array =
384+ arrow::ImportArray (projected_batch.first .get (), projected_batch.second .get ()).ValueOrDie ();
385+ auto expected_projected_array = arrow::ipc::internal::json::ArrayFromJSON (
386+ arrow::struct_ ({arrow::field (" f1" , arrow::int32 ())}),
387+ R"( [
388+ [10],
389+ [20]
390+ ])" )
391+ .ValueOrDie ();
392+ ASSERT_TRUE (projected_array->Equals (expected_projected_array)) << projected_array->ToString ();
393+ }
394+
346395TEST_F (AvroFileBatchReaderTest, TestGetNumberOfRows) {
347396 std::string file_path = PathUtil::JoinPath (dir_->Str (), " file.avro" );
348397
0 commit comments