This is an automated email from the ASF dual-hosted git repository. yiguolei pushed a commit to branch branch-4.2 in repository https://gitbox.apache.org/repos/asf/doris.git
commit ffa16d4b68de5cba085f13b039fa9425b6eb89e4 Author: Gabriel <[email protected]> AuthorDate: Tue Sep 29 16:35:21 2026 +0800 [fix](arrow) Validate Flight result arrays before publishing batches (#68596) ### What problem does this PR solve? Arrow Flight SQL can return a string array containing invalid UTF-8, for example from `SELECT unhex(encoded) FROM input_table` when a row contains the hex string `84`. Arrow string builders accept the bytes, but clients such as PyArrow fail when decoding the result. Validate each completed array in `ArrowFlightArrowBlockConvertor` before publishing the batch. Reuse Arrow's recursive validation to cover nested strings and per-value UTF-8 boundaries. Return an invalid-argument error identifying the column ordinal, field name and Arrow validation failure. Binary values retain their byte representation, and NULL payloads are handled through the serialized validity bitmap. Add six BE unit tests and a Flight SQL regression suite covering malformed encodings, separate rows whose concatenation is valid UTF-8, subsequent batches, nested arrays/structs/map keys and values, large strings, valid Unicode, NULLs and binary results. ### Validation - Reproduced four failing tests before the fix; the two valid-data/NULL tests passed. - ASAN: all 46 tests in `ArrowBlockConvertorTest` and `DataTypeSerDeArrowTest` passed after the fix. - clang-format 16 passed for all three affected C++ files. - Updated the Flight regression suite to bypass constant folding for the constant-input check, reuse the configured Flight connection, assert valid results directly, and clean up in a finally block. - Groovy 4.0.19: six mock JDBC execution checks passed, covering TLS on/off, absent expected errors, unrelated errors, and cleanup. These checks do not replace live Flight SQL execution; the updated regression is pending CI. - Re-ran the six focused Flight BE tests under ASAN after the regression-only update: all passed. The local BE unit-test build selected the relevant suites and their test support files; the build configuration was restored and is not part of this PR. ### Release note Arrow Flight SQL now rejects invalid UTF-8 string results with a descriptive server-side error instead of returning malformed Arrow data. ### Check List (For Author) - Test - [x] Regression test added (live execution pending CI) - [x] Unit Test - Behavior changed: - [x] Yes. Invalid Arrow Flight string payloads fail on the server before the affected batch is returned. - Does this need documentation? - [x] No. ### Check List (For Reviewer who merge this PR) - [ ] Confirm the release note - [ ] Confirm test cases - [ ] Confirm document - [ ] Add branch pick label --- be/src/format/arrow/arrow_block_convertor.cpp | 19 +++ be/src/format/arrow/arrow_block_convertor.h | 10 +- .../format/arrow/arrow_block_convertor_test.cpp | 133 +++++++++++++++++++++ .../test_flight_utf8_validation.groovy | 71 +++++++++++ 4 files changed, 230 insertions(+), 3 deletions(-) diff --git a/be/src/format/arrow/arrow_block_convertor.cpp b/be/src/format/arrow/arrow_block_convertor.cpp index 0046ad3dacc..b62ca98cde9 100644 --- a/be/src/format/arrow/arrow_block_convertor.cpp +++ b/be/src/format/arrow/arrow_block_convertor.cpp @@ -341,6 +341,25 @@ Status ArrowBlockConvertor::init() { return Status::OK(); } +Status ArrowFlightArrowBlockConvertor::convert_to_arrow(const Block& block, arrow::MemoryPool* pool, + std::shared_ptr<arrow::RecordBatch>* result, + size_t start_row, size_t end_row) const { + std::shared_ptr<arrow::RecordBatch> batch; + RETURN_IF_ERROR(ArrowBlockConvertor::convert_to_arrow(block, pool, &batch, start_row, end_row)); + // String builders accept arbitrary bytes, but Flight UTF-8 values must be valid per row, + // including nested children. Validate before publishing the batch to the reader. + for (int i = 0; i < batch->num_columns(); ++i) { + auto status = batch->column(i)->ValidateFull(); + if (!status.ok()) { + return Status::InvalidArgument("Invalid Arrow Flight result in column {} ('{}'): {}", + i + 1, batch->schema()->field(i)->name(), + status.ToString()); + } + } + *result = std::move(batch); + return Status::OK(); +} + Status DorisArrowBlockConvertor::init() { if (_arrow_schema == nullptr) { // cctz names fixed offsets as "Fixed/UTC+HH:MM:SS", which is not the Arrow diff --git a/be/src/format/arrow/arrow_block_convertor.h b/be/src/format/arrow/arrow_block_convertor.h index ec1083505a9..c923507786d 100644 --- a/be/src/format/arrow/arrow_block_convertor.h +++ b/be/src/format/arrow/arrow_block_convertor.h @@ -58,9 +58,9 @@ public: virtual Status init(); const std::shared_ptr<arrow::Schema>& arrow_schema() const { return _arrow_schema; } - Status convert_to_arrow(const Block& block, arrow::MemoryPool* pool, - std::shared_ptr<arrow::RecordBatch>* result, size_t start_row = 0, - size_t end_row = 0) const; + virtual Status convert_to_arrow(const Block& block, arrow::MemoryPool* pool, + std::shared_ptr<arrow::RecordBatch>* result, + size_t start_row = 0, size_t end_row = 0) const; virtual Status convert_from_arrow(const std::shared_ptr<arrow::RecordBatch>& batch, const DataTypes& types, Block* block) const; @@ -115,6 +115,10 @@ private: class ArrowFlightArrowBlockConvertor final : public DorisArrowBlockConvertor { public: using DorisArrowBlockConvertor::DorisArrowBlockConvertor; + + Status convert_to_arrow(const Block& block, arrow::MemoryPool* pool, + std::shared_ptr<arrow::RecordBatch>* result, size_t start_row = 0, + size_t end_row = 0) const override; }; class PythonArrowBlockConvertor final : public DorisArrowBlockConvertor { diff --git a/be/test/format/arrow/arrow_block_convertor_test.cpp b/be/test/format/arrow/arrow_block_convertor_test.cpp index f3db16d59de..b02b60cb7a7 100644 --- a/be/test/format/arrow/arrow_block_convertor_test.cpp +++ b/be/test/format/arrow/arrow_block_convertor_test.cpp @@ -23,6 +23,7 @@ #include <arrow/ipc/api.h> #include <gtest/gtest.h> +#include "core/column/column_nullable.h" #include "core/column/column_vector.h" #include "core/data_type/data_type_array.h" #include "core/data_type/data_type_factory.hpp" @@ -289,4 +290,136 @@ TEST_F(ArrowBlockConvertorTest, TableConvertersRejectMismatchedNestedSchemas) { } } +TEST_F(ArrowBlockConvertorTest, FlightRejectsInvalidUtf8BeforeReturningBatch) { + auto type = DataTypeFactory::instance().create_data_type(TYPE_STRING, false); + for (const std::string& value : + {std::string("\x84"), std::string("\xc0\xaf"), std::string("\xed\xa0\x80"), + std::string("\xf4\x90\x80\x80")}) { + auto column = type->create_column(); + column->insert(Field::create_field<TYPE_STRING>(value)); + Block block {{std::move(column), type, "payload"}}; + ArrowFlightArrowBlockConvertor flight(block, "UTC", cctz::utc_time_zone()); + ASSERT_TRUE(flight.init().ok()); + const ArrowBlockConvertor& converter = flight; + std::shared_ptr<arrow::RecordBatch> batch; + const auto status = converter.convert_to_arrow(block, arrow::default_memory_pool(), &batch); + EXPECT_EQ(ErrorCode::INVALID_ARGUMENT, status.code()) << status; + EXPECT_NE(std::string::npos, status.to_string().find("payload")); + EXPECT_NE(std::string::npos, status.to_string().find("UTF8")); + EXPECT_EQ(nullptr, batch); + + DorisArrowBlockConvertor ordinary(block, "UTC", cctz::utc_time_zone()); + ASSERT_TRUE(ordinary.init().ok()); + EXPECT_TRUE(ordinary.convert_to_arrow(block, arrow::default_memory_pool(), &batch).ok()); + } +} + +TEST_F(ArrowBlockConvertorTest, FlightChecksEachStringAndEveryBatch) { + auto type = DataTypeFactory::instance().create_data_type(TYPE_STRING, false); + auto column = type->create_column(); + for (const std::string& value : + {std::string("valid"), std::string("\xc2"), std::string("\xa2")}) { + column->insert(Field::create_field<TYPE_STRING>(value)); + } + Block block {{std::move(column), type, "payload"}}; + ArrowFlightArrowBlockConvertor converter(block, "UTC", cctz::utc_time_zone()); + ASSERT_TRUE(converter.init().ok()); + std::shared_ptr<arrow::RecordBatch> batch; + ASSERT_TRUE(converter.convert_to_arrow(block, arrow::default_memory_pool(), &batch, 0, 1).ok()); + ASSERT_EQ(1, batch->num_rows()); + batch.reset(); + // Adjacent invalid strings form valid UTF-8 when concatenated; row boundaries matter. + EXPECT_FALSE( + converter.convert_to_arrow(block, arrow::default_memory_pool(), &batch, 1, 3).ok()); + EXPECT_EQ(nullptr, batch); +} + +TEST_F(ArrowBlockConvertorTest, FlightPreservesValidTextNullsAndBinary) { + auto text = make_nullable(DataTypeFactory::instance().create_data_type(TYPE_STRING, false)); + auto strings = text->create_column(); + strings->insert(Field::create_field<TYPE_STRING>(std::string("\xe4\xb8\xad\xf0\x9f\x98\x80"))); + strings->insert(Field::create_field<TYPE_STRING>(std::string())); + strings->insert(Field::create_field<TYPE_STRING>(std::string("a\0b", 3))); + strings->insert_default(); + auto binary = DataTypeFactory::instance().create_data_type(TYPE_VARBINARY, false); + auto bytes = binary->create_column(); + for (int i = 0; i < 4; ++i) { + bytes->insert_data("\x84\0\xff", 3); + } + Block block {{std::move(strings), text, "text"}, {std::move(bytes), binary, "binary"}}; + ArrowFlightArrowBlockConvertor converter(block, "UTC", cctz::utc_time_zone()); + ASSERT_TRUE(converter.init().ok()); + std::shared_ptr<arrow::RecordBatch> batch; + ASSERT_TRUE(converter.convert_to_arrow(block, arrow::default_memory_pool(), &batch).ok()); + ASSERT_TRUE(batch->ValidateFull().ok()); + const auto& values = static_cast<const arrow::StringArray&>(*batch->column(0)); + EXPECT_EQ("\xe4\xb8\xad\xf0\x9f\x98\x80", values.GetString(0)); + EXPECT_EQ("", values.GetString(1)); + EXPECT_EQ(std::string("a\0b", 3), values.GetString(2)); + EXPECT_TRUE(values.IsNull(3)); + ASSERT_EQ(arrow::Type::BINARY, batch->column(1)->type_id()); + const auto& binary_values = static_cast<const arrow::BinaryArray&>(*batch->column(1)); + EXPECT_EQ(std::string("\x84\0\xff", 3), binary_values.GetString(0)); +} + +TEST_F(ArrowBlockConvertorTest, FlightRejectsNestedInvalidUtf8) { + auto text = make_nullable(DataTypeFactory::instance().create_data_type(TYPE_STRING, false)); + const auto invalid = Field::create_field<TYPE_STRING>(std::string("\x84")); + const auto valid = Field::create_field<TYPE_STRING>(std::string("key")); + DataTypes types {std::make_shared<DataTypeArray>(text), + std::make_shared<DataTypeStruct>(DataTypes {text}, Strings {"child"}), + std::make_shared<DataTypeMap>(text, text), + std::make_shared<DataTypeMap>(text, text)}; + FieldVector fields { + Field::create_field<TYPE_ARRAY>(Array {invalid}), + Field::create_field<TYPE_STRUCT>(Struct {invalid}), + Field::create_field<TYPE_MAP>(Map {Field::create_field<TYPE_ARRAY>(Array {valid}), + Field::create_field<TYPE_ARRAY>(Array {invalid})}), + Field::create_field<TYPE_MAP>(Map {Field::create_field<TYPE_ARRAY>(Array {invalid}), + Field::create_field<TYPE_ARRAY>(Array {valid})})}; + for (size_t i = 0; i < types.size(); ++i) { + SCOPED_TRACE(types[i]->get_name()); + auto column = types[i]->create_column(); + column->insert(fields[i]); + Block block {{std::move(column), types[i], "nested"}}; + ArrowFlightArrowBlockConvertor converter(block, "UTC", cctz::utc_time_zone()); + ASSERT_TRUE(converter.init().ok()); + std::shared_ptr<arrow::RecordBatch> batch; + const auto status = converter.convert_to_arrow(block, arrow::default_memory_pool(), &batch); + EXPECT_EQ(ErrorCode::INVALID_ARGUMENT, status.code()) << status; + EXPECT_NE(std::string::npos, status.to_string().find("nested")); + EXPECT_EQ(nullptr, batch); + } +} + +TEST_F(ArrowBlockConvertorTest, FlightRejectsInvalidLargeString) { + auto type = DataTypeFactory::instance().create_data_type(TYPE_STRING, false); + auto column = type->create_column(); + column->insert(Field::create_field<TYPE_STRING>(std::string("\x84"))); + Block block {{std::move(column), type, "large_text"}}; + auto schema = arrow::schema({arrow::field("large_text", arrow::large_utf8(), false)}); + ArrowFlightArrowBlockConvertor converter(schema, cctz::utc_time_zone()); + std::shared_ptr<arrow::RecordBatch> batch; + const auto status = converter.convert_to_arrow(block, arrow::default_memory_pool(), &batch); + EXPECT_EQ(ErrorCode::INVALID_ARGUMENT, status.code()) << status; + EXPECT_NE(std::string::npos, status.to_string().find("large_text")); + EXPECT_EQ(nullptr, batch); +} + +TEST_F(ArrowBlockConvertorTest, FlightIgnoresBytesMaskedByNull) { + auto text = DataTypeFactory::instance().create_data_type(TYPE_STRING, false); + auto values = text->create_column(); + values->insert_data("\x84", 1); + auto nulls = ColumnUInt8::create(); + nulls->insert_value(1); + Block block {{ColumnNullable::create(std::move(values), std::move(nulls)), make_nullable(text), + "nullable_text"}}; + ArrowFlightArrowBlockConvertor converter(block, "UTC", cctz::utc_time_zone()); + ASSERT_TRUE(converter.init().ok()); + std::shared_ptr<arrow::RecordBatch> batch; + ASSERT_TRUE(converter.convert_to_arrow(block, arrow::default_memory_pool(), &batch).ok()); + ASSERT_TRUE(batch->ValidateFull().ok()); + EXPECT_TRUE(batch->column(0)->IsNull(0)); +} + } // namespace doris diff --git a/regression-test/suites/arrow_flight_sql_p0/test_flight_utf8_validation.groovy b/regression-test/suites/arrow_flight_sql_p0/test_flight_utf8_validation.groovy new file mode 100644 index 00000000000..f19312b70d3 --- /dev/null +++ b/regression-test/suites/arrow_flight_sql_p0/test_flight_utf8_validation.groovy @@ -0,0 +1,71 @@ +// Licensed to the Apache Software Foundation (ASF) under one +// or more contributor license agreements. See the NOTICE file +// distributed with this work for additional information +// regarding copyright ownership. The ASF licenses this file +// to you under the Apache License, Version 2.0 (the +// "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, +// software distributed under the License is distributed on an +// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +// KIND, either express or implied. See the License for the +// specific language governing permissions and limitations +// under the License. + +suite("test_flight_utf8_validation", "arrow_flight_sql") { + // Reuse Flight credentials and avoid the TLS-dependent early return in connect(). + def flight = context.getArrowFlightSqlConnection() + def input = "${context.dbName}.flight_utf8_input" + def expectInvalidUtf8 = { String query -> + try { + arrow_flight_sql(query) + assertTrue(false, "Expected an invalid UTF-8 Flight result error") + } catch (Exception error) { + assertTrue(error.toString().contains("Invalid UTF8"), error.toString()) + } + } + arrow_flight_sql "DROP TABLE IF EXISTS ${input}" + arrow_flight_sql """CREATE TABLE ${input} (id INT, encoded STRING) + DUPLICATE KEY(id) DISTRIBUTED BY HASH(id) BUCKETS 1 + PROPERTIES ("replication_num" = "1")""" + try { + arrow_flight_sql """INSERT INTO ${input} VALUES + (1, '616263'), (2, 'E4B8ADF09F9880'), (3, ''), (4, NULL), + (5, '84'), (6, 'C0AF'), (7, 'EDA080'), (8, 'F4908080'), + (9, 'C2'), (10, 'A2')""" + for (def id : [5, 6, 7, 8]) { + expectInvalidUtf8("SELECT unhex(encoded) AS payload FROM ${input} WHERE id = ${id}") + } + // Constant folding can replace invalid bytes with valid text before Flight conversion. + // Keep a table scan and skip folding so the constant is evaluated and validated on the BE. + expectInvalidUtf8("""SELECT /*+ SET_VAR(debug_skip_fold_constant=true) */ + unhex('84') AS payload FROM ${input} WHERE id = 1""") + // These bytes are valid only when concatenated; each row must be valid independently. + expectInvalidUtf8("SELECT unhex(encoded) AS payload FROM ${input} WHERE id IN (9, 10) ORDER BY id") + for (def expression : ["array(unhex(encoded))", + "named_struct('child', unhex(encoded))", + "map('key', unhex(encoded))", + "map(unhex(encoded), 'value')"]) { + expectInvalidUtf8("SELECT ${expression} AS nested FROM ${input} WHERE id = 5") + } + // Reaching an invalid value after valid rows must fail rather than publish malformed text. + expectInvalidUtf8("SELECT unhex(encoded) AS payload FROM ${input} ORDER BY id") + assertEquals([['abc'], ['中😀'], [''], [null]], + arrow_flight_sql("SELECT unhex(encoded) FROM ${input} WHERE id <= 4 ORDER BY id")) + // Binary payloads may contain arbitrary bytes and must keep their binary Arrow type. + flight.createStatement().withCloseable { statement -> + statement.executeQuery("SELECT to_binary(encoded) AS payload FROM ${input} WHERE id = 5") + .withCloseable { rows -> + assertTrue(rows.next()) + assertTrue(rows.getObject(1) instanceof byte[]) + assertEquals([0x84], rows.getBytes(1).collect { it & 0xff }) + assertFalse(rows.next()) + } + } + } finally { + arrow_flight_sql "DROP TABLE IF EXISTS ${input}" + } +} --------------------------------------------------------------------- To unsubscribe, e-mail: [email protected] For additional commands, e-mail: [email protected]
