diff --git a/be/src/format_v2/lance/lance_reader_helper.cpp b/be/src/format_v2/lance/lance_reader_helper.cpp new file mode 100644 index 00000000000000..446a4b40179e4d --- /dev/null +++ b/be/src/format_v2/lance/lance_reader_helper.cpp @@ -0,0 +1,321 @@ +// 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. + +#include "format_v2/lance/lance_reader_helper.h" + +#include +#include +#include +#include + +#include +#include + +#include "common/logging.h" +#include "core/data_type/data_type_array.h" +#include "core/data_type/data_type_factory.hpp" +#include "core/data_type/data_type_map.h" +#include "core/data_type/data_type_nothing.h" +#include "core/data_type/data_type_nullable.h" +#include "core/data_type/data_type_struct.h" + +namespace doris::format::lance { +namespace { + +constexpr std::string_view ARROW_EXTENSION_NAME = "ARROW:extension:name"; + +int arrow_time_precision(arrow::TimeUnit::type unit) { + switch (unit) { + case arrow::TimeUnit::SECOND: + return 0; + case arrow::TimeUnit::MILLI: + return 3; + case arrow::TimeUnit::MICRO: + case arrow::TimeUnit::NANO: + return 6; + } + return 6; +} + +Status check_arrow_field_semantics(const std::shared_ptr& field) { + if (field->HasMetadata()) { + const auto extension_name = field->metadata()->Get(ARROW_EXTENSION_NAME); + if (extension_name.ok() && !extension_name.ValueUnsafe().empty()) { + return Status::NotSupported( + "unsupported Lance Arrow extension type '{}' for field '{}'", + extension_name.ValueUnsafe(), field->name()); + } + } + if (field->type()->id() == arrow::Type::DICTIONARY) { + return Status::NotSupported("unsupported Lance Arrow dictionary type for field '{}': {}", + field->name(), field->type()->ToString()); + } + return Status::OK(); +} + +Status arrow_field_to_doris_type(const std::shared_ptr& field, + DataTypePtr* doris_type) { + RETURN_IF_ERROR(check_arrow_field_semantics(field)); + const auto& arrow_type = field->type(); + const auto nullable_primitive = [&](PrimitiveType type, int precision = 0, int scale = 0, + int len = -1) { + *doris_type = + DataTypeFactory::instance().create_data_type(type, true, precision, scale, len); + return Status::OK(); + }; + + switch (arrow_type->id()) { + case arrow::Type::BOOL: + return nullable_primitive(TYPE_BOOLEAN); + case arrow::Type::INT8: + return nullable_primitive(TYPE_TINYINT); + case arrow::Type::UINT8: + case arrow::Type::INT16: + return nullable_primitive(TYPE_SMALLINT); + case arrow::Type::UINT16: + case arrow::Type::INT32: + return nullable_primitive(TYPE_INT); + case arrow::Type::UINT32: + case arrow::Type::INT64: + return nullable_primitive(TYPE_BIGINT); + case arrow::Type::UINT64: + return nullable_primitive(TYPE_LARGEINT); + case arrow::Type::HALF_FLOAT: + case arrow::Type::FLOAT: + return nullable_primitive(TYPE_FLOAT); + case arrow::Type::DOUBLE: + return nullable_primitive(TYPE_DOUBLE); + case arrow::Type::STRING: + case arrow::Type::LARGE_STRING: + return nullable_primitive(TYPE_STRING); + case arrow::Type::BINARY: + case arrow::Type::LARGE_BINARY: + return nullable_primitive(TYPE_VARBINARY, 0, 0, std::numeric_limits::max()); + case arrow::Type::FIXED_SIZE_BINARY: { + const auto binary = std::static_pointer_cast(arrow_type); + return nullable_primitive(TYPE_VARBINARY, 0, 0, binary->byte_width()); + } + case arrow::Type::DATE32: + case arrow::Type::DATE64: + return nullable_primitive(TYPE_DATEV2); + case arrow::Type::TIME32: + case arrow::Type::TIME64: { + const auto time = std::static_pointer_cast(arrow_type); + return nullable_primitive(TYPE_TIMEV2, 0, arrow_time_precision(time->unit())); + } + case arrow::Type::TIMESTAMP: { + const auto timestamp = std::static_pointer_cast(arrow_type); + const auto doris_type = timestamp->timezone().empty() ? TYPE_DATETIMEV2 : TYPE_TIMESTAMPTZ; + return nullable_primitive(doris_type, 0, arrow_time_precision(timestamp->unit())); + } + case arrow::Type::DECIMAL128: + case arrow::Type::DECIMAL256: { + const auto decimal = std::static_pointer_cast(arrow_type); + const int precision = decimal->precision(); + const int scale = decimal->scale(); + if (precision <= 0 || precision > arrow::Decimal256Type::kMaxPrecision || scale < 0 || + scale > precision) { + return Status::NotSupported( + "unsupported Lance Arrow decimal type for field '{}': precision={}, scale={}", + field->name(), precision, scale); + } + const PrimitiveType doris_decimal_type = precision <= 9 ? TYPE_DECIMAL32 + : precision <= 18 ? TYPE_DECIMAL64 + : precision <= 38 ? TYPE_DECIMAL128I + : TYPE_DECIMAL256; + return nullable_primitive(doris_decimal_type, precision, scale); + } + case arrow::Type::LIST: + case arrow::Type::LARGE_LIST: + case arrow::Type::FIXED_SIZE_LIST: { + const auto list = std::static_pointer_cast(arrow_type); + DataTypePtr value_type; + RETURN_IF_ERROR(arrow_field_to_doris_type(list->value_field(), &value_type)); + *doris_type = make_nullable(std::make_shared(value_type)); + return Status::OK(); + } + case arrow::Type::MAP: { + const auto map = std::static_pointer_cast(arrow_type); + RETURN_IF_ERROR(check_arrow_field_semantics(map->value_field())); + DataTypePtr key_type; + DataTypePtr item_type; + RETURN_IF_ERROR(arrow_field_to_doris_type(map->key_field(), &key_type)); + RETURN_IF_ERROR(arrow_field_to_doris_type(map->item_field(), &item_type)); + *doris_type = make_nullable(std::make_shared(key_type, item_type)); + return Status::OK(); + } + case arrow::Type::STRUCT: { + const auto struct_type = std::static_pointer_cast(arrow_type); + DataTypes field_types; + Strings field_names; + field_types.reserve(struct_type->num_fields()); + field_names.reserve(struct_type->num_fields()); + for (const auto& child : struct_type->fields()) { + DataTypePtr field_type; + RETURN_IF_ERROR(arrow_field_to_doris_type(child, &field_type)); + field_types.emplace_back(std::move(field_type)); + field_names.emplace_back(child->name()); + } + *doris_type = make_nullable(std::make_shared(field_types, field_names)); + return Status::OK(); + } + default: + return Status::NotSupported("unsupported Lance Arrow type: {}", arrow_type->ToString()); + } +} + +} // namespace + +void LanceDatasetDeleter::operator()(LanceDataset* dataset) const { + lance_dataset_close(dataset); +} + +void LanceScannerDeleter::operator()(LanceScanner* scanner) const { + lance_scanner_close(scanner); +} + +void LanceBatchDeleter::operator()(LanceBatch* batch) const { + lance_batch_free(batch); +} + +size_t lance_vector_element_width(TVectorElementType::type type) { + switch (type) { + case TVectorElementType::FLOAT16: + return sizeof(uint16_t); + case TVectorElementType::FLOAT32: + return sizeof(float); + case TVectorElementType::FLOAT64: + return sizeof(double); + case TVectorElementType::UINT8: + case TVectorElementType::INT8: + return sizeof(uint8_t); + } + return 0; +} + +Status parse_fragment_ids(const TLanceFileDesc& lance_params, std::vector* fragment_ids) { + DORIS_CHECK(fragment_ids != nullptr); + fragment_ids->clear(); + if (!lance_params.__isset.fragment_ids || lance_params.fragment_ids.empty()) { + return Status::OK(); + } + fragment_ids->reserve(lance_params.fragment_ids.size()); + for (const auto fragment_id : lance_params.fragment_ids) { + if (fragment_id < 0) { + return Status::InvalidArgument("Lance fragment id must be non-negative: {}", + fragment_id); + } + fragment_ids->emplace_back(static_cast(fragment_id)); + } + return Status::OK(); +} + +Status parse_index_segment_uuids(const TLanceFileDesc& lance_params, + std::vector* segment_uuids, size_t* segment_count) { + DORIS_CHECK(segment_uuids != nullptr); + DORIS_CHECK(segment_count != nullptr); + segment_uuids->clear(); + *segment_count = 0; + if (!lance_params.__isset.index_segment_uuids || lance_params.index_segment_uuids.empty()) { + return Status::OK(); + } + constexpr size_t UUID_SIZE = 16; + if (lance_params.index_segment_uuids.size() > std::numeric_limits::max() / UUID_SIZE) { + return Status::InvalidArgument("too many Lance index segment UUIDs"); + } + segment_uuids->reserve(lance_params.index_segment_uuids.size() * UUID_SIZE); + for (const auto& uuid : lance_params.index_segment_uuids) { + if (uuid.size() != UUID_SIZE) { + return Status::InvalidArgument("Lance index segment UUID must contain 16 bytes, got {}", + uuid.size()); + } + segment_uuids->insert(segment_uuids->end(), uuid.begin(), uuid.end()); + } + *segment_count = lance_params.index_segment_uuids.size(); + return Status::OK(); +} + +Status convert_arrow_schema_to_doris(const std::shared_ptr& arrow_schema, + std::vector* column_names, + std::vector* column_types) { + DORIS_CHECK(arrow_schema != nullptr); + DORIS_CHECK(column_names != nullptr); + DORIS_CHECK(column_types != nullptr); + + std::vector parsed_names; + std::vector parsed_types; + parsed_names.reserve(arrow_schema->num_fields()); + parsed_types.reserve(arrow_schema->num_fields()); + std::unordered_set unique_names; + unique_names.reserve(arrow_schema->num_fields()); + for (const auto& field : arrow_schema->fields()) { + if (!unique_names.emplace(field->name()).second) { + return Status::InvalidArgument("duplicate Lance schema column: {}", field->name()); + } + DataTypePtr doris_type; + const auto type_status = arrow_field_to_doris_type(field, &doris_type); + if (type_status.is()) { + parsed_types.emplace_back(std::make_shared()); + } else { + RETURN_IF_ERROR(type_status); + DORIS_CHECK(doris_type != nullptr); + parsed_types.emplace_back(std::move(doris_type)); + } + parsed_names.emplace_back(field->name()); + } + *column_names = std::move(parsed_names); + *column_types = std::move(parsed_types); + return Status::OK(); +} + +Status build_lance_storage_options(const TFileScanRangeParams* scan_params, + std::vector* options) { + DORIS_CHECK(options != nullptr); + options->clear(); + if (scan_params == nullptr || !scan_params->__isset.lance_scan_params || + !scan_params->lance_scan_params.__isset.lance_storage_options) { + return Status::OK(); + } + const auto& storage_options = scan_params->lance_scan_params.lance_storage_options; + options->reserve(storage_options.size() * 2); + for (const auto& [key, value] : storage_options) { + // Both values cross a C-string boundary. Reject embedded NULs instead of silently opening + // a different dataset configuration from the one validated and used by the FE. + if (key.find('\0') != std::string::npos || value.find('\0') != std::string::npos) { + return Status::InvalidArgument( + "Lance storage option '{}' contains a NUL and cannot reach lance-c", + key.substr(0, key.find('\0'))); + } + options->emplace_back(key); + options->emplace_back(value); + } + return Status::OK(); +} + +Status lance_error(std::string_view operation) { + const char* raw_message = lance_last_error_message(); + std::string message = raw_message == nullptr ? "" : raw_message; + if (raw_message != nullptr) { + lance_free_string(raw_message); + } + if (message.empty()) { + message = fmt::format("error_code={}", static_cast(lance_last_error_code())); + } + return Status::InternalError("{} failed: {}", operation, message); +} + +} // namespace doris::format::lance diff --git a/be/src/format_v2/lance/lance_reader_helper.h b/be/src/format_v2/lance/lance_reader_helper.h new file mode 100644 index 00000000000000..689e4f4fd9f492 --- /dev/null +++ b/be/src/format_v2/lance/lance_reader_helper.h @@ -0,0 +1,81 @@ +// 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. + +#pragma once + +#include +#include +#include +#include +#include +#include + +#include "common/status.h" +#include "core/data_type/data_type.h" +#include "gen_cpp/PlanNodes_types.h" + +struct LanceBatch; +struct LanceDataset; +struct LanceScanner; + +namespace arrow { +class Schema; +} // namespace arrow + +namespace doris::format::lance { + +inline constexpr std::string_view LANCE_DISTANCE_COLUMN = "_distance"; +inline constexpr std::string_view LANCE_SCORE_COLUMN = "_score"; +inline constexpr std::string_view LANCE_ROW_ID_COLUMN = "_rowid"; +inline constexpr const char* LANCE_READER_PROFILE = "LanceReader"; + +struct LanceDatasetDeleter { + void operator()(LanceDataset* dataset) const; +}; + +struct LanceScannerDeleter { + void operator()(LanceScanner* scanner) const; +}; + +struct LanceBatchDeleter { + void operator()(LanceBatch* batch) const; +}; + +size_t lance_vector_element_width(TVectorElementType::type type); + +// Validate and convert the fragment and index-segment identifiers carried by the FE into the +// unsigned and packed representations expected by lance-c. +Status parse_fragment_ids(const TLanceFileDesc& lance_params, std::vector* fragment_ids); +Status parse_index_segment_uuids(const TLanceFileDesc& lance_params, + std::vector* segment_uuids, size_t* segment_count); + +// Convert every top-level field without discarding unsupported columns. Malformed schemas still +// return an error and leave both output vectors unchanged. DataTypeNothing is the local sentinel +// for a valid Arrow field whose logical type Doris does not support. +Status convert_arrow_schema_to_doris(const std::shared_ptr& arrow_schema, + std::vector* column_names, + std::vector* column_types); + +// The FE sends storage options in Lance's own vocabulary. Preserve the key-value sequence exactly +// while validating that every value can cross the C-string boundary into lance-c. +Status build_lance_storage_options(const TFileScanRangeParams* scan_params, + std::vector* options); + +// Copy and release lance-c's thread-local error message before returning a Doris status. +Status lance_error(std::string_view operation); + +} // namespace doris::format::lance diff --git a/be/src/format_v2/lance/lance_runtime_filter_helper.cpp b/be/src/format_v2/lance/lance_runtime_filter_helper.cpp new file mode 100644 index 00000000000000..26b7c6402e78c7 --- /dev/null +++ b/be/src/format_v2/lance/lance_runtime_filter_helper.cpp @@ -0,0 +1,363 @@ +// 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. + +#include "format_v2/lance/lance_runtime_filter_helper.h" + +#include + +#include +#include +#include +#include +#include + +#include "common/logging.h" +#include "core/data_type/data_type_nullable.h" +#include "core/field.h" +#include "exprs/hybrid_set.h" +#include "exprs/runtime_filter_expr.h" +#include "exprs/vdirect_in_predicate.h" +#include "exprs/vexpr_context.h" +#include "exprs/vliteral.h" +#include "exprs/vslot_ref.h" +#include "format/format_common.h" +#include "runtime/runtime_profile.h" + +namespace doris::format::lance { +namespace { + +constexpr std::string_view LANCE_RUNTIME_FILTER_CACHE_KEY_PREFIX = "lance-runtime-filter-sql:"; + +std::string format_filter_ids(const std::vector& filter_ids) { + std::string result; + for (const auto filter_id : filter_ids) { + if (!result.empty()) { + result.append(","); + } + result.append(std::to_string(filter_id)); + } + return result; +} + +std::string quote_sql_identifier(std::string_view identifier) { + // Lance SQL uses backticks for delimited identifiers. Escape an embedded backtick by doubling + // it, matching the SQL parser's quoted-identifier syntax. + std::string quoted("`"); + quoted.reserve(identifier.size() + 2); + for (const char ch : identifier) { + if (ch == '`') { + quoted.append("``"); + } else { + quoted.push_back(ch); + } + } + quoted.push_back('`'); + return quoted; +} + +const RuntimeFilterExpr* get_runtime_filter(const VExprContextSPtr& conjunct) { + if (conjunct == nullptr || conjunct->root() == nullptr) { + return nullptr; + } + return dynamic_cast(conjunct->root().get()); +} + +void append_sql_conjunct(std::string_view conjunct, std::string* expression) { + if (!expression->empty()) { + expression->append(" AND "); + } + expression->append(conjunct); +} + +std::string lowercase_ascii(std::string value) { + std::ranges::transform(value, value.begin(), [](const unsigned char ch) { + return static_cast(std::tolower(ch)); + }); + return value; +} + +std::string quote_string_value(std::string_view value) { + std::string quoted("'"); + quoted.reserve(value.size() + 2); + for (const char ch : value) { + if (ch == '\'') { + quoted.append("''"); + } else { + quoted.push_back(ch); + } + } + quoted.push_back('\''); + return quoted; +} + +std::optional to_lance_sql_literal(const VLiteral& literal) { + const auto type = remove_nullable(literal.get_data_type()); + auto options = DataTypeSerDe::get_default_format_options(); + auto timezone = cctz::utc_time_zone(); + options.timezone = &timezone; + const auto value = literal.value(options); + switch (type->get_primitive_type()) { + case TYPE_BOOLEAN: { + const auto normalized = lowercase_ascii(value); + if (normalized == "0" || normalized == "false") { + return "FALSE"; + } + if (normalized == "1" || normalized == "true") { + return "TRUE"; + } + return std::nullopt; + } + case TYPE_TINYINT: + case TYPE_SMALLINT: + case TYPE_INT: + case TYPE_BIGINT: + case TYPE_LARGEINT: + return value; + case TYPE_FLOAT: + case TYPE_DOUBLE: { + const auto normalized = lowercase_ascii(value); + if (normalized.find("nan") != std::string::npos || + normalized.find("inf") != std::string::npos) { + return std::nullopt; + } + return value; + } + case TYPE_DECIMALV2: + case TYPE_DECIMAL32: + case TYPE_DECIMAL64: + case TYPE_DECIMAL128I: + case TYPE_DECIMAL256: + return value; + case TYPE_CHAR: + case TYPE_VARCHAR: + case TYPE_STRING: + return quote_string_value(value); + case TYPE_DATE: + case TYPE_DATEV2: + return "DATE " + quote_string_value(value); + case TYPE_DATETIME: + case TYPE_DATETIMEV2: + return "TIMESTAMP " + quote_string_value(value); + default: + return std::nullopt; + } +} + +template +std::optional in_value_to_lance_sql_literal(const void* raw_value, + const DataTypePtr& data_type) { + if (raw_value == nullptr || data_type == nullptr) { + return std::nullopt; + } + Field field; + if constexpr (is_string_type(PT)) { + const auto* value = static_cast(raw_value); + using CppType = typename PrimitiveTypeTraits::CppType; + field = Field::create_field(CppType(value->data, value->size)); + } else { + using CppType = typename PrimitiveTypeTraits::CppType; + field = Field::create_field(*static_cast(raw_value)); + } + return to_lance_sql_literal(VLiteral(data_type, field)); +} + +std::optional in_value_to_lance_sql_literal(PrimitiveType primitive_type, + const void* raw_value, + const DataTypePtr& data_type) { +#define DISPATCH_IN_VALUE(TYPE) \ + case TYPE: \ + return in_value_to_lance_sql_literal(raw_value, data_type) + switch (primitive_type) { + DISPATCH_IN_VALUE(TYPE_BOOLEAN); + DISPATCH_IN_VALUE(TYPE_TINYINT); + DISPATCH_IN_VALUE(TYPE_SMALLINT); + DISPATCH_IN_VALUE(TYPE_INT); + DISPATCH_IN_VALUE(TYPE_BIGINT); + DISPATCH_IN_VALUE(TYPE_LARGEINT); + DISPATCH_IN_VALUE(TYPE_FLOAT); + DISPATCH_IN_VALUE(TYPE_DOUBLE); + DISPATCH_IN_VALUE(TYPE_DATE); + DISPATCH_IN_VALUE(TYPE_DATETIME); + DISPATCH_IN_VALUE(TYPE_DATEV2); + DISPATCH_IN_VALUE(TYPE_DATETIMEV2); + DISPATCH_IN_VALUE(TYPE_CHAR); + DISPATCH_IN_VALUE(TYPE_VARCHAR); + DISPATCH_IN_VALUE(TYPE_STRING); + DISPATCH_IN_VALUE(TYPE_DECIMALV2); + DISPATCH_IN_VALUE(TYPE_DECIMAL32); + DISPATCH_IN_VALUE(TYPE_DECIMAL64); + DISPATCH_IN_VALUE(TYPE_DECIMAL128I); + DISPATCH_IN_VALUE(TYPE_DECIMAL256); + default: + return std::nullopt; + } +#undef DISPATCH_IN_VALUE +} + +std::optional build_in_filter_sql(const VDirectInPredicate& predicate) { + if (predicate.get_num_children() != 1) { + return std::nullopt; + } + const auto slot = std::dynamic_pointer_cast(predicate.get_child(0)); + const auto values = predicate.get_set_func(); + if (slot == nullptr || slot->data_type() == nullptr || values == nullptr || + values->contain_null() || values->size() == 0) { + return std::nullopt; + } + + const auto data_type = remove_nullable(slot->data_type()); + std::string expression("(" + quote_sql_identifier(slot->column_name()) + " IN ("); + auto* iterator = values->begin(); + bool first_value = true; + while (iterator != nullptr && iterator->has_next()) { + auto value = in_value_to_lance_sql_literal(data_type->get_primitive_type(), + iterator->get_value(), data_type); + if (!value.has_value()) { + return std::nullopt; + } + if (!first_value) { + expression.append(", "); + } + expression.append(*value); + first_value = false; + iterator->next(); + } + if (first_value) { + return std::nullopt; + } + expression.append("))"); + return expression; +} + +std::optional build_range_filter_sql(const VExpr& predicate) { + if ((predicate.op() != TExprOpcode::GE && predicate.op() != TExprOpcode::LE) || + predicate.get_num_children() != 2) { + return std::nullopt; + } + const auto slot = std::dynamic_pointer_cast(predicate.get_child(0)); + const auto literal = std::dynamic_pointer_cast(predicate.get_child(1)); + if (slot == nullptr || literal == nullptr) { + return std::nullopt; + } + const auto sql_literal = to_lance_sql_literal(*literal); + if (!sql_literal.has_value()) { + return std::nullopt; + } + const auto* sql_operator = predicate.op() == TExprOpcode::GE ? ">=" : "<="; + return "(" + quote_sql_identifier(slot->column_name()) + " " + sql_operator + " " + + *sql_literal + ")"; +} + +std::optional runtime_filter_to_lance_sql(const RuntimeFilterExpr& runtime_filter) { + const auto impl = runtime_filter.get_impl(); + if (impl == nullptr) { + return std::nullopt; + } + if (const auto* in_predicate = dynamic_cast(impl.get()); + in_predicate != nullptr) { + return build_in_filter_sql(*in_predicate); + } + return build_range_filter_sql(*impl); +} + +std::shared_ptr build_runtime_filter_sql( + const VExprContextSPtrs& conjuncts) { + auto result = std::make_shared(); + std::set seen_filter_ids; + std::set pushed_filter_ids; + for (const auto& conjunct : conjuncts) { + const auto* runtime_filter = get_runtime_filter(conjunct); + if (runtime_filter == nullptr) { + continue; + } + const auto filter_id = runtime_filter->filter_id(); + seen_filter_ids.emplace(filter_id); + + const auto expression = runtime_filter_to_lance_sql(*runtime_filter); + if (!expression.has_value()) { + continue; + } + append_sql_conjunct(*expression, &result->expression); + pushed_filter_ids.emplace(filter_id); + } + + result->pushable_filter_ids.assign(pushed_filter_ids.begin(), pushed_filter_ids.end()); + for (const auto filter_id : seen_filter_ids) { + if (!pushed_filter_ids.contains(filter_id)) { + result->skipped_filter_ids.emplace_back(filter_id); + } + } + return result; +} + +std::optional build_cache_key(const VExprContextSPtrs& conjuncts) { + // This cache is scoped to one FileScanLocalState. An RF is immutable after it is published, so + // the sorted RF IDs uniquely identify the snapshot shared by its parallel scanners. + std::set filter_ids; + for (const auto& conjunct : conjuncts) { + if (const auto* runtime_filter = get_runtime_filter(conjunct); runtime_filter != nullptr) { + filter_ids.emplace(runtime_filter->filter_id()); + } + } + if (filter_ids.empty()) { + return std::nullopt; + } + + std::string key(LANCE_RUNTIME_FILTER_CACHE_KEY_PREFIX); + for (const auto filter_id : filter_ids) { + key.append(std::to_string(filter_id)).append(","); + } + return key; +} + +} // namespace + +std::shared_ptr get_or_create_lance_runtime_filter_sql( + const VExprContextSPtrs& conjuncts, ShardedKVCache* cache) { + const auto cache_key = build_cache_key(conjuncts); + if (!cache_key.has_value()) { + return nullptr; + } + if (cache == nullptr) { + return build_runtime_filter_sql(conjuncts); + } + + auto* cached = cache->get>( + *cache_key, [&]() -> std::shared_ptr* { + return new std::shared_ptr( + build_runtime_filter_sql(conjuncts)); + }); + return cached == nullptr ? nullptr : *cached; +} + +void record_lance_runtime_filter_pushdown(RuntimeProfile* profile, + const LanceRuntimeFilterSql& runtime_filter_sql) { + DORIS_CHECK(profile != nullptr); + const auto pushed_ids = format_filter_ids(runtime_filter_sql.pushable_filter_ids); + const auto skipped_ids = format_filter_ids(runtime_filter_sql.skipped_filter_ids); + + if (!pushed_ids.empty()) { + profile->add_info_string("LanceRuntimeFilterPushedIds", pushed_ids); + } + if (!skipped_ids.empty()) { + profile->add_info_string("LanceRuntimeFilterSkippedIds", skipped_ids); + } + VLOG_DEBUG << "Lance runtime filter pushdown: pushed_ids=[" << pushed_ids << "], skipped_ids=[" + << skipped_ids << "]"; +} + +} // namespace doris::format::lance diff --git a/be/src/format_v2/lance/lance_runtime_filter_helper.h b/be/src/format_v2/lance/lance_runtime_filter_helper.h new file mode 100644 index 00000000000000..c40631bfdbcc09 --- /dev/null +++ b/be/src/format_v2/lance/lance_runtime_filter_helper.h @@ -0,0 +1,50 @@ +// 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. + +#pragma once + +#include +#include +#include + +#include "exprs/vexpr_fwd.h" + +namespace doris { + +class ShardedKVCache; +class RuntimeProfile; + +namespace format::lance { + +struct LanceRuntimeFilterSql { + std::string expression; + std::vector pushable_filter_ids; + std::vector skipped_filter_ids; +}; + +// Build one immutable SQL snapshot for all supported Doris runtime filters. When cache is non-null, +// equivalent RF snapshots from parallel Lance readers in the same FileScanLocalState share the +// conversion result. The returned snapshot also identifies RFs that cannot be represented exactly +// by Lance SQL. A null result means that conjuncts contain no Doris runtime filter. +std::shared_ptr get_or_create_lance_runtime_filter_sql( + const VExprContextSPtrs& conjuncts, ShardedKVCache* cache); + +void record_lance_runtime_filter_pushdown(RuntimeProfile* profile, + const LanceRuntimeFilterSql& runtime_filter_sql); + +} // namespace format::lance +} // namespace doris diff --git a/be/src/format_v2/table/lance_reader.cpp b/be/src/format_v2/table/lance_reader.cpp index e8ad0c5229aca3..c46700fd0fe352 100644 --- a/be/src/format_v2/table/lance_reader.cpp +++ b/be/src/format_v2/table/lance_reader.cpp @@ -21,236 +21,26 @@ #include #include #include -#include #include +#include #include #include #include #include +#include #include "common/consts.h" #include "common/logging.h" #include "core/column/column_nullable.h" #include "core/column/column_string.h" -#include "core/data_type/data_type_array.h" -#include "core/data_type/data_type_factory.hpp" -#include "core/data_type/data_type_map.h" -#include "core/data_type/data_type_nothing.h" -#include "core/data_type/data_type_struct.h" #include "exec/common/endian.h" +#include "format_v2/lance/lance_reader_helper.h" +#include "format_v2/lance/lance_runtime_filter_helper.h" #include "runtime/file_scan_profile.h" #include "storage/utils.h" namespace doris::format::lance { -namespace { - -struct LanceDatasetDeleter { - void operator()(LanceDataset* dataset) const { lance_dataset_close(dataset); } -}; - -struct LanceScannerDeleter { - void operator()(LanceScanner* scanner) const { lance_scanner_close(scanner); } -}; - -struct LanceBatchDeleter { - void operator()(LanceBatch* batch) const { lance_batch_free(batch); } -}; - -constexpr std::string_view DISTANCE_COLUMN = "_distance"; -constexpr std::string_view ROW_ID_COLUMN = "_rowid"; -constexpr std::string_view ARROW_EXTENSION_NAME = "ARROW:extension:name"; -constexpr const char* LANCE_READER_PROFILE = "LanceReader"; - -size_t vector_element_width(TVectorElementType::type type) { - switch (type) { - case TVectorElementType::FLOAT16: - return sizeof(uint16_t); - case TVectorElementType::FLOAT32: - return sizeof(float); - case TVectorElementType::FLOAT64: - return sizeof(double); - case TVectorElementType::UINT8: - case TVectorElementType::INT8: - return sizeof(uint8_t); - } - return 0; -} - -int arrow_time_precision(arrow::TimeUnit::type unit) { - switch (unit) { - case arrow::TimeUnit::SECOND: - return 0; - case arrow::TimeUnit::MILLI: - return 3; - case arrow::TimeUnit::MICRO: - case arrow::TimeUnit::NANO: - return 6; - } - return 6; -} - -Status check_arrow_field_semantics(const std::shared_ptr& field) { - if (field->HasMetadata()) { - const auto extension_name = field->metadata()->Get(ARROW_EXTENSION_NAME); - if (extension_name.ok() && !extension_name.ValueUnsafe().empty()) { - return Status::NotSupported( - "unsupported Lance Arrow extension type '{}' for field '{}'", - extension_name.ValueUnsafe(), field->name()); - } - } - if (field->type()->id() == arrow::Type::DICTIONARY) { - return Status::NotSupported("unsupported Lance Arrow dictionary type for field '{}': {}", - field->name(), field->type()->ToString()); - } - return Status::OK(); -} - -Status arrow_field_to_doris_type(const std::shared_ptr& field, - DataTypePtr* doris_type) { - RETURN_IF_ERROR(check_arrow_field_semantics(field)); - const auto& arrow_type = field->type(); - const auto nullable_primitive = [&](PrimitiveType type, int precision = 0, int scale = 0, - int len = -1) { - *doris_type = - DataTypeFactory::instance().create_data_type(type, true, precision, scale, len); - return Status::OK(); - }; - - switch (arrow_type->id()) { - case arrow::Type::BOOL: - return nullable_primitive(TYPE_BOOLEAN); - case arrow::Type::INT8: - return nullable_primitive(TYPE_TINYINT); - case arrow::Type::UINT8: - case arrow::Type::INT16: - return nullable_primitive(TYPE_SMALLINT); - case arrow::Type::UINT16: - case arrow::Type::INT32: - return nullable_primitive(TYPE_INT); - case arrow::Type::UINT32: - case arrow::Type::INT64: - return nullable_primitive(TYPE_BIGINT); - case arrow::Type::UINT64: - return nullable_primitive(TYPE_LARGEINT); - case arrow::Type::HALF_FLOAT: - case arrow::Type::FLOAT: - return nullable_primitive(TYPE_FLOAT); - case arrow::Type::DOUBLE: - return nullable_primitive(TYPE_DOUBLE); - case arrow::Type::STRING: - case arrow::Type::LARGE_STRING: - return nullable_primitive(TYPE_STRING); - case arrow::Type::BINARY: - case arrow::Type::LARGE_BINARY: - return nullable_primitive(TYPE_VARBINARY, 0, 0, std::numeric_limits::max()); - case arrow::Type::FIXED_SIZE_BINARY: { - const auto binary = std::static_pointer_cast(arrow_type); - return nullable_primitive(TYPE_VARBINARY, 0, 0, binary->byte_width()); - } - case arrow::Type::DATE32: - case arrow::Type::DATE64: - return nullable_primitive(TYPE_DATEV2); - case arrow::Type::TIME32: - case arrow::Type::TIME64: { - const auto time = std::static_pointer_cast(arrow_type); - return nullable_primitive(TYPE_TIMEV2, 0, arrow_time_precision(time->unit())); - } - case arrow::Type::TIMESTAMP: { - const auto timestamp = std::static_pointer_cast(arrow_type); - const auto doris_type = timestamp->timezone().empty() ? TYPE_DATETIMEV2 : TYPE_TIMESTAMPTZ; - return nullable_primitive(doris_type, 0, arrow_time_precision(timestamp->unit())); - } - case arrow::Type::DECIMAL128: - case arrow::Type::DECIMAL256: { - const auto decimal = std::static_pointer_cast(arrow_type); - const int precision = decimal->precision(); - const int scale = decimal->scale(); - if (precision <= 0 || precision > arrow::Decimal256Type::kMaxPrecision || scale < 0 || - scale > precision) { - return Status::NotSupported( - "unsupported Lance Arrow decimal type for field '{}': precision={}, scale={}", - field->name(), precision, scale); - } - const PrimitiveType doris_decimal_type = precision <= 9 ? TYPE_DECIMAL32 - : precision <= 18 ? TYPE_DECIMAL64 - : precision <= 38 ? TYPE_DECIMAL128I - : TYPE_DECIMAL256; - return nullable_primitive(doris_decimal_type, precision, scale); - } - case arrow::Type::LIST: - case arrow::Type::LARGE_LIST: - case arrow::Type::FIXED_SIZE_LIST: { - const auto list = std::static_pointer_cast(arrow_type); - DataTypePtr value_type; - RETURN_IF_ERROR(arrow_field_to_doris_type(list->value_field(), &value_type)); - *doris_type = make_nullable(std::make_shared(value_type)); - return Status::OK(); - } - case arrow::Type::MAP: { - const auto map = std::static_pointer_cast(arrow_type); - RETURN_IF_ERROR(check_arrow_field_semantics(map->value_field())); - DataTypePtr key_type; - DataTypePtr item_type; - RETURN_IF_ERROR(arrow_field_to_doris_type(map->key_field(), &key_type)); - RETURN_IF_ERROR(arrow_field_to_doris_type(map->item_field(), &item_type)); - *doris_type = make_nullable(std::make_shared(key_type, item_type)); - return Status::OK(); - } - case arrow::Type::STRUCT: { - const auto struct_type = std::static_pointer_cast(arrow_type); - DataTypes field_types; - Strings field_names; - field_types.reserve(struct_type->num_fields()); - field_names.reserve(struct_type->num_fields()); - for (const auto& child : struct_type->fields()) { - DataTypePtr field_type; - RETURN_IF_ERROR(arrow_field_to_doris_type(child, &field_type)); - field_types.emplace_back(std::move(field_type)); - field_names.emplace_back(child->name()); - } - *doris_type = make_nullable(std::make_shared(field_types, field_names)); - return Status::OK(); - } - default: - return Status::NotSupported("unsupported Lance Arrow type: {}", arrow_type->ToString()); - } -} - -} // namespace - -Status convert_arrow_schema_to_doris(const std::shared_ptr& arrow_schema, - std::vector* column_names, - std::vector* column_types) { - DORIS_CHECK(arrow_schema != nullptr); - DORIS_CHECK(column_names != nullptr); - DORIS_CHECK(column_types != nullptr); - - std::vector parsed_names; - std::vector parsed_types; - parsed_names.reserve(arrow_schema->num_fields()); - parsed_types.reserve(arrow_schema->num_fields()); - std::unordered_set unique_names; - unique_names.reserve(arrow_schema->num_fields()); - for (const auto& field : arrow_schema->fields()) { - if (!unique_names.emplace(field->name()).second) { - return Status::InvalidArgument("duplicate Lance schema column: {}", field->name()); - } - DataTypePtr doris_type; - const auto type_status = arrow_field_to_doris_type(field, &doris_type); - if (type_status.is()) { - parsed_types.emplace_back(std::make_shared()); - } else { - RETURN_IF_ERROR(type_status); - DORIS_CHECK(doris_type != nullptr); - parsed_types.emplace_back(std::move(doris_type)); - } - parsed_names.emplace_back(field->name()); - } - *column_names = std::move(parsed_names); - *column_types = std::move(parsed_types); - return Status::OK(); -} LanceTableReader::~LanceTableReader() { static_cast(close()); @@ -265,7 +55,7 @@ Status LanceTableReader::fetch_schema(const TFileRangeDesc& range, } const auto& params = range.table_format_params.lance_params; std::vector storage_options; - RETURN_IF_ERROR(_storage_options(&scan_params, &storage_options)); + RETURN_IF_ERROR(build_lance_storage_options(&scan_params, &storage_options)); std::vector storage_option_ptrs; storage_option_ptrs.reserve(storage_options.size() + 1); for (const auto& option : storage_options) { @@ -278,12 +68,12 @@ Status LanceTableReader::fetch_schema(const TFileRangeDesc& range, storage_options.empty() ? nullptr : storage_option_ptrs.data(), static_cast(params.version))); if (dataset == nullptr) { - return _lance_error("open Lance dataset for schema"); + return lance_error("open Lance dataset for schema"); } ArrowSchema arrow_schema {}; if (lance_dataset_schema(dataset.get(), &arrow_schema) != 0) { - return _lance_error("get Lance dataset schema"); + return lance_error("get Lance dataset schema"); } auto imported_schema = arrow::ImportSchema(&arrow_schema); if (!imported_schema.ok()) { @@ -303,6 +93,7 @@ Status LanceTableReader::init(TableReadOptions&& options) { DORIS_CHECK(_runtime_state != nullptr); DORIS_CHECK(_scanner_profile != nullptr); DORIS_CHECK(_scan_params != nullptr); + RETURN_IF_ERROR(_resolve_search_kind()); _ctz = _runtime_state->timezone_obj(); const auto& lance_scan_params = _scan_params->lance_scan_params; @@ -358,16 +149,29 @@ Status LanceTableReader::init(TableReadOptions&& options) { ADD_CHILD_TIMER_WITH_LEVEL(_scanner_profile, "LanceIVFPartitionRankingTime", LANCE_READER_PROFILE, 1)}, }; - _vector_search = _scan_params->__isset.lance_scan_params && - lance_scan_params.__isset.external_search_request; - if (_vector_search) { + if (_search_kind != SearchKind::NORMAL) { RETURN_IF_ERROR(_validate_external_search_request()); const auto& request = lance_scan_params.external_search_request; - const auto& vector = request.search_query.vector_search; - _scanner_profile->add_info_string("LanceTopK", std::to_string(vector.top_k)); - _scanner_profile->add_info_string("LanceOffset", std::to_string(vector.offset)); - _scanner_profile->add_info_string("LanceTopKPlusOffset", - std::to_string(vector.top_k + vector.offset)); + int64_t top_k; + int64_t offset; + if (_search_kind == SearchKind::VECTOR) { + const auto& vector = request.search_query.vector_search; + top_k = vector.top_k; + offset = vector.offset; + _scanner_profile->add_info_string("LanceSearchType", "VECTOR"); + } else { + DORIS_CHECK(_search_kind == SearchKind::FULL_TEXT); + const auto& full_text = request.search_query.full_text_search; + top_k = full_text.top_k; + offset = full_text.offset; + _scanner_profile->add_info_string("LanceSearchType", "FULL_TEXT"); + _scanner_profile->add_info_string( + "LanceFtsCoverageMode", + full_text.coverage_mode == TFtsCoverageMode::STRICT ? "STRICT" : "INDEX_ONLY"); + } + _scanner_profile->add_info_string("LanceTopK", std::to_string(top_k)); + _scanner_profile->add_info_string("LanceOffset", std::to_string(offset)); + _scanner_profile->add_info_string("LanceTopKPlusOffset", std::to_string(top_k + offset)); _planned_index_segment_count = ADD_CHILD_COUNTER_WITH_LEVEL(_scanner_profile, "LancePlannedIndexSegmentCount", TUnit::UNIT, LANCE_READER_PROFILE, 1); @@ -395,9 +199,9 @@ Status LanceTableReader::init(TableReadOptions&& options) { return Status::InvalidArgument("Lance projected column '{}' has no type", column.name); } if (column.name.starts_with(BeConsts::GLOBAL_ROWID_COL)) { - if (!_vector_search) { + if (_search_kind == SearchKind::NORMAL) { return Status::NotSupported( - "Lance global row id is currently supported only for vector search"); + "Lance global row id is currently supported only for external search"); } if (_global_rowid_output_idx.has_value()) { return Status::InvalidArgument("duplicate Lance global row id projected column: {}", @@ -414,12 +218,19 @@ Status LanceTableReader::init(TableReadOptions&& options) { if (!_output_name_to_idx.emplace(column.name, idx).second) { return Status::InvalidArgument("duplicate Lance projected column: {}", column.name); } - if (_vector_search && column.name == DISTANCE_COLUMN) { + if (_search_kind == SearchKind::VECTOR && column.name == LANCE_DISTANCE_COLUMN) { const auto distance_type = remove_nullable(column.type); if (distance_type->get_primitive_type() != TYPE_FLOAT) { return Status::InvalidArgument( "Lance vector search column '{}' must have Doris FLOAT type, but was {}", - DISTANCE_COLUMN, column.type->get_name()); + LANCE_DISTANCE_COLUMN, column.type->get_name()); + } + } else if (_search_kind == SearchKind::FULL_TEXT && column.name == LANCE_SCORE_COLUMN) { + const auto score_type = remove_nullable(column.type); + if (score_type->get_primitive_type() != TYPE_FLOAT) { + return Status::InvalidArgument( + "Lance full-text search column '{}' must have Doris FLOAT type, but was {}", + LANCE_SCORE_COLUMN, column.type->get_name()); } } } @@ -429,6 +240,7 @@ Status LanceTableReader::init(TableReadOptions&& options) { Status LanceTableReader::prepare_split(const SplitReadOptions& options) { _close_scanner(); _eof = false; + _runtime_filter_cache = options.cache; RETURN_IF_ERROR(TableReader::prepare_split(options)); // Lance does not currently provide metadata aggregate pushdown. Do not let a generic @@ -485,7 +297,7 @@ Status LanceTableReader::get_block(Block* block, bool* eos) { break; } if (scan_status != 0 || raw_batch == nullptr) { - return _lance_error("read next Lance batch"); + return lance_error("read next Lance batch"); } std::unique_ptr batch(raw_batch); @@ -533,7 +345,9 @@ Status LanceTableReader::read_by_row_ids(const TFileRangeDesc& range, } SCOPED_TIMER(_row_id_fetch_total_time); - RETURN_IF_ERROR(_ensure_dataset_open(range)); + // Phase-two row fetch does not execute FTS, so a reader created only for take_rows must not + // collect query-specific global statistics. + RETURN_IF_ERROR(_ensure_dataset_open(range, false)); std::vector columns; columns.reserve(_projected_columns.size() + 1); for (const auto& column : _projected_columns) { @@ -552,7 +366,7 @@ Status LanceTableReader::read_by_row_ids(const TFileRangeDesc& range, if (stream.release != nullptr) { stream.release(&stream); } - return _lance_error("take Lance rows by row id"); + return lance_error("take Lance rows by row id"); } auto imported_reader = arrow::ImportRecordBatchReader(&stream); if (!imported_reader.ok()) { @@ -608,8 +422,31 @@ Status LanceTableReader::close() { return TableReader::close(); } +Status LanceTableReader::_resolve_search_kind() { + DORIS_CHECK(_scan_params != nullptr); + _search_kind = SearchKind::NORMAL; + if (!_scan_params->__isset.lance_scan_params) { + return Status::OK(); + } + const auto& lance_scan_params = _scan_params->lance_scan_params; + if (!lance_scan_params.__isset.external_search_request) { + return Status::OK(); + } + const auto& request = lance_scan_params.external_search_request; + if (!request.__isset.search_query) { + return Status::InvalidArgument("external search request requires search_query"); + } + const bool has_vector = request.search_query.__isset.vector_search; + const bool has_full_text = request.search_query.__isset.full_text_search; + if (has_vector == has_full_text) { + return Status::InvalidArgument("external search query must set exactly one search kind"); + } + _search_kind = has_vector ? SearchKind::VECTOR : SearchKind::FULL_TEXT; + return Status::OK(); +} + Status LanceTableReader::_validate_external_search_request() const { - // FE validates requests produced by vector_search(), but this reader consumes a deserialized + // FE validates requests produced by the search TVFs, but this reader consumes a deserialized // Thrift boundary. Recheck structural invariants and values used for allocation, pointer // arithmetic, C-string calls, and narrowing conversions before accessing them below. DORIS_CHECK(_scan_params != nullptr); @@ -618,7 +455,7 @@ Status LanceTableReader::_validate_external_search_request() const { DORIS_CHECK(lance_scan_params.__isset.external_search_request); if (lance_scan_params.__isset.lance_substrait_filter) { return Status::InvalidArgument( - "Lance vector search cannot combine its pre-search filter with " + "Lance external search cannot combine its pre-search filter with " "lance_substrait_filter"); } @@ -627,58 +464,87 @@ Status LanceTableReader::_validate_external_search_request() const { return Status::NotSupported("unsupported external search schema version: {}", request.schema_version); } - if (!request.__isset.search_query) { - return Status::InvalidArgument("external search request requires search_query"); - } - - const bool has_vector = request.search_query.__isset.vector_search; - const bool has_full_text = request.search_query.__isset.full_text_search; - if (has_vector == has_full_text) { - return Status::InvalidArgument("external search query must set exactly one search kind"); - } - if (has_full_text) { - return Status::NotSupported("Lance Format V2 reader does not yet support full-text search"); - } - - const auto& vector = request.search_query.vector_search; - if (!vector.__isset.column || vector.column.empty() || - vector.column.find('\0') != std::string::npos) { - return Status::InvalidArgument("Lance vector search requires a non-empty column"); - } - if (!vector.__isset.query_vector) { - return Status::InvalidArgument("Lance vector search requires a query vector"); - } - const auto& query_vector = vector.query_vector; - if (!query_vector.__isset.element_type || !query_vector.__isset.dimension || - !query_vector.__isset.values) { - return Status::InvalidArgument( - "Lance query vector requires element_type, dimension, and values"); - } - if (query_vector.dimension <= 0) { - return Status::InvalidArgument("Lance query vector dimension must be positive: {}", - query_vector.dimension); - } - const auto element_width = vector_element_width(query_vector.element_type); - if (element_width == 0) { - return Status::NotSupported("unsupported Lance query vector element type: {}", - static_cast(query_vector.element_type)); - } - const auto dimension = static_cast(query_vector.dimension); - if (dimension > std::numeric_limits::max() / element_width || - query_vector.values.size() != dimension * element_width) { - return Status::InvalidArgument( - "Lance query vector byte size {} does not match dimension {} and element width {}", - query_vector.values.size(), dimension, element_width); - } - if (!vector.__isset.top_k || vector.top_k <= 0) { - return Status::InvalidArgument("Lance vector search top_k must be positive"); - } - if (!vector.__isset.offset || vector.offset < 0) { - return Status::InvalidArgument("Lance vector search offset must be non-negative"); - } + DORIS_CHECK(request.__isset.search_query); + DORIS_CHECK(_search_kind != SearchKind::NORMAL); constexpr auto UINT32_MAX_VALUE = static_cast(std::numeric_limits::max()); - if (vector.offset > UINT32_MAX_VALUE || vector.top_k > UINT32_MAX_VALUE - vector.offset) { - return Status::InvalidArgument("Lance vector search top_k + offset exceeds uint32 range"); + if (_search_kind == SearchKind::VECTOR) { + const auto& vector = request.search_query.vector_search; + if (!vector.__isset.column || vector.column.empty() || + vector.column.find('\0') != std::string::npos) { + return Status::InvalidArgument("Lance vector search requires a non-empty column"); + } + if (!vector.__isset.query_vector) { + return Status::InvalidArgument("Lance vector search requires a query vector"); + } + const auto& query_vector = vector.query_vector; + if (!query_vector.__isset.element_type || !query_vector.__isset.dimension || + !query_vector.__isset.values) { + return Status::InvalidArgument( + "Lance query vector requires element_type, dimension, and values"); + } + if (query_vector.dimension <= 0) { + return Status::InvalidArgument("Lance query vector dimension must be positive: {}", + query_vector.dimension); + } + const auto element_width = lance_vector_element_width(query_vector.element_type); + if (element_width == 0) { + return Status::NotSupported("unsupported Lance query vector element type: {}", + static_cast(query_vector.element_type)); + } + const auto dimension = static_cast(query_vector.dimension); + if (dimension > std::numeric_limits::max() / element_width || + query_vector.values.size() != dimension * element_width) { + return Status::InvalidArgument( + "Lance query vector byte size {} does not match dimension {} and element width " + "{}", + query_vector.values.size(), dimension, element_width); + } + if (!vector.__isset.top_k || vector.top_k <= 0) { + return Status::InvalidArgument("Lance vector search top_k must be positive"); + } + if (!vector.__isset.offset || vector.offset < 0) { + return Status::InvalidArgument("Lance vector search offset must be non-negative"); + } + if (vector.offset > UINT32_MAX_VALUE || vector.top_k > UINT32_MAX_VALUE - vector.offset) { + return Status::InvalidArgument( + "Lance vector search top_k + offset exceeds uint32 range"); + } + } else { + DORIS_CHECK(_search_kind == SearchKind::FULL_TEXT); + const auto& full_text = request.search_query.full_text_search; + if (!full_text.__isset.column || full_text.column.empty() || + full_text.column.find('\0') != std::string::npos) { + return Status::InvalidArgument("Lance full-text search requires a non-empty column"); + } + if (!full_text.__isset.query || full_text.query.empty() || + full_text.query.find('\0') != std::string::npos) { + return Status::InvalidArgument("Lance full-text search requires a non-empty query"); + } + if (!full_text.__isset.top_k || full_text.top_k <= 0) { + return Status::InvalidArgument("Lance full-text search top_k must be positive"); + } + if (!full_text.__isset.offset || full_text.offset < 0) { + return Status::InvalidArgument("Lance full-text search offset must be non-negative"); + } + if (full_text.offset > UINT32_MAX_VALUE || + full_text.top_k > UINT32_MAX_VALUE - full_text.offset) { + return Status::InvalidArgument( + "Lance full-text search top_k + offset exceeds uint32 range"); + } + if (!full_text.__isset.coverage_mode || + (full_text.coverage_mode != TFtsCoverageMode::STRICT && + full_text.coverage_mode != TFtsCoverageMode::INDEX_ONLY)) { + return Status::InvalidArgument( + "Lance full-text search requires STRICT or INDEX_ONLY coverage_mode"); + } + if (full_text.__isset.global_statistics && full_text.global_statistics.empty()) { + return Status::InvalidArgument( + "Lance full-text search global_statistics must not be empty when set"); + } + if (request.__isset.vector_search_options) { + return Status::InvalidArgument( + "Lance full-text search cannot set vector_search_options"); + } } if (request.__isset.search_filter) { @@ -687,22 +553,16 @@ Status LanceTableReader::_validate_external_search_request() const { return Status::InvalidArgument( "external search filter requires format and non-empty payload"); } - switch (filter.format) { - case TSearchFilterFormat::SQL: - if (filter.payload.find('\0') != std::string::npos) { - return Status::InvalidArgument( - "Lance SQL search filter contains an embedded NUL byte"); - } - break; - case TSearchFilterFormat::SUBSTRAIT: - break; - default: + if (filter.format != TSearchFilterFormat::SQL) { return Status::NotSupported("unsupported external search filter format: {}", static_cast(filter.format)); } + if (filter.payload.find('\0') != std::string::npos) { + return Status::InvalidArgument("Lance SQL search filter contains an embedded NUL byte"); + } } - if (request.__isset.vector_search_options) { + if (_search_kind == SearchKind::VECTOR && request.__isset.vector_search_options) { const auto& options = request.vector_search_options; if (options.__isset.nprobes && options.nprobes <= 0) { return Status::InvalidArgument("Lance nprobes must be positive"); @@ -717,7 +577,8 @@ Status LanceTableReader::_validate_external_search_request() const { return Status::OK(); } -Status LanceTableReader::_ensure_dataset_open(const TFileRangeDesc& range) { +Status LanceTableReader::_ensure_dataset_open(const TFileRangeDesc& range, + bool prepare_fts_context) { DatasetKey key; RETURN_IF_ERROR(_dataset_key(range, &key)); if (_dataset == nullptr) { @@ -727,6 +588,10 @@ Status LanceTableReader::_ensure_dataset_open(const TFileRangeDesc& range) { return Status::InvalidArgument( "Lance reader cannot mix dataset snapshots or storage options"); } + if (_search_kind == SearchKind::FULL_TEXT && prepare_fts_context && + _fts_query_context == nullptr) { + RETURN_IF_ERROR(_prepare_fts_query_context()); + } return Status::OK(); } @@ -745,7 +610,31 @@ Status LanceTableReader::_open_dataset(const DatasetKey& key) { static_cast(key.version)); } if (_dataset == nullptr) { - return _lance_error("open Lance dataset"); + return lance_error("open Lance dataset"); + } + return Status::OK(); +} + +Status LanceTableReader::_prepare_fts_query_context() { + DORIS_CHECK(_dataset != nullptr); + DORIS_CHECK(_fts_query_context == nullptr); + DORIS_CHECK(_scan_params != nullptr); + const auto& full_text = + _scan_params->lance_scan_params.external_search_request.search_query.full_text_search; + if (full_text.__isset.global_statistics) { + return Status::NotSupported( + "Lance FE-provided FTS global statistics require a lance-c consumer API"); + } + const auto coverage_mode = full_text.coverage_mode == TFtsCoverageMode::STRICT + ? LANCE_FTS_COVERAGE_STRICT + : LANCE_FTS_COVERAGE_INDEX_ONLY; + // Keep statistics preparation at the reader/scanner lifetime today. A future FE-provided + // opaque statistics payload should enter through this boundary and create the same context, + // leaving segment-scoped scanner execution unchanged. + _fts_query_context = lance_dataset_prepare_fts_query(_dataset, full_text.column.c_str(), + full_text.query.c_str(), 0, coverage_mode); + if (_fts_query_context == nullptr) { + return lance_error("prepare Lance FTS query context"); } return Status::OK(); } @@ -761,26 +650,32 @@ Status LanceTableReader::_open_scanner(const TFileRangeDesc& range) { const auto& column = _projected_columns[idx]; columns.emplace_back(column.name.c_str()); } - if (_vector_search && columns.empty()) { + if (_search_kind != SearchKind::NORMAL && columns.empty()) { // Keep an explicit empty user projection from becoming `nullptr`, which means all dataset - // columns to lance-c. nearest() already returns this optional system column. - columns.emplace_back(DISTANCE_COLUMN.data()); + // columns to lance-c. Search execution already returns its generated result column. + columns.emplace_back(_search_kind == SearchKind::VECTOR ? LANCE_DISTANCE_COLUMN.data() + : LANCE_SCORE_COLUMN.data()); } columns.emplace_back(nullptr); const auto& lance_scan_params = _scan_params->lance_scan_params; - const char* sql_filter = nullptr; - if (_vector_search) { + std::string sql_filter; + std::shared_ptr runtime_filter_sql; + if (_search_kind == SearchKind::NORMAL) { + runtime_filter_sql = + get_or_create_lance_runtime_filter_sql(_conjuncts, _runtime_filter_cache); + } else { const auto& request = lance_scan_params.external_search_request; if (request.__isset.search_filter && request.search_filter.format == TSearchFilterFormat::SQL) { - sql_filter = request.search_filter.payload.c_str(); + sql_filter = request.search_filter.payload; } } LanceScanner* scanner = - lance_scanner_new(_dataset, columns.size() == 1 ? nullptr : columns.data(), sql_filter); + lance_scanner_new(_dataset, columns.size() == 1 ? nullptr : columns.data(), + sql_filter.empty() ? nullptr : sql_filter.c_str()); if (scanner == nullptr) { - return _lance_error("create Lance scanner"); + return lance_error("create Lance scanner"); } std::unique_ptr scanner_guard(scanner); const auto collect_scan_statistics = [](void* callback_ctx, @@ -788,114 +683,93 @@ Status LanceTableReader::_open_scanner(const TFileRangeDesc& range) { LanceTableReader::_collect_scan_statistics(callback_ctx, statistics); }; if (lance_scanner_set_statistics_callback(scanner, collect_scan_statistics, this) != 0) { - return _lance_error("set Lance scanner statistics callback"); + return lance_error("set Lance scanner statistics callback"); } if (_global_rowid_output_idx.has_value() && lance_scanner_with_row_id(scanner, true) != 0) { - return _lance_error("enable Lance row id output"); + return lance_error("enable Lance row id output"); } - if (_scan_params->__isset.lance_scan_params && - lance_scan_params.__isset.lance_substrait_filter && - !lance_scan_params.lance_substrait_filter.empty()) { - const auto& filter = lance_scan_params.lance_substrait_filter; - if (lance_scanner_set_substrait_filter( - scanner, reinterpret_cast(filter.data()), filter.size()) != 0) { - return _lance_error("set Lance Substrait filter"); - } + if (lance_scan_params.__isset.lance_substrait_filter && + lance_scanner_set_substrait_filter( + scanner, + reinterpret_cast(lance_scan_params.lance_substrait_filter.data()), + lance_scan_params.lance_substrait_filter.size()) != 0) { + return lance_error("set Lance Substrait filter"); } - if (_vector_search) { - const auto& request = lance_scan_params.external_search_request; - if (request.__isset.search_filter && - request.search_filter.format == TSearchFilterFormat::SUBSTRAIT) { - const auto& filter = request.search_filter.payload; - if (lance_scanner_set_substrait_filter(scanner, - reinterpret_cast(filter.data()), - filter.size()) != 0) { - return _lance_error("set Lance vector search Substrait filter"); - } + if (runtime_filter_sql != nullptr && !runtime_filter_sql->expression.empty()) { + if (lance_scanner_additional_sql_filter(scanner, runtime_filter_sql->expression.c_str()) != + 0) { + return lance_error("set Lance additional SQL filter"); } + record_lance_runtime_filter_pushdown(_scanner_profile, *runtime_filter_sql); } const auto batch_size = _batch_size > 0 ? _batch_size : _runtime_state->batch_size(); if (lance_scanner_set_batch_size(scanner, static_cast(batch_size)) != 0) { - return _lance_error("set Lance scanner batch size"); + return lance_error("set Lance scanner batch size"); } const auto& lance_params = range.table_format_params.lance_params; - if (lance_params.__isset.fragment_ids && !lance_params.fragment_ids.empty()) { - const auto& thrift_ids = lance_params.fragment_ids; - std::vector fragment_ids; - fragment_ids.reserve(thrift_ids.size()); - for (const auto fragment_id : thrift_ids) { - fragment_ids.emplace_back(static_cast(fragment_id)); - } - if (lance_scanner_set_fragment_ids(scanner, fragment_ids.data(), fragment_ids.size()) != - 0) { - return _lance_error("set Lance scanner fragment ids"); - } - } - if (lance_params.__isset.index_segment_uuids && !lance_params.index_segment_uuids.empty()) { - if (!_vector_search) { - return Status::InvalidArgument( - "Lance index segments are only supported for vector search splits"); - } - constexpr size_t UUID_SIZE = 16; - if (lance_params.index_segment_uuids.size() > - std::numeric_limits::max() / UUID_SIZE) { - return Status::InvalidArgument("too many Lance index segment UUIDs"); - } - std::vector segment_uuids; - segment_uuids.reserve(lance_params.index_segment_uuids.size() * UUID_SIZE); - for (const auto& uuid : lance_params.index_segment_uuids) { - if (uuid.size() != UUID_SIZE) { - return Status::InvalidArgument( - "Lance index segment UUID must contain 16 bytes, got {}", uuid.size()); - } - segment_uuids.insert(segment_uuids.end(), uuid.begin(), uuid.end()); - } - if (lance_scanner_set_index_segments(scanner, segment_uuids.data(), - lance_params.index_segment_uuids.size()) != 0) { - return _lance_error("set Lance scanner index segments"); - } - } - // Ordinary scans may carry a pushed-down LIMIT. The FE only sets it when all predicates are - // pushed into Lance, so the scanner can safely stop after `limit` rows. Vector search manages - // its own top_k limit in _configure_vector_search, so skip it here. - if (!_vector_search && lance_params.__isset.limit && lance_params.limit > 0) { - if (lance_scanner_set_limit(scanner, lance_params.limit) != 0) { - return _lance_error("set Lance scanner limit"); - } - } - if (_vector_search) { - // Distributed vector search always restricts each scanner to an explicit fragment set. - // Tell Lance that this fragment scan is the input to nearest() before installing the - // query. The same prefilter path also applies the TVF search filter, when present. - if (lance_scanner_set_prefilter(scanner, true) != 0) { - return _lance_error("enable Lance vector prefilter"); - } - RETURN_IF_ERROR(_configure_vector_search(scanner)); - const int64_t fragment_count = - lance_params.__isset.fragment_ids - ? static_cast(lance_params.fragment_ids.size()) - : 0; - if (lance_params.__isset.index_segment_uuids && !lance_params.index_segment_uuids.empty()) { - COUNTER_UPDATE(_planned_index_segment_count, - static_cast(lance_params.index_segment_uuids.size())); - COUNTER_UPDATE(_planned_indexed_fragment_count, fragment_count); - } else { - COUNTER_UPDATE(_planned_flat_search_fragment_count, fragment_count); - } + switch (_search_kind) { + case SearchKind::NORMAL: + RETURN_IF_ERROR(_configure_normal_scan(scanner, lance_params)); + break; + case SearchKind::VECTOR: + RETURN_IF_ERROR(_configure_vector_search(scanner, lance_params)); + break; + case SearchKind::FULL_TEXT: + RETURN_IF_ERROR(_configure_full_text_search(scanner, lance_params)); + break; } _scanner = scanner_guard.release(); _scanner_batch_size = batch_size; return Status::OK(); } -Status LanceTableReader::_configure_vector_search(LanceScanner* scanner) const { +Status LanceTableReader::_configure_normal_scan(LanceScanner* scanner, + const TLanceFileDesc& lance_params) const { + DORIS_CHECK(scanner != nullptr); + std::vector fragment_ids; + RETURN_IF_ERROR(parse_fragment_ids(lance_params, &fragment_ids)); + if (!fragment_ids.empty() && + lance_scanner_set_fragment_ids(scanner, fragment_ids.data(), fragment_ids.size()) != 0) { + return lance_error("set Lance scanner fragment ids"); + } + if (lance_params.__isset.index_segment_uuids && !lance_params.index_segment_uuids.empty()) { + return Status::InvalidArgument("normal Lance scan cannot contain index segment UUIDs"); + } + // FE sets this only when every predicate has been pushed into Lance. + if (lance_params.__isset.limit && lance_params.limit > 0 && + lance_scanner_set_limit(scanner, lance_params.limit) != 0) { + return lance_error("set Lance scanner limit"); + } + return Status::OK(); +} + +Status LanceTableReader::_configure_vector_search(LanceScanner* scanner, + const TLanceFileDesc& lance_params) const { DORIS_CHECK(scanner != nullptr); DORIS_CHECK(_scan_params != nullptr); DORIS_CHECK(_scan_params->__isset.lance_scan_params); + std::vector fragment_ids; + RETURN_IF_ERROR(parse_fragment_ids(lance_params, &fragment_ids)); + if (!fragment_ids.empty() && + lance_scanner_set_fragment_ids(scanner, fragment_ids.data(), fragment_ids.size()) != 0) { + return lance_error("set Lance vector scanner fragment ids"); + } + std::vector segment_uuids; + size_t segment_count = 0; + RETURN_IF_ERROR(parse_index_segment_uuids(lance_params, &segment_uuids, &segment_count)); + if (segment_count > 0 && + lance_scanner_set_index_segments(scanner, segment_uuids.data(), segment_count) != 0) { + return lance_error("set Lance vector scanner index segments"); + } + // Fragment-scoped nearest queries require prefiltering before installing the query. The same + // path applies the TVF search filter, when present. + if (lance_scanner_set_prefilter(scanner, true) != 0) { + return lance_error("enable Lance vector prefilter"); + } const auto& lance_scan_params = _scan_params->lance_scan_params; DORIS_CHECK(lance_scan_params.__isset.external_search_request); const auto& request = lance_scan_params.external_search_request; @@ -908,7 +782,7 @@ Status LanceTableReader::_configure_vector_search(LanceScanner* scanner) const { const auto set_nearest = [&](const void* values, LanceDataType type) -> Status { if (lance_scanner_nearest(scanner, vector.column.c_str(), values, dimension, type, candidate_k) != 0) { - return _lance_error("set Lance nearest query"); + return lance_error("set Lance nearest query"); } return Status::OK(); }; @@ -977,7 +851,7 @@ Status LanceTableReader::_configure_vector_search(LanceScanner* scanner) const { static_cast(vector.metric)); } if (lance_scanner_set_metric(scanner, metric) != 0) { - return _lance_error("set Lance vector metric"); + return lance_error("set Lance vector metric"); } } @@ -985,28 +859,68 @@ Status LanceTableReader::_configure_vector_search(LanceScanner* scanner) const { const auto& options = request.vector_search_options; if (options.__isset.nprobes && lance_scanner_set_nprobes(scanner, static_cast(options.nprobes)) != 0) { - return _lance_error("set Lance vector nprobes"); + return lance_error("set Lance vector nprobes"); } if (options.__isset.refine_factor && lance_scanner_set_refine_factor(scanner, static_cast(options.refine_factor)) != 0) { - return _lance_error("set Lance vector refine factor"); + return lance_error("set Lance vector refine factor"); } if (options.__isset.ef && lance_scanner_set_ef(scanner, static_cast(options.ef)) != 0) { - return _lance_error("set Lance vector ef"); + return lance_error("set Lance vector ef"); } if (options.__isset.use_index && lance_scanner_set_use_index(scanner, options.use_index) != 0) { - return _lance_error("set Lance vector use_index"); + return lance_error("set Lance vector use_index"); } } if (lance_scanner_set_offset(scanner, vector.offset) != 0) { - return _lance_error("set Lance vector offset"); + return lance_error("set Lance vector offset"); } if (lance_scanner_set_limit(scanner, vector.top_k) != 0) { - return _lance_error("set Lance vector result limit"); + return lance_error("set Lance vector result limit"); + } + const auto fragment_count = static_cast(fragment_ids.size()); + if (segment_count > 0) { + COUNTER_UPDATE(_planned_index_segment_count, static_cast(segment_count)); + COUNTER_UPDATE(_planned_indexed_fragment_count, fragment_count); + } else { + COUNTER_UPDATE(_planned_flat_search_fragment_count, fragment_count); + } + return Status::OK(); +} + +Status LanceTableReader::_configure_full_text_search(LanceScanner* scanner, + const TLanceFileDesc& lance_params) const { + DORIS_CHECK(scanner != nullptr); + DORIS_CHECK(_fts_query_context != nullptr); + DORIS_CHECK(_scan_params != nullptr); + // FTS fragment IDs describe the selected segment's coverage for planning and profiling. They + // are not installed as a generic fragment filter because lance-c rejects combining one with a + // prepared FTS context; the segment UUID is the execution boundary. + std::vector fragment_ids; + RETURN_IF_ERROR(parse_fragment_ids(lance_params, &fragment_ids)); + std::vector segment_uuids; + size_t segment_count = 0; + RETURN_IF_ERROR(parse_index_segment_uuids(lance_params, &segment_uuids, &segment_count)); + if (segment_count == 0) { + return Status::InvalidArgument( + "Lance full-text search split requires at least one FTS index segment UUID"); + } + const auto& full_text = + _scan_params->lance_scan_params.external_search_request.search_query.full_text_search; + if (lance_scanner_set_fts_query_context(scanner, _fts_query_context) != 0) { + return lance_error("attach Lance FTS query context"); } + if (lance_scanner_set_fts_index_segments(scanner, segment_uuids.data(), segment_count) != 0) { + return lance_error("set Lance FTS scanner index segments"); + } + if (lance_scanner_set_limit(scanner, full_text.top_k + full_text.offset) != 0) { + return lance_error("set Lance FTS scanner candidate limit"); + } + COUNTER_UPDATE(_planned_index_segment_count, static_cast(segment_count)); + COUNTER_UPDATE(_planned_indexed_fragment_count, static_cast(fragment_ids.size())); return Status::OK(); } @@ -1095,6 +1009,10 @@ void LanceTableReader::_close_scanner() { } void LanceTableReader::_close_dataset() { + if (_fts_query_context != nullptr) { + lance_fts_query_context_close(_fts_query_context); + _fts_query_context = nullptr; + } if (_dataset != nullptr) { lance_dataset_close(_dataset); _dataset = nullptr; @@ -1109,7 +1027,7 @@ Status LanceTableReader::_fill_block_from_lance_batch(LanceBatch* batch, Block* ArrowArray array {}; ArrowSchema schema {}; if (lance_batch_to_arrow(batch, &array, &schema) != 0) { - return _lance_error("export Lance batch to Arrow"); + return lance_error("export Lance batch to Arrow"); } auto result = arrow::ImportRecordBatch(&array, &schema); if (!result.ok()) { @@ -1176,11 +1094,12 @@ Status LanceTableReader::_fill_block_from_record_batch( auto& columns = columns_guard.mutable_columns(); for (int arrow_idx = 0; arrow_idx < record_batch->num_columns(); ++arrow_idx) { const auto& field = record_batch->schema()->field(arrow_idx); - if (field->name() == ROW_ID_COLUMN && _global_rowid_output_idx.has_value()) { + if (field->name() == LANCE_ROW_ID_COLUMN && _global_rowid_output_idx.has_value()) { const auto output_idx = *_global_rowid_output_idx; const auto& output_name = _projected_columns[output_idx].name; if (!materialized_columns.emplace(output_name).second) { - return Status::InternalError("Lance returned duplicate column '{}'", ROW_ID_COLUMN); + return Status::InternalError("Lance returned duplicate column '{}'", + LANCE_ROW_ID_COLUMN); } RETURN_IF_ERROR( _append_global_row_ids(record_batch->column(arrow_idx), columns[output_idx])); @@ -1188,9 +1107,10 @@ Status LanceTableReader::_fill_block_from_record_batch( } const auto output_it = _output_name_to_idx.find(field->name()); if (output_it == _output_name_to_idx.end()) { - if (_vector_search && field->name() == DISTANCE_COLUMN) { - // Lance currently auto-projects _distance for nearest queries. It is valid for - // Doris slot pruning to omit that optional result column. + if ((_search_kind == SearchKind::VECTOR && field->name() == LANCE_DISTANCE_COLUMN) || + (_search_kind == SearchKind::FULL_TEXT && field->name() == LANCE_SCORE_COLUMN)) { + // Lance auto-projects the generated search result column. It is valid for Doris + // slot pruning to omit that optional result column. continue; } return Status::InternalError("Lance returned unknown column '{}'", field->name()); @@ -1219,53 +1139,11 @@ Status LanceTableReader::_fill_block_from_record_batch( return Status::OK(); } -// The FE sends these already in Lance's own vocabulary, merged from the catalog properties and -// from whatever the namespace vended. Re-encoding them here would drop every option this list did -// not anticipate, so they are handed to lance-c as they arrive. -Status LanceTableReader::_storage_options(const TFileScanRangeParams* scan_params, - std::vector* options) { - options->clear(); - if (scan_params == nullptr || !scan_params->__isset.lance_scan_params || - !scan_params->lance_scan_params.__isset.lance_storage_options) { - return Status::OK(); - } - const auto& storage_options = scan_params->lance_scan_params.lance_storage_options; - options->reserve(storage_options.size() * 2); - for (const auto& [key, value] : storage_options) { - // These become C strings below, so a NUL would truncate the option here while the FE went - // on using the whole thing, and the two halves would open the dataset with different - // configuration. The FE rejects these on both paths it builds options from - its own - // storage configuration and what a namespace vends - so this is the last line of defence, - // for an FE that predates those checks. Dropping one here instead of failing would just - // recreate the divergence it exists to prevent. - if (key.find('\0') != std::string::npos || value.find('\0') != std::string::npos) { - return Status::InvalidArgument( - "Lance storage option '{}' contains a NUL and cannot reach lance-c", - key.substr(0, key.find('\0'))); - } - options->emplace_back(key); - options->emplace_back(value); - } - return Status::OK(); -} - Status LanceTableReader::_dataset_key(const TFileRangeDesc& range, DatasetKey* key) const { const auto& params = range.table_format_params.lance_params; key->uri = params.dataset_uri; key->version = params.version; - return _storage_options(_scan_params, &key->storage_options); -} - -Status LanceTableReader::_lance_error(std::string_view operation) { - const char* raw_message = lance_last_error_message(); - std::string message = raw_message == nullptr ? "" : raw_message; - if (raw_message != nullptr) { - lance_free_string(raw_message); - } - if (message.empty()) { - message = fmt::format("error_code={}", static_cast(lance_last_error_code())); - } - return Status::InternalError("{} failed: {}", operation, message); + return build_lance_storage_options(_scan_params, &key->storage_options); } } // namespace doris::format::lance diff --git a/be/src/format_v2/table/lance_reader.h b/be/src/format_v2/table/lance_reader.h index 892aaf518e5d83..f3c334fbd785ea 100644 --- a/be/src/format_v2/table/lance_reader.h +++ b/be/src/format_v2/table/lance_reader.h @@ -34,23 +34,20 @@ struct LanceBatch; struct LanceDataset; +struct LanceFtsQueryContext; struct LanceScanner; +namespace doris { +class ShardedKVCache; +} + namespace arrow { class Array; class RecordBatch; -class Schema; } // namespace arrow namespace doris::format::lance { -// Convert every top-level field without discarding unsupported columns. Malformed schemas still -// return an error and leave both output vectors unchanged. DataTypeNothing is the local sentinel -// for a valid Arrow field whose logical type Doris does not support. -Status convert_arrow_schema_to_doris(const std::shared_ptr& arrow_schema, - std::vector* column_names, - std::vector* column_types); - // A FORMAT_LANCE table reader. Unlike file formats such as Parquet, a Lance split is not a // physical-file range. It either selects fragments from a fixed snapshot or scans the whole // latest snapshot, so the dataset is owned by this table reader and each split owns its scanner. @@ -83,11 +80,17 @@ class LanceTableReader final : public TableReader { bool operator==(const DatasetKey&) const = default; }; + Status _resolve_search_kind(); Status _validate_external_search_request() const; - Status _ensure_dataset_open(const TFileRangeDesc& range); + Status _ensure_dataset_open(const TFileRangeDesc& range, bool prepare_fts_context = true); Status _open_dataset(const DatasetKey& key); + Status _prepare_fts_query_context(); Status _open_scanner(const TFileRangeDesc& range); - Status _configure_vector_search(LanceScanner* scanner) const; + Status _configure_normal_scan(LanceScanner* scanner, const TLanceFileDesc& lance_params) const; + Status _configure_vector_search(LanceScanner* scanner, + const TLanceFileDesc& lance_params) const; + Status _configure_full_text_search(LanceScanner* scanner, + const TLanceFileDesc& lance_params) const; // Keep lance-c's anonymous statistics typedef out of this header. _open_scanner installs the // strongly typed C callback adapter before forwarding the borrowed value here. static void _collect_scan_statistics(void* callback_ctx, const void* opaque_statistics); @@ -98,13 +101,11 @@ class LanceTableReader final : public TableReader { Block* block, size_t* rows); Status _append_global_row_ids(const std::shared_ptr& row_ids, MutableColumnPtr& output_column) const; - static Status _storage_options(const TFileScanRangeParams* scan_params, - std::vector* options); Status _dataset_key(const TFileRangeDesc& range, DatasetKey* key) const; - static Status _lance_error(std::string_view operation); LanceDataset* _dataset = nullptr; LanceScanner* _scanner = nullptr; + ShardedKVCache* _runtime_filter_cache = nullptr; std::optional _opened_dataset_key; std::unordered_map _output_name_to_idx; std::optional _global_rowid_output_idx; @@ -126,7 +127,9 @@ class LanceTableReader final : public TableReader { RuntimeProfile::Counter* _index_comparisons = nullptr; std::unordered_map _lance_count_metrics; std::unordered_map _lance_time_metrics; - bool _vector_search = false; + LanceFtsQueryContext* _fts_query_context = nullptr; + enum class SearchKind { NORMAL, VECTOR, FULL_TEXT }; + SearchKind _search_kind = SearchKind::NORMAL; bool _eof = false; }; diff --git a/be/test/format_v2/lance/lance_runtime_filter_helper_test.cpp b/be/test/format_v2/lance/lance_runtime_filter_helper_test.cpp new file mode 100644 index 00000000000000..af65092f83e3ba --- /dev/null +++ b/be/test/format_v2/lance/lance_runtime_filter_helper_test.cpp @@ -0,0 +1,200 @@ +// 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. + +#include "format_v2/lance/lance_runtime_filter_helper.h" + +#include + +#include +#include +#include +#include +#include + +#include "core/data_type/data_type_nullable.h" +#include "core/data_type/data_type_number.h" +#include "core/data_type/data_type_string.h" +#include "core/field.h" +#include "exprs/hybrid_set.h" +#include "exprs/runtime_filter_expr.h" +#include "exprs/vbloom_predicate.h" +#include "exprs/vdirect_in_predicate.h" +#include "exprs/vectorized_fn_call.h" +#include "exprs/vexpr_context.h" +#include "exprs/vliteral.h" +#include "exprs/vslot_ref.h" +#include "format/format_common.h" +#include "runtime/runtime_profile.h" + +namespace doris::format::lance { +namespace { + +TExprNode runtime_in_node() { + TExprNode node; + node.__set_type(std::make_shared()->to_thrift()); + node.__set_node_type(TExprNodeType::IN_PRED); + node.in_predicate.__set_is_not_in(false); + node.__set_opcode(TExprOpcode::FILTER_IN); + node.__set_is_nullable(false); + return node; +} + +VExprContextSPtr wrap_runtime_filter(VExprSPtr impl, const TExprNode& node, int filter_id) { + return VExprContext::create_shared( + RuntimeFilterExpr::create_shared(node, std::move(impl), 0.0, false, filter_id)); +} + +VExprContextSPtr int64_runtime_in(std::string column_name, std::vector values, + int filter_id) { + std::shared_ptr filter(create_set(TYPE_BIGINT, false)); + for (const auto value : values) { + filter->insert(&value); + } + auto node = runtime_in_node(); + auto predicate = VDirectInPredicate::create_shared(node, std::move(filter), true); + predicate->add_child(VSlotRef::create_shared( + 0, 0, -1, make_nullable(std::make_shared()), std::move(column_name))); + return wrap_runtime_filter(std::move(predicate), node, filter_id); +} + +VExprContextSPtr string_runtime_in(std::string column_name, const std::string& value, + int filter_id) { + std::shared_ptr filter(create_set(TYPE_STRING, false)); + StringRef value_ref(value.data(), value.size()); + filter->insert(&value_ref); + auto node = runtime_in_node(); + auto predicate = VDirectInPredicate::create_shared(node, std::move(filter), true); + predicate->add_child(VSlotRef::create_shared( + 0, 0, -1, make_nullable(std::make_shared()), std::move(column_name))); + return wrap_runtime_filter(std::move(predicate), node, filter_id); +} + +VExprContextSPtr int64_runtime_range(std::string column_name, TExprOpcode::type opcode, + int64_t value, int filter_id) { + const auto value_type = std::make_shared(); + const auto nullable_value_type = make_nullable(value_type); + const auto result_type = make_nullable(std::make_shared()); + + TFunctionName function_name; + function_name.__set_function_name(opcode == TExprOpcode::GE ? "ge" : "le"); + TFunction function; + function.__set_name(function_name); + function.__set_binary_type(TFunctionBinaryType::BUILTIN); + function.__set_arg_types({nullable_value_type->to_thrift(), value_type->to_thrift()}); + function.__set_ret_type(result_type->to_thrift()); + function.__set_has_var_args(false); + + TExprNode predicate_node; + predicate_node.__set_node_type(TExprNodeType::BINARY_PRED); + predicate_node.__set_opcode(opcode); + predicate_node.__set_type(result_type->to_thrift()); + predicate_node.__set_fn(function); + predicate_node.__set_num_children(2); + predicate_node.__set_is_nullable(true); + auto predicate = VectorizedFnCall::create_shared(predicate_node); + predicate->add_child( + VSlotRef::create_shared(0, 0, -1, nullable_value_type, std::move(column_name))); + predicate->add_child( + VLiteral::create_shared(value_type, Field::create_field(value))); + + TExprNode wrapper_node; + wrapper_node.__set_type(std::make_shared()->to_thrift()); + wrapper_node.__set_is_nullable(false); + return wrap_runtime_filter(std::move(predicate), wrapper_node, filter_id); +} + +VExprContextSPtr unsupported_bloom_runtime_filter(std::string column_name, int filter_id) { + auto node = runtime_in_node(); + node.__set_node_type(TExprNodeType::BLOOM_PRED); + node.__set_opcode(TExprOpcode::RT_FILTER); + auto predicate = VBloomPredicate::create_shared(node); + predicate->add_child(VSlotRef::create_shared( + 0, 0, -1, make_nullable(std::make_shared()), std::move(column_name))); + return wrap_runtime_filter(std::move(predicate), node, filter_id); +} + +TEST(LanceRuntimeFilterHelperTest, ConvertsSupportedFiltersToLanceSql) { + const VExprContextSPtrs conjuncts { + int64_runtime_in("order`key", {7}, 3), + string_runtime_in("author", "O'Reilly", 5), + int64_runtime_range("score", TExprOpcode::GE, 10, 7), + int64_runtime_range("score", TExprOpcode::LE, 20, 7), + }; + + const auto result = get_or_create_lance_runtime_filter_sql(conjuncts, nullptr); + ASSERT_NE(result, nullptr); + EXPECT_EQ( + "(`order``key` IN (7)) AND (`author` IN ('O''Reilly')) AND " + "(`score` >= 10) AND (`score` <= 20)", + result->expression); + EXPECT_EQ((std::vector {3, 5, 7}), result->pushable_filter_ids); + EXPECT_TRUE(result->skipped_filter_ids.empty()); +} + +TEST(LanceRuntimeFilterHelperTest, RecordsUnsupportedRuntimeFilters) { + const VExprContextSPtrs conjuncts { + int64_runtime_in("id", {2}, 3), + unsupported_bloom_runtime_filter("id", 8), + }; + + const auto result = get_or_create_lance_runtime_filter_sql(conjuncts, nullptr); + ASSERT_NE(result, nullptr); + EXPECT_EQ("(`id` IN (2))", result->expression); + EXPECT_EQ((std::vector {3}), result->pushable_filter_ids); + EXPECT_EQ((std::vector {8}), result->skipped_filter_ids); + + RuntimeProfile profile("lance_runtime_filter_profile"); + record_lance_runtime_filter_pushdown(&profile, *result); + ASSERT_NE(profile.get_info_string("LanceRuntimeFilterPushedIds"), nullptr); + EXPECT_EQ("3", *profile.get_info_string("LanceRuntimeFilterPushedIds")); + ASSERT_NE(profile.get_info_string("LanceRuntimeFilterSkippedIds"), nullptr); + EXPECT_EQ("8", *profile.get_info_string("LanceRuntimeFilterSkippedIds")); +} + +TEST(LanceRuntimeFilterHelperTest, IgnoresNonRuntimeFilterConjuncts) { + const VExprContextSPtrs conjuncts {VExprContext::create_shared(VSlotRef::create_shared( + 0, 0, -1, std::make_shared(), "ordinary_column"))}; + + EXPECT_EQ(nullptr, get_or_create_lance_runtime_filter_sql(conjuncts, nullptr)); +} + +TEST(LanceRuntimeFilterHelperTest, ReusesSnapshotAcrossParallelReaders) { + ShardedKVCache cache(2); + const VExprContextSPtrs first_conjuncts { + int64_runtime_in("id", {2}, 12), + int64_runtime_range("score", TExprOpcode::GE, 10, 13), + }; + const VExprContextSPtrs reordered_conjuncts { + int64_runtime_range("score", TExprOpcode::GE, 10, 13), + int64_runtime_in("id", {2}, 12), + }; + + const auto first = get_or_create_lance_runtime_filter_sql(first_conjuncts, &cache); + const auto reused = get_or_create_lance_runtime_filter_sql(reordered_conjuncts, &cache); + ASSERT_NE(first, nullptr); + ASSERT_NE(reused, nullptr); + EXPECT_EQ(first.get(), reused.get()); + EXPECT_EQ("(`id` IN (2)) AND (`score` >= 10)", reused->expression); + + const auto different = + get_or_create_lance_runtime_filter_sql({int64_runtime_in("id", {2}, 14)}, &cache); + ASSERT_NE(different, nullptr); + EXPECT_NE(first.get(), different.get()); +} + +} // namespace +} // namespace doris::format::lance diff --git a/be/test/format_v2/table/lance_reader_test.cpp b/be/test/format_v2/table/lance_reader_test.cpp index 9a2e802217eeb4..4a6335422bc7ec 100644 --- a/be/test/format_v2/table/lance_reader_test.cpp +++ b/be/test/format_v2/table/lance_reader_test.cpp @@ -57,7 +57,13 @@ #include "core/data_type/data_type_number.h" #include "core/data_type/data_type_varbinary.h" #include "exec/common/endian.h" +#include "exprs/hybrid_set.h" +#include "exprs/runtime_filter_expr.h" +#include "exprs/vdirect_in_predicate.h" #include "exprs/vexpr.h" +#include "exprs/vexpr_context.h" +#include "exprs/vslot_ref.h" +#include "format_v2/lance/lance_reader_helper.h" #include "runtime/runtime_profile.h" #include "runtime/runtime_state.h" #include "storage/utils.h" @@ -166,6 +172,26 @@ ColumnDefinition projected_column(std::string name, PrimitiveType type, bool nul DataTypeFactory::instance().create_data_type(type, nullable)); } +VExprContextSPtr create_int64_runtime_in_conjunct(std::string column_name, + const std::vector& values, + int filter_id) { + std::shared_ptr filter(create_set(TYPE_BIGINT, false)); + for (const auto value : values) { + filter->insert(&value); + } + TExprNode node; + node.__set_type(std::make_shared()->to_thrift()); + node.__set_node_type(TExprNodeType::IN_PRED); + node.in_predicate.__set_is_not_in(false); + node.__set_opcode(TExprOpcode::FILTER_IN); + node.__set_is_nullable(false); + auto predicate = VDirectInPredicate::create_shared(node, std::move(filter), true); + predicate->add_child(VSlotRef::create_shared( + 0, 0, -1, make_nullable(std::make_shared()), std::move(column_name))); + return VExprContext::create_shared( + RuntimeFilterExpr::create_shared(node, std::move(predicate), 0.0, false, filter_id)); +} + void add_output_columns(Block* block, const Columns& columns) { for (const auto& column : columns) { block->insert({column.type->create_column(), column.type, column.name}); @@ -635,7 +661,7 @@ TEST(LanceTableReaderVectorSearchTest, ReadsOnlyGlobalRowIdVirtualColumn) { EXPECT_TRUE(reader.close().ok()); } -TEST(LanceTableReaderFilterTest, PushesFilterOnNonProjectedColumn) { +TEST(LanceTableReaderFilterTest, CombinesStaticSubstraitFilterWithRuntimeFilter) { const std::filesystem::path dataset_uri = "./be/test/format_v2/table/lance/data/all_types.lance"; LanceFixtureInfo fixture; @@ -712,6 +738,85 @@ TEST(LanceTableReaderFilterTest, PushesFilterOnNonProjectedColumn) { std::ranges::sort(labels); EXPECT_EQ((std::vector {"extra", "mixed"}), labels); EXPECT_TRUE(reader.close().ok()); + + // The static Substrait prefilter and a later runtime filter must both remain active. Static + // row_id >= 3 yields {3, 4}; runtime row_id IN (2, 4) yields {2, 4}; their intersection is {4}. + std::string combined_substrait_filter; + ASSERT_TRUE(base64_decode(substrait_filter_base64, &combined_substrait_filter)); + TLanceScanParams combined_lance_scan_params; + combined_lance_scan_params.__set_lance_substrait_filter(std::move(combined_substrait_filter)); + TFileScanRangeParams combined_scan_params; + combined_scan_params.__set_lance_scan_params(std::move(combined_lance_scan_params)); + const Columns row_id_columns {projected_column("row_id", TYPE_BIGINT, false)}; + RuntimeProfile combined_profile("lance_substrait_and_runtime_filter_fixture"); + const auto combined_runtime_filter = create_int64_runtime_in_conjunct("row_id", {2, 4}, 42); + + LanceTableReader combined_reader; + ASSERT_TRUE(init_reader(&combined_reader, row_id_columns, &state, &combined_profile, + &combined_scan_params, {combined_runtime_filter}) + .ok()); + ASSERT_TRUE(prepare_fixture(&combined_reader, dataset_uri, fixture, fixture.fragment_ids).ok()); + + Block combined_block; + add_output_columns(&combined_block, row_id_columns); + std::vector combined_row_ids; + eos = false; + while (!eos) { + ASSERT_TRUE(combined_reader.get_block(&combined_block, &eos).ok()); + if (eos) { + continue; + } + const auto& row_ids = + assert_cast(*combined_block.get_by_position(0).column); + combined_row_ids.insert(combined_row_ids.end(), row_ids.get_data().begin(), + row_ids.get_data().end()); + } + EXPECT_EQ((std::vector {4}), combined_row_ids); + ASSERT_NE(combined_profile.get_info_string("LanceRuntimeFilterPushedIds"), nullptr); + EXPECT_EQ("42", *combined_profile.get_info_string("LanceRuntimeFilterPushedIds")); + EXPECT_TRUE(combined_reader.close().ok()); +} + +TEST(LanceTableReaderFilterTest, PushesRuntimeInFilterIntoLanceScanner) { + const std::filesystem::path dataset_uri = + "./be/test/format_v2/table/lance/data/all_types.lance"; + LanceFixtureInfo fixture; + ASSERT_TRUE(get_fixture_info(dataset_uri, &fixture).ok()); + + const Columns columns {projected_column("row_id", TYPE_BIGINT, false)}; + TQueryOptions query_options; + query_options.__set_batch_size(4); + TQueryGlobals query_globals; + RuntimeState state(query_globals); + state.set_query_options(query_options); + RuntimeProfile profile("lance_runtime_filter_pushdown_fixture"); + TFileScanRangeParams scan_params; + const auto runtime_filter = create_int64_runtime_in_conjunct("row_id", {2, 4}, 41); + + LanceTableReader reader; + ASSERT_TRUE( + init_reader(&reader, columns, &state, &profile, &scan_params, {runtime_filter}).ok()); + ASSERT_TRUE(prepare_fixture(&reader, dataset_uri, fixture, fixture.fragment_ids).ok()); + + Block block; + add_output_columns(&block, columns); + std::vector actual_row_ids; + bool eos = false; + while (!eos) { + ASSERT_TRUE(reader.get_block(&block, &eos).ok()); + if (eos) { + continue; + } + const auto& row_ids = assert_cast(*block.get_by_position(0).column); + actual_row_ids.insert(actual_row_ids.end(), row_ids.get_data().begin(), + row_ids.get_data().end()); + } + std::ranges::sort(actual_row_ids); + EXPECT_EQ((std::vector {2, 4}), actual_row_ids); + ASSERT_NE(profile.get_info_string("LanceRuntimeFilterPushedIds"), nullptr); + EXPECT_EQ("41", *profile.get_info_string("LanceRuntimeFilterPushedIds")); + EXPECT_EQ(profile.get_info_string("LanceRuntimeFilterSkippedIds"), nullptr); + EXPECT_TRUE(reader.close().ok()); } TEST(LanceTableReaderFilterTest, LeavesResidualPredicatesToScanner) { diff --git a/fe/fe-core/src/main/java/org/apache/doris/catalog/BuiltinTableValuedFunctions.java b/fe/fe-core/src/main/java/org/apache/doris/catalog/BuiltinTableValuedFunctions.java index 06b7242d161875..741710801f3b23 100644 --- a/fe/fe-core/src/main/java/org/apache/doris/catalog/BuiltinTableValuedFunctions.java +++ b/fe/fe-core/src/main/java/org/apache/doris/catalog/BuiltinTableValuedFunctions.java @@ -23,6 +23,7 @@ import org.apache.doris.nereids.trees.expressions.functions.table.File; import org.apache.doris.nereids.trees.expressions.functions.table.Frontends; import org.apache.doris.nereids.trees.expressions.functions.table.FrontendsDisks; +import org.apache.doris.nereids.trees.expressions.functions.table.FullTextSearch; import org.apache.doris.nereids.trees.expressions.functions.table.GroupCommit; import org.apache.doris.nereids.trees.expressions.functions.table.Hdfs; import org.apache.doris.nereids.trees.expressions.functions.table.Http; @@ -77,7 +78,8 @@ public class BuiltinTableValuedFunctions implements FunctionHelper { tableValued(ParquetKvMetadata.class, "parquet_kv_metadata"), tableValued(ParquetBloomProbe.class, "parquet_bloom_probe"), tableValued(CdcStream.class, "cdc_stream"), - tableValued(VectorSearch.class, "vector_search") + tableValued(VectorSearch.class, "vector_search"), + tableValued(FullTextSearch.class, "full_text_search") ); public static final BuiltinTableValuedFunctions INSTANCE = new BuiltinTableValuedFunctions(); diff --git a/fe/fe-core/src/main/java/org/apache/doris/datasource/lance/LanceExternalCatalog.java b/fe/fe-core/src/main/java/org/apache/doris/datasource/lance/LanceExternalCatalog.java index c2fa3d1a6ee8a4..33ebc39a30ef85 100644 --- a/fe/fe-core/src/main/java/org/apache/doris/datasource/lance/LanceExternalCatalog.java +++ b/fe/fe-core/src/main/java/org/apache/doris/datasource/lance/LanceExternalCatalog.java @@ -291,7 +291,7 @@ public LanceTableMetadata loadTableMetadata(String dbName, String tableName) { return loadTableMetadata(dbName, tableName, Optional.empty(), false); } - public LanceTableMetadata loadTableMetadataForVectorSearch(String dbName, String tableName) { + public LanceTableMetadata loadTableMetadataForSearch(String dbName, String tableName) { return loadTableMetadata(dbName, tableName, Optional.empty(), true); } diff --git a/fe/fe-core/src/main/java/org/apache/doris/datasource/lance/LanceExternalTable.java b/fe/fe-core/src/main/java/org/apache/doris/datasource/lance/LanceExternalTable.java index acba68b63fcbe5..0d1f124fce6f88 100644 --- a/fe/fe-core/src/main/java/org/apache/doris/datasource/lance/LanceExternalTable.java +++ b/fe/fe-core/src/main/java/org/apache/doris/datasource/lance/LanceExternalTable.java @@ -66,8 +66,8 @@ public LanceTableMetadata loadMetadata() { return ((LanceExternalCatalog) catalog).loadTableMetadata(db.getRemoteName(), remoteName); } - public LanceTableMetadata loadMetadataForVectorSearch() { - return ((LanceExternalCatalog) catalog).loadTableMetadataForVectorSearch( + public LanceTableMetadata loadMetadataForSearch() { + return ((LanceExternalCatalog) catalog).loadTableMetadataForSearch( db.getRemoteName(), remoteName); } diff --git a/fe/fe-core/src/main/java/org/apache/doris/datasource/lance/LanceIndexSegmentInfo.java b/fe/fe-core/src/main/java/org/apache/doris/datasource/lance/LanceIndexSegmentInfo.java index 7fb373f9246035..f0f1667cbee573 100644 --- a/fe/fe-core/src/main/java/org/apache/doris/datasource/lance/LanceIndexSegmentInfo.java +++ b/fe/fe-core/src/main/java/org/apache/doris/datasource/lance/LanceIndexSegmentInfo.java @@ -17,27 +17,32 @@ package org.apache.doris.datasource.lance; +import org.lance.index.IndexType; + import java.util.ArrayList; import java.util.Collections; import java.util.List; +import java.util.Objects; import java.util.Optional; import java.util.UUID; -/** Immutable metadata for one physical segment of a logical Lance vector index. */ +/** Immutable metadata for one physical segment of a logical Lance search index. */ public final class LanceIndexSegmentInfo { private final UUID uuid; private final String indexName; private final List fieldIds; private final List fragmentIds; + private final IndexType indexType; private final String metric; public LanceIndexSegmentInfo(UUID uuid, String indexName, List fieldIds, - List fragmentIds, String metric) { + List fragmentIds, IndexType indexType, String metric) { this.uuid = uuid; this.indexName = indexName; this.fieldIds = Collections.unmodifiableList(new ArrayList<>(fieldIds)); this.fragmentIds = fragmentIds == null ? null : Collections.unmodifiableList(new ArrayList<>(fragmentIds)); + this.indexType = Objects.requireNonNull(indexType, "indexType must not be null"); this.metric = metric; } @@ -56,14 +61,26 @@ public List getFieldIds() { /** * Returns the fragment bitmap recorded in the manifest. * - *

Legacy index segments may not have a bitmap. Callers must not infer coverage from the - * segment's dataset version in that case. + *

Index segments without a fragment bitmap have unknown coverage. Callers must not infer + * coverage from the segment's dataset version in that case. */ public Optional> getFragmentIds() { return Optional.ofNullable(fragmentIds); } - /** Returns the normalized Lance metric name, or an empty optional for legacy metadata. */ + public IndexType getIndexType() { + return indexType; + } + + public boolean isVectorIndex() { + return indexType.getValue() >= IndexType.VECTOR.getValue(); + } + + public boolean isFullTextIndex() { + return indexType == IndexType.INVERTED; + } + + /** Returns the normalized Lance metric name when the index metadata supplies one. */ public Optional getMetric() { return Optional.ofNullable(metric); } diff --git a/fe/fe-core/src/main/java/org/apache/doris/datasource/lance/LanceMetadataLoader.java b/fe/fe-core/src/main/java/org/apache/doris/datasource/lance/LanceMetadataLoader.java index 9c4803af7100d6..a3c27e7b6bbdf6 100644 --- a/fe/fe-core/src/main/java/org/apache/doris/datasource/lance/LanceMetadataLoader.java +++ b/fe/fe-core/src/main/java/org/apache/doris/datasource/lance/LanceMetadataLoader.java @@ -72,7 +72,7 @@ public static LanceTableMetadata loadLatest(String datasetUri, datasetUri, lanceStorageOptions, OptionalLong.empty(), allocator, false); } - /** Loads the latest fixed snapshot together with vector index segment coverage. */ + /** Loads the latest fixed snapshot together with search-index segment coverage. */ public static LanceTableMetadata loadLatestWithIndexSegments( String datasetUri, Map lanceStorageOptions, BufferAllocator allocator) throws Exception { return loadInternal( @@ -109,7 +109,7 @@ private static LanceTableMetadata loadInternal(String datasetUri, Map lanceFieldIds = loadIndexSegments ? loadTopLevelFieldIds(dataset) : Collections.emptyMap(); List indexSegments = loadIndexSegments - ? loadVectorIndexSegments(dataset) : Collections.emptyList(); + ? loadSearchIndexSegments(dataset) : Collections.emptyList(); return loadIndexSegments ? LanceTableMetadata.withIndexSegments(datasetUri, resolvedVersion, dataset.getSchema(), fragments, lanceFieldIds, @@ -134,12 +134,13 @@ private static Map loadTopLevelFieldIds(Dataset dataset) { return result; } - private static List loadVectorIndexSegments(Dataset dataset) { + private static List loadSearchIndexSegments(Dataset dataset) { List result = new ArrayList<>(); for (IndexDescription description : dataset.describeIndices()) { String metric = parseMetric(description.getDetailsJson()); for (Index segment : description.getSegments()) { - if (segment.indexType() == null || segment.indexType().getValue() < 100) { + if (segment.indexType() == null || (segment.indexType().getValue() < 100 + && segment.indexType() != org.lance.index.IndexType.INVERTED)) { continue; } List fragmentIds = segment.fragments() @@ -152,7 +153,7 @@ private static List loadVectorIndexSegments(Dataset datas }) .orElse(null); result.add(new LanceIndexSegmentInfo(segment.uuid(), description.getName(), - description.getFieldIds(), fragmentIds, metric)); + description.getFieldIds(), fragmentIds, segment.indexType(), metric)); } } return result; @@ -166,8 +167,8 @@ private static String parseMetric(String detailsJson) { JsonNode metric = JsonUtil.readTree(detailsJson).get("metric_type"); return metric == null || !metric.isTextual() ? null : metric.asText().toUpperCase(); } catch (RuntimeException e) { - // Index details are optional compatibility metadata. An unknown legacy encoding should - // disable metric-sensitive segment planning rather than prevent ordinary table access. + // Index details are optional metadata. Malformed details disable metric-sensitive + // segment planning rather than preventing ordinary table access. return null; } } diff --git a/fe/fe-core/src/main/java/org/apache/doris/datasource/lance/source/IndexSegmentSplitPlan.java b/fe/fe-core/src/main/java/org/apache/doris/datasource/lance/source/IndexSegmentSplitPlan.java index c8c61fa46a33d1..b890242ddff16c 100644 --- a/fe/fe-core/src/main/java/org/apache/doris/datasource/lance/source/IndexSegmentSplitPlan.java +++ b/fe/fe-core/src/main/java/org/apache/doris/datasource/lance/source/IndexSegmentSplitPlan.java @@ -26,7 +26,7 @@ import java.util.Set; import java.util.UUID; -/** Builds vector-search splits from physical Lance index segments and unindexed fragments. */ +/** Builds external-search splits from physical Lance index segments and optional fragments. */ final class IndexSegmentSplitPlan { private final String datasetUri; private final long version; diff --git a/fe/fe-core/src/main/java/org/apache/doris/datasource/lance/source/LanceScanNode.java b/fe/fe-core/src/main/java/org/apache/doris/datasource/lance/source/LanceScanNode.java index 394cb6ebf068a1..16a20a6176218a 100644 --- a/fe/fe-core/src/main/java/org/apache/doris/datasource/lance/source/LanceScanNode.java +++ b/fe/fe-core/src/main/java/org/apache/doris/datasource/lance/source/LanceScanNode.java @@ -39,10 +39,13 @@ import org.apache.doris.thrift.TExternalSearchRequest; import org.apache.doris.thrift.TFileFormatType; import org.apache.doris.thrift.TFileRangeDesc; +import org.apache.doris.thrift.TFtsCoverageMode; +import org.apache.doris.thrift.TFullTextSearchParams; import org.apache.doris.thrift.TLanceFileDesc; import org.apache.doris.thrift.TLanceScanParams; import org.apache.doris.thrift.TTableFormatFileDesc; import org.apache.doris.thrift.TVectorMetric; +import org.apache.doris.thrift.TVectorSearchOptions; import org.apache.doris.thrift.TVectorSearchParams; import java.nio.ByteBuffer; @@ -61,14 +64,22 @@ * Keeping them in one node prevents those common parts from drifting apart. The search request is * also an explicit mode marker. Ordinary scans are split by fragment. Indexed vector searches are * split by physical index segment, with uncovered fragments retained as flat-search fallbacks. - * Each search split produces local candidates; a Doris TopN above this scan merges them into the - * requested snapshot-wide result. + * Full-text searches are split only by committed inverted-index segments, with coverage governed + * by the request's STRICT or INDEX_ONLY mode. Each search split produces local candidates; a Doris + * TopN above this scan merges them into the requested snapshot-wide result. */ public class LanceScanNode extends FileQueryScanNode { + private enum SearchKind { + NORMAL, + VECTOR, + FULL_TEXT + } + private LanceExternalTable lanceTable; private LanceTableMetadata plannedMetadata; - private int vectorFieldId = -1; - private TExternalSearchRequest externalSearchRequest; + private final int searchFieldId; + private final TExternalSearchRequest externalSearchRequest; + private final SearchKind searchKind; private byte[] lanceSubstraitFilter = new byte[0]; private String lancePushdownPredicate = ""; private long plannedVersion = -1; @@ -81,38 +92,45 @@ public LanceScanNode(PlanNodeId id, TupleDescriptor desc, boolean needCheckColum SessionVariable sessionVariable, ScanContext scanContext) { super(id, desc, "LANCE_SCAN_NODE", StatisticalType.LANCE_SCAN_NODE, scanContext, needCheckColumnPriv, sessionVariable); + this.searchFieldId = -1; + this.externalSearchRequest = null; + this.searchKind = SearchKind.NORMAL; } /** * Creates the search mode of this node. * *

The tuple descriptor belongs to a FunctionGenTable and contains generated columns such as - * {@code _distance}. Therefore the real Lance table and the metadata snapshot selected while - * analyzing the TVF must be passed separately. + * {@code _distance} or {@code _score}. Therefore the real Lance table and the metadata + * snapshot selected while analyzing the TVF must be passed separately. */ - public static LanceScanNode forVectorSearch(PlanNodeId id, TupleDescriptor desc, - LanceExternalTable lanceTable, LanceTableMetadata plannedMetadata, int vectorFieldId, + public static LanceScanNode forExternalSearch(PlanNodeId id, TupleDescriptor desc, + LanceExternalTable lanceTable, LanceTableMetadata plannedMetadata, int searchFieldId, TExternalSearchRequest externalSearchRequest, SessionVariable sessionVariable) { - return new LanceScanNode(id, desc, lanceTable, plannedMetadata, vectorFieldId, + return new LanceScanNode(id, desc, lanceTable, plannedMetadata, searchFieldId, externalSearchRequest, sessionVariable); } private LanceScanNode(PlanNodeId id, TupleDescriptor desc, LanceExternalTable lanceTable, - LanceTableMetadata plannedMetadata, int vectorFieldId, + LanceTableMetadata plannedMetadata, int searchFieldId, TExternalSearchRequest externalSearchRequest, SessionVariable sessionVariable) { super(id, desc, "LANCE_SCAN_NODE", StatisticalType.LANCE_SCAN_NODE, ScanContext.builder().clusterName(sessionVariable.resolveCloudClusterName()).build(), false, sessionVariable); this.lanceTable = lanceTable; this.plannedMetadata = plannedMetadata; - this.vectorFieldId = vectorFieldId; + this.searchFieldId = searchFieldId; + if (externalSearchRequest == null) { + throw new IllegalArgumentException("Lance external search request must not be null"); + } this.externalSearchRequest = externalSearchRequest.deepCopy(); + this.searchKind = resolveSearchKind(this.externalSearchRequest); } @Override protected void doInitialize() throws UserException { List sourceColumns; - if (isExternalSearch()) { + if (searchKind != SearchKind.NORMAL) { sourceColumns = desc.getTable().getColumns(); } else { lanceTable = (LanceExternalTable) desc.getTable(); @@ -123,11 +141,11 @@ protected void doInitialize() throws UserException { super.doInitialize(); ExternalUtil.initSchemaInfo(params, -1L, sourceColumns); - if (isExternalSearch()) { + if (searchKind != SearchKind.NORMAL) { // Search output comes from the FunctionGenTable because it adds generated columns such - // as _distance. The real Lance table is still retained for storage and metadata access. + // as _distance or _score. The real Lance table is retained for storage and metadata. getOrCreateLanceScanParams() - .setExternalSearchRequest(createFragmentSearchRequest(externalSearchRequest)); + .setExternalSearchRequest(createSplitSearchRequest()); } } @@ -153,9 +171,9 @@ private boolean canPushDownLimit() { @Override protected void convertPredicate() { - if (isExternalSearch()) { + if (searchKind != SearchKind.NORMAL) { // The TVF "filter" property is already serialized in externalSearchRequest and is - // evaluated by Lance before vector search. Outer WHERE conjuncts have different + // evaluated by Lance before candidate search. Outer WHERE conjuncts have different // semantics: keep them as Doris scan residuals. Each fragment first returns its Lance // ANN candidates, then Doris evaluates these conjuncts before the local/global TopN. } else { @@ -186,20 +204,31 @@ public List getSplits(int numBackends) throws UserException { LanceTableMetadata metadata = plannedMetadata; plannedVersion = metadata.getVersion(); plannedFragments = metadata.getFragments().size(); - plannedUnindexedFragments = isExternalSearch() ? plannedFragments : 0; + plannedUnindexedFragments = searchKind == SearchKind.NORMAL ? 0 : plannedFragments; plannedIndexSegments = 0; plannedIndexFragments = 0; - if (isExternalSearch() && plannedVersion <= 0) { + if (searchKind != SearchKind.NORMAL && plannedVersion <= 0) { throw new UserException( - "Lance vector search requires a fixed positive dataset version"); + "Lance external search requires a fixed positive dataset version"); } Map visibleFragments = getVisibleFragments(metadata); - if (isExternalSearch() && shouldUseIndex()) { - Optional> indexSplits = createIndexSegmentSplits(metadata, visibleFragments); - if (indexSplits.isPresent()) { - return indexSplits.get(); - } + switch (searchKind) { + case FULL_TEXT: + return createFullTextIndexSegmentSplits(metadata, visibleFragments); + case VECTOR: + if (isVectorIndexEnabled()) { + Optional> indexSplits = createVectorIndexSegmentSplits( + metadata, visibleFragments); + if (indexSplits.isPresent()) { + return indexSplits.get(); + } + } + break; + case NORMAL: + break; + default: + throw new IllegalStateException("Unsupported Lance search kind " + searchKind); } return createFragmentSplits(metadata, visibleFragments); } @@ -236,25 +265,25 @@ private List createFragmentSplits(LanceTableMetadata metadata, return splits; } - private Optional> createIndexSegmentSplits(LanceTableMetadata metadata, + private Optional> createVectorIndexSegmentSplits(LanceTableMetadata metadata, Map visibleFragments) throws UserException { if (metadata.getIndexSegments().isEmpty()) { return Optional.empty(); } TVectorSearchParams vectorSearchParam = externalSearchRequest.getSearchQuery().getVectorSearch(); - if (vectorFieldId < 0) { + if (searchFieldId < 0) { throw new UserException("Lance vector column '" + vectorSearchParam.getColumn() + "' has no field ID in the Lance schema"); } - List matchingSegments = selectIndexSegments( - metadata.getIndexSegments(), vectorFieldId); + List matchingSegments = selectVectorIndexSegments( + metadata.getIndexSegments(), searchFieldId); if (matchingSegments.isEmpty() || !metricMatches(vectorSearchParam, matchingSegments)) { return Optional.empty(); } Optional indexPlan = planIndexSegments( - metadata, matchingSegments, visibleFragments); + metadata, matchingSegments, visibleFragments, false); if (!indexPlan.isPresent()) { return Optional.empty(); } @@ -266,12 +295,44 @@ private Optional> createIndexSegmentSplits(LanceTableMetadata metada return Optional.of(plan.buildSplits()); } - private static List selectIndexSegments( - List indexSegments, int vectorFieldId) { + private List createFullTextIndexSegmentSplits(LanceTableMetadata metadata, + Map visibleFragments) throws UserException { + TFullTextSearchParams fullText = + externalSearchRequest.getSearchQuery().getFullTextSearch(); + if (searchFieldId < 0) { + throw new UserException("Lance full-text column '" + fullText.getColumn() + + "' has no field ID in the Lance schema"); + } + List matchingSegments = selectFullTextIndexSegments( + metadata.getIndexSegments(), searchFieldId, fullText.getColumn()); + if (matchingSegments.isEmpty()) { + throw new UserException("No committed Lance FTS index exists for column '" + + fullText.getColumn() + "' at dataset version " + metadata.getVersion()); + } + IndexSegmentSplitPlan plan = planIndexSegments( + metadata, matchingSegments, visibleFragments, true) + .orElseThrow(() -> new UserException("Lance FTS index for column '" + + fullText.getColumn() + "' has no visible indexed fragments at dataset version " + + metadata.getVersion())); + plannedIndexSegments = plan.splitCount(); + plannedIndexFragments = plan.indexSegmentFragmentCount(); + plannedUnindexedFragments = plannedFragments - plannedIndexFragments; + if (fullText.getCoverageMode() == TFtsCoverageMode.STRICT + && plannedUnindexedFragments != 0) { + throw new UserException("Lance FTS coverage_mode=STRICT requires every fragment at " + + "dataset version " + metadata.getVersion() + " to be indexed; column '" + + fullText.getColumn() + "' has " + plannedUnindexedFragments + + " unindexed fragments. Rebuild the index or use coverage_mode=index_only"); + } + return plan.buildSplits(); + } + + private static List selectVectorIndexSegments( + List indexSegments, int fieldId) { List selectedSegments = new ArrayList<>(); String selectedIndexName = null; for (LanceIndexSegmentInfo segment : indexSegments) { - if (!segment.getFieldIds().contains(vectorFieldId)) { + if (!segment.isVectorIndex() || !segment.getFieldIds().contains(fieldId)) { continue; } if (selectedIndexName == null) { @@ -284,19 +345,52 @@ private static List selectIndexSegments( return selectedSegments; } + private static List selectFullTextIndexSegments( + List indexSegments, int fieldId, String column) + throws UserException { + List selectedSegments = new ArrayList<>(); + String selectedIndexName = null; + for (LanceIndexSegmentInfo segment : indexSegments) { + if (!segment.isFullTextIndex() || !segment.getFieldIds().contains(fieldId)) { + continue; + } + if (selectedIndexName == null) { + selectedIndexName = segment.getIndexName(); + } else if (!selectedIndexName.equals(segment.getIndexName())) { + throw new UserException("Multiple Lance FTS indexes exist for column '" + column + + "'; distributed FTS requires one unambiguous logical index"); + } + selectedSegments.add(segment); + } + return selectedSegments; + } + private static Optional planIndexSegments( LanceTableMetadata metadata, List indexSegments, - Map visibleFragments) { + Map visibleFragments, + boolean requireKnownCoverage) throws UserException { IndexSegmentSplitPlan plan = new IndexSegmentSplitPlan( metadata.getDatasetUri(), metadata.getVersion(), indexSegments.size()); for (LanceIndexSegmentInfo segment : indexSegments) { Optional> segmentFragments = segment.getFragmentIds(); if (!segmentFragments.isPresent()) { + if (requireKnownCoverage) { + throw new UserException("Lance FTS segment " + segment.getUuid() + + " has no fragment coverage metadata"); + } return Optional.empty(); } List visibleIndexSegmentFragmentIds = effectiveFragmentIds( segmentFragments.get(), visibleFragments); + if (requireKnownCoverage) { + for (Long fragmentId : visibleIndexSegmentFragmentIds) { + if (plan.isCoveredByIndexSegment(fragmentId)) { + throw new UserException("Lance FTS fragment " + fragmentId + + " is covered by multiple physical index segments"); + } + } + } if (!visibleIndexSegmentFragmentIds.isEmpty()) { plan.addIndexSegmentSplit( segment.getUuid(), visibleIndexSegmentFragmentIds, @@ -339,10 +433,13 @@ private static void appendUnindexedFragmentSplits(IndexSegmentSplitPlan plan, } } - private boolean shouldUseIndex() { - return !externalSearchRequest.isSetVectorSearchOptions() - || !externalSearchRequest.getVectorSearchOptions().isSetUseIndex() - || externalSearchRequest.getVectorSearchOptions().isUseIndex(); + private boolean isVectorIndexEnabled() { + // default use_index is true + if (!externalSearchRequest.isSetVectorSearchOptions()) { + return true; + } + TVectorSearchOptions options = externalSearchRequest.getVectorSearchOptions(); + return !options.isSetUseIndex() || options.isUseIndex(); } private static boolean metricMatches(TVectorSearchParams vector, @@ -372,11 +469,15 @@ protected void setScanParams(TFileRangeDesc rangeDesc, Split split) { if (lanceSplit.getFragmentIds().isEmpty()) { throw new IllegalArgumentException("Lance scan split must contain fragments"); } - if (!isExternalSearch() && (lanceSplit.getFragmentIds().size() != 1 + if (searchKind == SearchKind.NORMAL && (lanceSplit.getFragmentIds().size() != 1 || lanceSplit.hasIndexSegmentUuids())) { throw new IllegalArgumentException( "Ordinary Lance scan split must contain one fragment and no index segment"); } + if (searchKind == SearchKind.FULL_TEXT && !lanceSplit.hasIndexSegmentUuids()) { + throw new IllegalArgumentException( + "Lance full-text search split must contain an FTS index segment"); + } lanceParams.setFragmentIds(lanceSplit.getFragmentIds()); if (lanceSplit.hasIndexSegmentUuids()) { List uuids = new ArrayList<>(lanceSplit.getIndexSegmentUuids().size()); @@ -390,8 +491,8 @@ protected void setScanParams(TFileRangeDesc rangeDesc, Split split) { lanceParams.setIndexSegmentUuids(uuids); } // Push LIMIT into each ordinary fragment scanner only when it is safe to truncate that - // fragment early. Vector search uses its own per-split candidate bound. - if (!isExternalSearch() && canPushDownLimit()) { + // fragment early. External searches use their own per-split candidate bound. + if (searchKind == SearchKind.NORMAL && canPushDownLimit()) { lanceParams.setLimit(getLimit()); } @@ -413,7 +514,7 @@ protected List getPathPartitionKeys() { @Override protected TableIf getTargetTable() { - if (isExternalSearch()) { + if (searchKind != SearchKind.NORMAL) { // In search mode desc.getTable() is a FunctionGenTable, but default-value expressions // and storage access still belong to the underlying Lance table. return lanceTable; @@ -432,13 +533,28 @@ protected Map getLocationProperties() { @Override public String getNodeExplainString(String prefix, TExplainLevel detailLevel) { StringBuilder result = new StringBuilder(super.getNodeExplainString(prefix, detailLevel)); - if (isExternalSearch()) { - TVectorSearchParams vector = externalSearchRequest.getSearchQuery().getVectorSearch(); - result.append(prefix).append("externalSearchType=VECTOR\n"); - result.append(prefix).append("lanceVectorColumn=").append(vector.getColumn()).append("\n"); - result.append(prefix).append("lanceMetric=") - .append(vector.isSetMetric() ? metricName(vector.getMetric()) : "default") - .append("\n"); + if (searchKind != SearchKind.NORMAL) { + if (searchKind == SearchKind.VECTOR) { + TVectorSearchParams vector = + externalSearchRequest.getSearchQuery().getVectorSearch(); + result.append(prefix).append("externalSearchType=VECTOR\n"); + result.append(prefix).append("lanceVectorColumn=") + .append(vector.getColumn()).append("\n"); + result.append(prefix).append("lanceMetric=") + .append(vector.isSetMetric() ? metricName(vector.getMetric()) : "default") + .append("\n"); + } else { + if (searchKind != SearchKind.FULL_TEXT) { + throw new IllegalStateException("Unsupported Lance search kind " + searchKind); + } + TFullTextSearchParams fullText = + externalSearchRequest.getSearchQuery().getFullTextSearch(); + result.append(prefix).append("externalSearchType=FULL_TEXT\n"); + result.append(prefix).append("lanceFullTextColumn=") + .append(fullText.getColumn()).append("\n"); + result.append(prefix).append("lanceFtsCoverageMode=") + .append(fullText.getCoverageMode()).append("\n"); + } result.append(prefix).append("lanceVersion=") .append(plannedMetadata.getVersion()).append("\n"); result.append(prefix).append("lanceSearchFragments=") @@ -465,19 +581,40 @@ public String getNodeExplainString(String prefix, TExplainLevel detailLevel) { return result.toString(); } - private boolean isExternalSearch() { - return externalSearchRequest != null; + TExternalSearchRequest createSplitSearchRequest() { + TExternalSearchRequest splitRequest = externalSearchRequest.deepCopy(); + // Every split must retain enough rows for the later global OFFSET/LIMIT. Applying the + // logical offset independently inside each split could discard rows that belong to the + // snapshot-wide result. + switch (searchKind) { + case VECTOR: + TVectorSearchParams vector = splitRequest.getSearchQuery().getVectorSearch(); + vector.setTopK(vector.getTopK() + vector.getOffset()); + vector.setOffset(0); + break; + case FULL_TEXT: + TFullTextSearchParams fullText = splitRequest.getSearchQuery().getFullTextSearch(); + fullText.setTopK(fullText.getTopK() + fullText.getOffset()); + fullText.setOffset(0); + break; + case NORMAL: + default: + throw new IllegalStateException("Cannot create a search split for " + searchKind); + } + return splitRequest; } - static TExternalSearchRequest createFragmentSearchRequest(TExternalSearchRequest searchRequest) { - TExternalSearchRequest fragmentRequest = searchRequest.deepCopy(); - TVectorSearchParams vector = fragmentRequest.getSearchQuery().getVectorSearch(); - // Every fragment must retain enough rows for the later global OFFSET/LIMIT. Applying the - // logical offset independently inside each fragment could discard rows that belong to the - // snapshot-wide result. - vector.setTopK(vector.getTopK() + vector.getOffset()); - vector.setOffset(0); - return fragmentRequest; + private static SearchKind resolveSearchKind(TExternalSearchRequest searchRequest) { + if (!searchRequest.isSetSearchQuery()) { + throw new IllegalArgumentException("Lance external search request requires search_query"); + } + boolean hasVector = searchRequest.getSearchQuery().isSetVectorSearch(); + boolean hasFullText = searchRequest.getSearchQuery().isSetFullTextSearch(); + if (hasVector == hasFullText) { + throw new IllegalArgumentException( + "Lance external search query must set exactly one search kind"); + } + return hasVector ? SearchKind.VECTOR : SearchKind.FULL_TEXT; } private static String metricName(TVectorMetric metric) { diff --git a/fe/fe-core/src/main/java/org/apache/doris/nereids/processor/post/materialize/MaterializeProbeVisitor.java b/fe/fe-core/src/main/java/org/apache/doris/nereids/processor/post/materialize/MaterializeProbeVisitor.java index 8336105a004819..b8d1830a1dc590 100644 --- a/fe/fe-core/src/main/java/org/apache/doris/nereids/processor/post/materialize/MaterializeProbeVisitor.java +++ b/fe/fe-core/src/main/java/org/apache/doris/nereids/processor/post/materialize/MaterializeProbeVisitor.java @@ -38,6 +38,7 @@ import org.apache.doris.nereids.trees.plans.physical.PhysicalTVFRelation; import org.apache.doris.nereids.trees.plans.visitor.DefaultPlanVisitor; import org.apache.doris.qe.SessionVariable; +import org.apache.doris.tablefunction.FullTextSearchTableValuedFunction; import org.apache.doris.tablefunction.VectorSearchTableValuedFunction; import com.google.common.collect.ImmutableSet; @@ -127,7 +128,7 @@ boolean checkRelationTableSupportedType(PhysicalCatalogRelation relation) { } boolean checkTVFRelationTableSupportedType(PhysicalTVFRelation tvfRelation) { - if (isVectorSearch(tvfRelation)) { + if (isLanceExternalSearch(tvfRelation)) { return true; } @@ -144,8 +145,10 @@ boolean checkTVFRelationTableSupportedType(PhysicalTVFRelation tvfRelation) { return false; } - private boolean isVectorSearch(PhysicalTVFRelation tvfRelation) { - return VectorSearchTableValuedFunction.NAME.equals(tvfRelation.getFunction().getName()); + private boolean isLanceExternalSearch(PhysicalTVFRelation tvfRelation) { + String functionName = tvfRelation.getFunction().getName(); + return VectorSearchTableValuedFunction.NAME.equals(functionName) + || FullTextSearchTableValuedFunction.NAME.equals(functionName); } @Override @@ -185,7 +188,7 @@ public Optional visitPhysicalTVFRelation( PhysicalTVFRelation tvfRelation, ProbeContext context) { // The first Lance implementation fetches top-level columns by row ID. Keep nested // sub-column projections in the search phase until take_rows supports access paths. - if (isVectorSearch(tvfRelation) && context.slot.hasSubColPath()) { + if (isLanceExternalSearch(tvfRelation) && context.slot.hasSubColPath()) { return Optional.empty(); } if (checkTVFRelationTableSupportedType(tvfRelation) && tvfRelation.getOutput().contains(context.slot) diff --git a/fe/fe-core/src/main/java/org/apache/doris/nereids/rules/analysis/BindExpression.java b/fe/fe-core/src/main/java/org/apache/doris/nereids/rules/analysis/BindExpression.java index 41cce5539a30d3..d145e09c875e4e 100644 --- a/fe/fe-core/src/main/java/org/apache/doris/nereids/rules/analysis/BindExpression.java +++ b/fe/fe-core/src/main/java/org/apache/doris/nereids/rules/analysis/BindExpression.java @@ -68,6 +68,7 @@ import org.apache.doris.nereids.trees.expressions.functions.generator.Unnest; import org.apache.doris.nereids.trees.expressions.functions.scalar.ElementAt; import org.apache.doris.nereids.trees.expressions.functions.scalar.GroupingScalarFunction; +import org.apache.doris.nereids.trees.expressions.functions.table.FullTextSearch; import org.apache.doris.nereids.trees.expressions.functions.table.TableValuedFunction; import org.apache.doris.nereids.trees.expressions.functions.table.VectorSearch; import org.apache.doris.nereids.trees.expressions.literal.IntegerLikeLiteral; @@ -116,6 +117,7 @@ import org.apache.doris.nereids.util.TypeCoercionUtils; import org.apache.doris.nereids.util.Utils; import org.apache.doris.qe.SqlModeHelper; +import org.apache.doris.tablefunction.FullTextSearchTableValuedFunction; import org.apache.doris.tablefunction.VectorSearchTableValuedFunction; import com.google.common.base.Joiner; @@ -1788,7 +1790,8 @@ private Plan bindTableValuedFunction(MatchingContext ctx) { TableValuedFunction tableValuedFunction = (TableValuedFunction) bindResult.first; LogicalTVFRelation relation = new LogicalTVFRelation( unboundTVFRelation.getRelationId(), tableValuedFunction, ImmutableList.of()); - if (!(tableValuedFunction instanceof VectorSearch)) { + if (!(tableValuedFunction instanceof VectorSearch) + && !(tableValuedFunction instanceof FullTextSearch)) { return relation; } @@ -1796,17 +1799,31 @@ private Plan bindTableValuedFunction(MatchingContext ctx) { // relation with a Doris TopN to merge them into the snapshot-wide result. The predicate // pushdown rules move an outer WHERE below this synthetic TopN, where it is evaluated as // a Doris scan residual after each fragment's Lance search and before the global TopN. - VectorSearchTableValuedFunction vectorSearch = - (VectorSearchTableValuedFunction) tableValuedFunction.getCatalogFunction(); - Slot distance = relation.getOutput().stream() - .filter(slot -> slot.getName().equalsIgnoreCase( - VectorSearchTableValuedFunction.DISTANCE_COLUMN)) + if (tableValuedFunction instanceof VectorSearch) { + VectorSearchTableValuedFunction vectorSearch = + (VectorSearchTableValuedFunction) tableValuedFunction.getCatalogFunction(); + Slot distance = requireSearchResultSlot( + relation, VectorSearchTableValuedFunction.DISTANCE_COLUMN, + VectorSearchTableValuedFunction.NAME); + return new LogicalTopN<>(ImmutableList.of(new OrderKey(distance, true, false)), + vectorSearch.getTopK(), vectorSearch.getOffset(), relation); + } + FullTextSearchTableValuedFunction fullTextSearch = + (FullTextSearchTableValuedFunction) tableValuedFunction.getCatalogFunction(); + Slot score = requireSearchResultSlot( + relation, FullTextSearchTableValuedFunction.SCORE_COLUMN, + FullTextSearchTableValuedFunction.NAME); + return new LogicalTopN<>(ImmutableList.of(new OrderKey(score, false, false)), + fullTextSearch.getTopK(), fullTextSearch.getOffset(), relation); + } + + private Slot requireSearchResultSlot( + LogicalTVFRelation relation, String column, String functionName) { + return relation.getOutput().stream() + .filter(slot -> slot.getName().equalsIgnoreCase(column)) .findFirst() - .orElseThrow(() -> new AnalysisException("vector_search() output is missing '" - + VectorSearchTableValuedFunction.DISTANCE_COLUMN + "'")); - OrderKey distanceAscending = new OrderKey(distance, true, false); - return new LogicalTopN<>(ImmutableList.of(distanceAscending), - vectorSearch.getTopK(), vectorSearch.getOffset(), relation); + .orElseThrow(() -> new AnalysisException(functionName + "() output is missing '" + + column + "'")); } private void checkIfOutputAliasNameDuplicatedForGroupBy(Collection expressions, diff --git a/fe/fe-core/src/main/java/org/apache/doris/nereids/rules/rewrite/PushDownFilterThroughVectorSearchTopN.java b/fe/fe-core/src/main/java/org/apache/doris/nereids/rules/rewrite/PushDownFilterThroughVectorSearchTopN.java index f6c7feed6dbb7a..73f65e5a5ed8be 100644 --- a/fe/fe-core/src/main/java/org/apache/doris/nereids/rules/rewrite/PushDownFilterThroughVectorSearchTopN.java +++ b/fe/fe-core/src/main/java/org/apache/doris/nereids/rules/rewrite/PushDownFilterThroughVectorSearchTopN.java @@ -23,24 +23,25 @@ import org.apache.doris.nereids.trees.plans.logical.LogicalFilter; import org.apache.doris.nereids.trees.plans.logical.LogicalTVFRelation; import org.apache.doris.nereids.trees.plans.logical.LogicalTopN; +import org.apache.doris.tablefunction.FullTextSearchTableValuedFunction; import org.apache.doris.tablefunction.VectorSearchTableValuedFunction; /** - * Move an outer vector_search WHERE predicate below its Doris merge TopN. + * Move an outer Lance external-search WHERE predicate below its Doris merge TopN. * - *

The TopN immediately above a vector_search TVF is added by {@code BindExpression} to merge + *

The TopN immediately above a search TVF is added by {@code BindExpression} to merge * the candidates returned by all Lance fragment scans. The SQL WHERE predicate must therefore be * evaluated below this TopN so it can become a residual conjunct on the Doris Lance scan node: * *

  * Filter                         TopN
  *   TopN            ->            Filter
- *     vector_search                 vector_search
+ *     search TVF                    search TVF
  * 
* - *

This remains a postfilter relative to Lance nearest(): every fragment first returns its ANN - * candidates, and Doris filters those candidates before the local/global TopN. It is deliberately - * not converted into the Lance prefilter carried by the TVF's {@code filter} property. + *

This remains a postfilter relative to the Lance search: every split first returns candidates, + * and Doris filters those candidates before the local/global TopN. It is deliberately not + * converted into the Lance prefilter carried by the TVF's {@code filter} property. */ public class PushDownFilterThroughVectorSearchTopN extends OneRewriteRuleFactory { @Override @@ -48,8 +49,9 @@ public Rule build() { return logicalFilter(logicalTopN(logicalTVFRelation())) .then(filter -> { LogicalTopN topN = filter.child(); - if (!VectorSearchTableValuedFunction.NAME.equals( - topN.child().getFunction().getName())) { + String functionName = topN.child().getFunction().getName(); + if (!VectorSearchTableValuedFunction.NAME.equals(functionName) + && !FullTextSearchTableValuedFunction.NAME.equals(functionName)) { return null; } LogicalFilter scanFilter = new LogicalFilter<>( diff --git a/fe/fe-core/src/main/java/org/apache/doris/nereids/trees/expressions/functions/table/FullTextSearch.java b/fe/fe-core/src/main/java/org/apache/doris/nereids/trees/expressions/functions/table/FullTextSearch.java new file mode 100644 index 00000000000000..55b297ae6688c7 --- /dev/null +++ b/fe/fe-core/src/main/java/org/apache/doris/nereids/trees/expressions/functions/table/FullTextSearch.java @@ -0,0 +1,49 @@ +// 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. + +package org.apache.doris.nereids.trees.expressions.functions.table; + +import org.apache.doris.catalog.FunctionSignature; +import org.apache.doris.nereids.exceptions.AnalysisException; +import org.apache.doris.nereids.trees.expressions.Properties; +import org.apache.doris.nereids.types.coercion.AnyDataType; +import org.apache.doris.tablefunction.FullTextSearchTableValuedFunction; +import org.apache.doris.tablefunction.TableValuedFunctionIf; + +import java.util.Map; + +/** Lance full_text_search relation TVF. */ +public class FullTextSearch extends TableValuedFunction { + public FullTextSearch(Properties properties) { + super(FullTextSearchTableValuedFunction.NAME, properties); + } + + @Override + public FunctionSignature customSignature() { + return FunctionSignature.of(AnyDataType.INSTANCE_WITHOUT_INDEX, getArgumentsTypes()); + } + + @Override + protected TableValuedFunctionIf toCatalogFunction() { + try { + Map arguments = getTVFProperties().getMap(); + return new FullTextSearchTableValuedFunction(arguments); + } catch (Throwable t) { + throw new AnalysisException("Can not build full_text_search(): " + t.getMessage(), t); + } + } +} diff --git a/fe/fe-core/src/main/java/org/apache/doris/tablefunction/FullTextSearchTableValuedFunction.java b/fe/fe-core/src/main/java/org/apache/doris/tablefunction/FullTextSearchTableValuedFunction.java new file mode 100644 index 00000000000000..735b65d366fdf5 --- /dev/null +++ b/fe/fe-core/src/main/java/org/apache/doris/tablefunction/FullTextSearchTableValuedFunction.java @@ -0,0 +1,113 @@ +// 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. + +package org.apache.doris.tablefunction; + +import org.apache.doris.common.AnalysisException; +import org.apache.doris.datasource.lance.LanceTableMetadata; +import org.apache.doris.thrift.TExternalSearchQuery; +import org.apache.doris.thrift.TExternalSearchRequest; +import org.apache.doris.thrift.TFtsCoverageMode; +import org.apache.doris.thrift.TFullTextSearchParams; + +import com.google.common.collect.ImmutableSet; +import org.apache.arrow.vector.types.pojo.ArrowType; +import org.apache.arrow.vector.types.pojo.Field; + +import java.util.Locale; +import java.util.Map; +import java.util.Set; + +/** Relation TVF for a fixed-snapshot distributed Lance full-text search. */ +public class FullTextSearchTableValuedFunction extends LanceExternalSearchTableValuedFunction { + public static final String NAME = "full_text_search"; + public static final String SCORE_COLUMN = "_score"; + + private static final String QUERY = "query"; + private static final String COVERAGE_MODE = "coverage_mode"; + private static final Set PROPERTIES = ImmutableSet.of( + TABLE, COLUMN, QUERY, TOP_K, OFFSET, FILTER, COVERAGE_MODE); + + public FullTextSearchTableValuedFunction(Map properties) + throws AnalysisException { + super(prepare(properties)); + } + + private static PreparedSearch prepare(Map properties) + throws AnalysisException { + Map params = normalizeProperties(properties, PROPERTIES, NAME); + CommonSearch common = prepareCommon(params, NAME, + "FullTextSearchTableValuedFunction", "full-text search", true); + + Field field = findStringField(common.metadata(), required(params, COLUMN, NAME)); + int fieldId = requireLanceFieldId(common.metadata(), field, "full-text"); + String query = required(params, QUERY, NAME); + if (query.indexOf('\0') >= 0) { + throw new AnalysisException("'query' must not contain an embedded NUL byte"); + } + + TFullTextSearchParams fullTextParams = new TFullTextSearchParams() + .setColumn(field.getName()) + .setQuery(query) + .setTopK(common.topK()) + .setOffset(common.offset()) + .setCoverageMode(parseCoverageMode( + params.getOrDefault(COVERAGE_MODE, "strict"))); + TExternalSearchRequest searchRequest = new TExternalSearchRequest() + .setSchemaVersion(1) + .setSearchQuery(TExternalSearchQuery.full_text_search(fullTextParams)); + return prepareSearch( + common, fieldId, searchRequest, SCORE_COLUMN, "full-text search"); + } + + private static Field findStringField(LanceTableMetadata metadata, String column) + throws AnalysisException { + Field match = null; + for (Field field : metadata.getSchema().getFields()) { + if (field.getName().equalsIgnoreCase(column)) { + if (match != null) { + throw new AnalysisException("Lance full-text column '" + column + + "' is ambiguous under case-insensitive matching"); + } + match = field; + } + } + if (match == null) { + throw new AnalysisException("Lance full-text column '" + column + "' does not exist"); + } + ArrowType.ArrowTypeID typeId = match.getType().getTypeID(); + if (typeId != ArrowType.ArrowTypeID.Utf8 + && typeId != ArrowType.ArrowTypeID.LargeUtf8) { + throw new AnalysisException("Lance full-text column '" + match.getName() + + "' must be STRING"); + } + return match; + } + + private static TFtsCoverageMode parseCoverageMode(String value) throws AnalysisException { + switch (value.trim().toLowerCase(Locale.ROOT)) { + case "strict": + return TFtsCoverageMode.STRICT; + case "index_only": + case "index-only": + return TFtsCoverageMode.INDEX_ONLY; + default: + throw new AnalysisException("Unsupported FTS coverage_mode '" + value + + "': expected strict or index_only"); + } + } +} diff --git a/fe/fe-core/src/main/java/org/apache/doris/tablefunction/LanceExternalSearchTableValuedFunction.java b/fe/fe-core/src/main/java/org/apache/doris/tablefunction/LanceExternalSearchTableValuedFunction.java new file mode 100644 index 00000000000000..856ee108fc6c9a --- /dev/null +++ b/fe/fe-core/src/main/java/org/apache/doris/tablefunction/LanceExternalSearchTableValuedFunction.java @@ -0,0 +1,360 @@ +// 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. + +package org.apache.doris.tablefunction; + +import org.apache.doris.analysis.TableName; +import org.apache.doris.analysis.TupleDescriptor; +import org.apache.doris.catalog.Column; +import org.apache.doris.catalog.Env; +import org.apache.doris.catalog.TableIf; +import org.apache.doris.catalog.Type; +import org.apache.doris.common.AnalysisException; +import org.apache.doris.common.ErrorCode; +import org.apache.doris.common.ErrorReport; +import org.apache.doris.datasource.CatalogIf; +import org.apache.doris.datasource.lance.LanceExternalCatalog; +import org.apache.doris.datasource.lance.LanceExternalTable; +import org.apache.doris.datasource.lance.LanceTableMetadata; +import org.apache.doris.datasource.lance.LanceTypeConverter; +import org.apache.doris.datasource.lance.source.LanceScanNode; +import org.apache.doris.mysql.privilege.PrivPredicate; +import org.apache.doris.nereids.analyzer.UnboundSlot; +import org.apache.doris.nereids.exceptions.ParseException; +import org.apache.doris.nereids.parser.NereidsParser; +import org.apache.doris.nereids.trees.expressions.Expression; +import org.apache.doris.planner.PlanNodeId; +import org.apache.doris.planner.ScanNode; +import org.apache.doris.qe.ConnectContext; +import org.apache.doris.qe.SessionVariable; +import org.apache.doris.thrift.TExternalSearchRequest; +import org.apache.doris.thrift.TSearchFilter; +import org.apache.doris.thrift.TSearchFilterFormat; + +import org.apache.arrow.vector.types.pojo.Field; + +import java.nio.charset.StandardCharsets; +import java.util.ArrayList; +import java.util.Collections; +import java.util.List; +import java.util.Locale; +import java.util.Map; +import java.util.OptionalInt; +import java.util.Set; +import java.util.TreeMap; +import java.util.TreeSet; + +/** Common immutable planning state and validation for Lance external-search relation TVFs. */ +abstract class LanceExternalSearchTableValuedFunction extends TableValuedFunctionIf { + protected static final String TABLE = "table"; + protected static final String COLUMN = "column"; + protected static final String TOP_K = "top_k"; + protected static final String OFFSET = "offset"; + protected static final String FILTER = "filter"; + + private static final String FULLY_QUALIFIED_TABLE_NAME_ERROR = + "'table' must be a fully qualified catalog.database.table name"; + private static final long UINT32_MAX = 0xFFFF_FFFFL; + + private final String displayName; + private final TableName sourceTableName; + private final LanceExternalTable sourceTable; + private final LanceTableMetadata metadata; + private final int fieldId; + private final TExternalSearchRequest searchRequest; + private final List columns; + private final long topK; + private final long offset; + + protected LanceExternalSearchTableValuedFunction(PreparedSearch prepared) { + CommonSearch common = prepared.common; + this.displayName = common.displayName; + this.sourceTableName = common.sourceTableName; + this.sourceTable = common.sourceTable; + this.metadata = common.metadata; + this.fieldId = prepared.fieldId; + this.searchRequest = prepared.searchRequest.deepCopy(); + this.columns = Collections.unmodifiableList(new ArrayList<>(prepared.columns)); + this.topK = common.topK; + this.offset = common.offset; + } + + public final LanceExternalTable getSourceTable() { + return sourceTable; + } + + public final LanceTableMetadata getMetadata() { + return metadata; + } + + public final TExternalSearchRequest getSearchRequest() { + return searchRequest.deepCopy(); + } + + public final long getTopK() { + return topK; + } + + public final long getOffset() { + return offset; + } + + @Override + public final String getTableName() { + return displayName + "<" + sourceTableName + ">"; + } + + @Override + public final List getTableColumns() { + return columns; + } + + @Override + public final ScanNode getScanNode(PlanNodeId id, TupleDescriptor desc, SessionVariable sv) { + return LanceScanNode.forExternalSearch( + id, desc, sourceTable, metadata, fieldId, searchRequest, sv); + } + + protected static Map normalizeProperties(Map properties, + Set allowedProperties, String functionName) throws AnalysisException { + Map normalized = new TreeMap<>(String.CASE_INSENSITIVE_ORDER); + for (Map.Entry entry : properties.entrySet()) { + String key = entry.getKey().toLowerCase(Locale.ROOT); + if (!allowedProperties.contains(key)) { + throw new AnalysisException("'" + entry.getKey() + + "' is an invalid property for " + functionName + "()"); + } + if (normalized.put(key, entry.getValue()) != null) { + throw new AnalysisException( + "Duplicate " + functionName + "() property '" + key + "'"); + } + } + return normalized; + } + + protected static String required(Map params, String key, String functionName) + throws AnalysisException { + String value = params.get(key); + if (value == null || value.trim().isEmpty()) { + throw new AnalysisException( + "Missing required " + functionName + "() property '" + key + "'"); + } + return value.trim(); + } + + protected static CommonSearch prepareCommon(Map params, String functionName, + String displayName, String searchDescription, boolean loadIndexMetadata) + throws AnalysisException { + TableName sourceTableName = parseTableName(required(params, TABLE, functionName)); + LanceExternalTable sourceTable = findLanceExternalTable(sourceTableName); + LanceTableMetadata metadata; + try { + metadata = loadIndexMetadata + ? sourceTable.loadMetadataForSearch() : sourceTable.loadMetadata(); + } catch (RuntimeException e) { + throw new AnalysisException("Failed to load Lance metadata for " + searchDescription + + " on " + sourceTableName + ": " + e.getMessage(), e); + } + if (metadata.getVersion() <= 0) { + throw new AnalysisException("Lance " + searchDescription + + " requires a fixed positive dataset version"); + } + + long topK = parseLong(params.getOrDefault(TOP_K, "10"), TOP_K, 1, Long.MAX_VALUE); + long offset = parseLong(params.getOrDefault(OFFSET, "0"), OFFSET, 0, Long.MAX_VALUE); + if (offset > UINT32_MAX || topK > UINT32_MAX - offset) { + throw new AnalysisException("'top_k + offset' must not exceed " + UINT32_MAX); + } + return new CommonSearch(params, displayName, sourceTableName, sourceTable, metadata, + topK, offset); + } + + protected static PreparedSearch prepareSearch(CommonSearch common, int fieldId, + TExternalSearchRequest searchRequest, String resultColumn, String searchDescription) + throws AnalysisException { + if (common.params.containsKey(FILTER)) { + searchRequest.setSearchFilter(new TSearchFilter() + .setFormat(TSearchFilterFormat.SQL) + .setPayload(validateAndEncodeSqlFilter(common.params.get(FILTER)))); + } + List columns = buildOutputColumns( + common.metadata, resultColumn, searchDescription); + return new PreparedSearch(common, fieldId, searchRequest, columns); + } + + protected static int requireLanceFieldId(LanceTableMetadata metadata, Field field, + String searchDescription) throws AnalysisException { + OptionalInt fieldId = metadata.getLanceFieldId(field.getName()); + if (!fieldId.isPresent()) { + throw new AnalysisException("Lance " + searchDescription + " column '" + + field.getName() + "' has no field ID in the Lance schema"); + } + return fieldId.getAsInt(); + } + + protected static List buildOutputColumns(LanceTableMetadata metadata, + String resultColumn, String searchDescription) throws AnalysisException { + List result = new ArrayList<>(metadata.getSchema().getFields().size() + 1); + Set fieldNames = new TreeSet<>(String.CASE_INSENSITIVE_ORDER); + int position = 0; + for (Field field : metadata.getSchema().getFields()) { + if (!fieldNames.add(field.getName())) { + throw new AnalysisException("Duplicate Lance schema column under " + + "case-insensitive matching: '" + field.getName() + "'"); + } + if (field.getName().startsWith(Column.GLOBAL_ROWID_COL)) { + throw new AnalysisException("Lance table contains column '" + field.getName() + + "' using reserved Doris internal column prefix '" + + Column.GLOBAL_ROWID_COL + "'"); + } + if (field.getName().equalsIgnoreCase(resultColumn)) { + throw new AnalysisException("Lance table already contains reserved " + + searchDescription + " column '" + resultColumn + "'"); + } + String comment = field.getMetadata() == null + ? null : field.getMetadata().get("comment"); + Type type; + try { + type = LanceTypeConverter.toDorisType(field); + } catch (RuntimeException e) { + throw new AnalysisException("Invalid Lance type for column '" + field.getName() + + "': " + e.getMessage(), e); + } + result.add(new Column(field.getName(), type, false, null, + field.isNullable(), comment, true, position++)); + } + result.add(new Column(resultColumn, Type.FLOAT, false, null, + true, null, true, position)); + return result; + } + + protected static TableName parseTableName(String value) throws AnalysisException { + Expression expression; + try { + expression = new NereidsParser().parseExpression(value); + } catch (ParseException e) { + throw new AnalysisException(FULLY_QUALIFIED_TABLE_NAME_ERROR, e); + } + if (!(expression instanceof UnboundSlot)) { + throw new AnalysisException(FULLY_QUALIFIED_TABLE_NAME_ERROR); + } + List names = ((UnboundSlot) expression).getNameParts(); + if (names.size() != 3) { + throw new AnalysisException(FULLY_QUALIFIED_TABLE_NAME_ERROR); + } + return new TableName(names.get(0), names.get(1), names.get(2)); + } + + protected static LanceExternalTable findLanceExternalTable(TableName tableName) + throws AnalysisException { + ConnectContext context = ConnectContext.get(); + if (!Env.getCurrentEnv().getAccessManager() + .checkTblPriv(context, tableName, PrivPredicate.SELECT)) { + ErrorReport.reportAnalysisException(ErrorCode.ERR_TABLEACCESS_DENIED_ERROR, "SELECT", + context.getQualifiedUser(), context.getRemoteIP(), + tableName.getDb() + ": " + tableName.getTbl()); + } + CatalogIf catalog = Env.getCurrentEnv().getCatalogMgr().getCatalog(tableName.getCtl()); + if (!(catalog instanceof LanceExternalCatalog)) { + throw new AnalysisException("Catalog '" + tableName.getCtl() + + "' is not a Lance catalog"); + } + TableIf table = catalog.getDbOrAnalysisException(tableName.getDb()) + .getTableOrAnalysisException(tableName.getTbl()); + if (!(table instanceof LanceExternalTable)) { + throw new AnalysisException("Table '" + tableName + "' is not a Lance table"); + } + return (LanceExternalTable) table; + } + + protected static byte[] validateAndEncodeSqlFilter(String filter) throws AnalysisException { + if (filter == null || filter.trim().isEmpty()) { + throw new AnalysisException("'filter' must not be empty"); + } + if (filter.indexOf('\0') >= 0) { + throw new AnalysisException("'filter' must not contain an embedded NUL byte"); + } + return filter.getBytes(StandardCharsets.UTF_8); + } + + protected static long parseLong(String value, String property, long min, long max) + throws AnalysisException { + try { + long parsed = Long.parseLong(value); + if (parsed < min || parsed > max) { + throw new AnalysisException("'" + property + "' must be between " + + min + " and " + max); + } + return parsed; + } catch (NumberFormatException e) { + throw new AnalysisException("'" + property + "' must be an integer", e); + } + } + + protected static final class CommonSearch { + private final Map params; + private final String displayName; + private final TableName sourceTableName; + private final LanceExternalTable sourceTable; + private final LanceTableMetadata metadata; + private final long topK; + private final long offset; + + private CommonSearch(Map params, String displayName, + TableName sourceTableName, LanceExternalTable sourceTable, + LanceTableMetadata metadata, long topK, long offset) { + this.params = Collections.unmodifiableMap(new TreeMap<>(params)); + this.displayName = displayName; + this.sourceTableName = sourceTableName; + this.sourceTable = sourceTable; + this.metadata = metadata; + this.topK = topK; + this.offset = offset; + } + + protected Map params() { + return params; + } + + protected LanceTableMetadata metadata() { + return metadata; + } + + protected long topK() { + return topK; + } + + protected long offset() { + return offset; + } + } + + protected static final class PreparedSearch { + private final CommonSearch common; + private final int fieldId; + private final TExternalSearchRequest searchRequest; + private final List columns; + + private PreparedSearch(CommonSearch common, int fieldId, + TExternalSearchRequest searchRequest, List columns) { + this.common = common; + this.fieldId = fieldId; + this.searchRequest = searchRequest; + this.columns = columns; + } + } +} diff --git a/fe/fe-core/src/main/java/org/apache/doris/tablefunction/TableValuedFunctionIf.java b/fe/fe-core/src/main/java/org/apache/doris/tablefunction/TableValuedFunctionIf.java index 7a7569583d5283..79c10b465bee27 100644 --- a/fe/fe-core/src/main/java/org/apache/doris/tablefunction/TableValuedFunctionIf.java +++ b/fe/fe-core/src/main/java/org/apache/doris/tablefunction/TableValuedFunctionIf.java @@ -104,6 +104,8 @@ public static TableValuedFunctionIf getTableFunction(String funcName, Map PROPERTIES = ImmutableSet.of( TABLE, COLUMN, QUERY_VECTOR, TOP_K, OFFSET, METRIC, FILTER, NPROBES, REFINE_FACTOR, EF, USE_INDEX); - private final TableName sourceTableName; - private final LanceExternalTable sourceTable; - private final LanceTableMetadata metadata; - private final int vectorFieldId; - private final List columns; - private final TExternalSearchRequest searchRequest; - public VectorSearchTableValuedFunction(Map properties) throws AnalysisException { - Map params = normalizeProperties(properties); - sourceTableName = parseTableName(required(params, TABLE)); - sourceTable = findLanceExternalTable(sourceTableName); + super(prepare(properties)); + } + + private static PreparedSearch prepare(Map properties) + throws AnalysisException { + Map params = normalizeProperties(properties, PROPERTIES, NAME); boolean useIndex = !params.containsKey(USE_INDEX) || parseBoolean(params.get(USE_INDEX), USE_INDEX); - try { - metadata = useIndex - ? sourceTable.loadMetadataForVectorSearch() : sourceTable.loadMetadata(); - } catch (RuntimeException e) { - throw new AnalysisException("Failed to load Lance metadata for vector search on " - + sourceTableName + ": " + e.getMessage(), e); - } - if (metadata.getVersion() <= 0) { - throw new AnalysisException("Lance vector search requires a fixed positive dataset version"); - } + CommonSearch common = prepareCommon(params, NAME, + "VectorSearchTableValuedFunction", "vector search", useIndex); Field vectorField = LanceVectorQuery.findVectorColumnField( - metadata.getSchema(), required(params, COLUMN)); - vectorFieldId = useIndex ? requireLanceFieldId(metadata, vectorField) : -1; + common.metadata().getSchema(), required(params, COLUMN, NAME)); + int vectorFieldId = useIndex + ? requireLanceFieldId(common.metadata(), vectorField) : -1; TSearchVector queryVector = LanceVectorQuery.parseAndEncodeQueryVector( - vectorField, required(params, QUERY_VECTOR)); - long topK = parseLong(params.getOrDefault(TOP_K, "10"), TOP_K, 1, Long.MAX_VALUE); - long offset = parseLong(params.getOrDefault(OFFSET, "0"), OFFSET, 0, Long.MAX_VALUE); - if (offset > UINT32_MAX || topK > UINT32_MAX - offset) { - throw new AnalysisException("'top_k + offset' must not exceed " + UINT32_MAX); - } + vectorField, required(params, QUERY_VECTOR, NAME)); TVectorSearchParams vectorParams = new TVectorSearchParams() .setColumn(vectorField.getName()) .setQueryVector(queryVector) - .setTopK(topK) - .setOffset(offset); + .setTopK(common.topK()) + .setOffset(common.offset()); if (params.containsKey(METRIC)) { vectorParams.setMetric(parseMetric(params.get(METRIC))); } - searchRequest = new TExternalSearchRequest() + TExternalSearchRequest searchRequest = new TExternalSearchRequest() .setSchemaVersion(1) .setSearchQuery(TExternalSearchQuery.vector_search(vectorParams)); - if (params.containsKey(FILTER)) { - searchRequest.setSearchFilter(new TSearchFilter() - .setFormat(TSearchFilterFormat.SQL) - .setPayload(validateAndEncodeSqlFilter(params.get(FILTER)))); + TVectorSearchOptions vectorSearchOptions = buildVectorSearchOptions(params, useIndex); + if (vectorSearchOptions != null) { + searchRequest.setVectorSearchOptions(vectorSearchOptions); } + return prepareSearch( + common, vectorFieldId, searchRequest, DISTANCE_COLUMN, "vector search"); + } - TVectorSearchOptions vectorSearchOptions = new TVectorSearchOptions(); - boolean hasVectorSearchOptions = false; + private static TVectorSearchOptions buildVectorSearchOptions( + Map params, boolean useIndex) throws AnalysisException { + TVectorSearchOptions options = new TVectorSearchOptions(); + boolean configured = false; if (params.containsKey(NPROBES)) { - vectorSearchOptions.setNprobes(parsePositiveInt(params.get(NPROBES), NPROBES)); - hasVectorSearchOptions = true; + options.setNprobes(parsePositiveInt(params.get(NPROBES), NPROBES)); + configured = true; } if (params.containsKey(REFINE_FACTOR)) { - vectorSearchOptions.setRefineFactor( + options.setRefineFactor( parsePositiveInt(params.get(REFINE_FACTOR), REFINE_FACTOR)); - hasVectorSearchOptions = true; + configured = true; } if (params.containsKey(EF)) { - vectorSearchOptions.setEf(parsePositiveInt(params.get(EF), EF)); - hasVectorSearchOptions = true; + options.setEf(parsePositiveInt(params.get(EF), EF)); + configured = true; } if (params.containsKey(USE_INDEX)) { - vectorSearchOptions.setUseIndex(useIndex); - hasVectorSearchOptions = true; - } - if (hasVectorSearchOptions) { - searchRequest.setVectorSearchOptions(vectorSearchOptions); - } - columns = buildOutputColumns(metadata); - } - - public LanceExternalTable getSourceTable() { - return sourceTable; - } - - public LanceTableMetadata getMetadata() { - return metadata; - } - - public TExternalSearchRequest getSearchRequest() { - return searchRequest.deepCopy(); - } - - public long getTopK() { - return searchRequest.getSearchQuery().getVectorSearch().getTopK(); - } - - public long getOffset() { - return searchRequest.getSearchQuery().getVectorSearch().getOffset(); - } - - @Override - public String getTableName() { - return "VectorSearchTableValuedFunction<" + sourceTableName + ">"; - } - - @Override - public List getTableColumns() { - return columns; - } - - @Override - public ScanNode getScanNode(PlanNodeId id, TupleDescriptor desc, SessionVariable sv) { - return LanceScanNode.forVectorSearch(id, desc, sourceTable, metadata, - vectorFieldId, searchRequest, sv); - } - - private static Map normalizeProperties(Map properties) - throws AnalysisException { - Map normalized = new TreeMap<>(String.CASE_INSENSITIVE_ORDER); - for (Map.Entry entry : properties.entrySet()) { - String key = entry.getKey().toLowerCase(Locale.ROOT); - if (!PROPERTIES.contains(key)) { - throw new AnalysisException("'" + entry.getKey() - + "' is an invalid property for vector_search()"); - } - if (normalized.put(key, entry.getValue()) != null) { - throw new AnalysisException("Duplicate vector_search() property '" + key + "'"); - } - } - return normalized; - } - - private static String required(Map params, String key) - throws AnalysisException { - String value = params.get(key); - if (value == null || value.trim().isEmpty()) { - throw new AnalysisException("Missing required vector_search() property '" + key + "'"); + options.setUseIndex(useIndex); + configured = true; } - return value.trim(); - } - - @VisibleForTesting - static TableName parseTableName(String value) throws AnalysisException { - Expression expression; - try { - expression = new NereidsParser().parseExpression(value); - } catch (ParseException e) { - throw new AnalysisException(FULLY_QUALIFIED_TABLE_NAME_ERROR, e); - } - if (!(expression instanceof UnboundSlot)) { - throw new AnalysisException(FULLY_QUALIFIED_TABLE_NAME_ERROR); - } - List names = ((UnboundSlot) expression).getNameParts(); - if (names.size() != 3) { - throw new AnalysisException(FULLY_QUALIFIED_TABLE_NAME_ERROR); - } - return new TableName(names.get(0), names.get(1), names.get(2)); - } - - private static LanceExternalTable findLanceExternalTable(TableName tableName) - throws AnalysisException { - ConnectContext context = ConnectContext.get(); - if (!Env.getCurrentEnv().getAccessManager() - .checkTblPriv(context, tableName, PrivPredicate.SELECT)) { - ErrorReport.reportAnalysisException(ErrorCode.ERR_TABLEACCESS_DENIED_ERROR, "SELECT", - context.getQualifiedUser(), context.getRemoteIP(), - tableName.getDb() + ": " + tableName.getTbl()); - } - CatalogIf catalog = Env.getCurrentEnv().getCatalogMgr().getCatalog(tableName.getCtl()); - if (!(catalog instanceof LanceExternalCatalog)) { - throw new AnalysisException("Catalog '" + tableName.getCtl() - + "' is not a Lance catalog"); - } - TableIf table = catalog.getDbOrAnalysisException(tableName.getDb()) - .getTableOrAnalysisException(tableName.getTbl()); - if (!(table instanceof LanceExternalTable)) { - throw new AnalysisException("Table '" + tableName + "' is not a Lance table"); - } - return (LanceExternalTable) table; + return configured ? options : null; } @VisibleForTesting static List buildOutputColumns(LanceTableMetadata metadata) throws AnalysisException { - List result = new ArrayList<>(metadata.getSchema().getFields().size() + 1); - Set fieldNames = new TreeSet<>(String.CASE_INSENSITIVE_ORDER); - int position = 0; - for (Field field : metadata.getSchema().getFields()) { - if (!fieldNames.add(field.getName())) { - throw new AnalysisException("Duplicate Lance schema column under " - + "case-insensitive matching: '" + field.getName() + "'"); - } - if (field.getName().startsWith(Column.GLOBAL_ROWID_COL)) { - throw new AnalysisException("Lance table contains column '" + field.getName() - + "' using reserved Doris internal column prefix '" - + Column.GLOBAL_ROWID_COL + "'"); - } - if (field.getName().equalsIgnoreCase(DISTANCE_COLUMN)) { - throw new AnalysisException("Lance table already contains reserved vector search " - + "column '" + DISTANCE_COLUMN + "'"); - } - String comment = field.getMetadata() == null - ? null : field.getMetadata().get("comment"); - Type type; - try { - type = LanceTypeConverter.toDorisType(field); - } catch (RuntimeException e) { - throw new AnalysisException("Invalid Lance type for column '" + field.getName() - + "': " + e.getMessage(), e); - } - result.add(new Column(field.getName(), type, false, null, - field.isNullable(), comment, true, position++)); - } - result.add(new Column(DISTANCE_COLUMN, Type.FLOAT, false, null, - true, null, true, position)); - return result; + return buildOutputColumns(metadata, DISTANCE_COLUMN, "vector search"); } @VisibleForTesting static int requireLanceFieldId(LanceTableMetadata metadata, Field field) throws AnalysisException { - OptionalInt fieldId = metadata.getLanceFieldId(field.getName()); - if (!fieldId.isPresent()) { - throw new AnalysisException("Lance vector column '" + field.getName() - + "' has no field ID in the Lance schema"); - } - return fieldId.getAsInt(); - } - - @VisibleForTesting - static byte[] validateAndEncodeSqlFilter(String filter) throws AnalysisException { - if (filter == null || filter.trim().isEmpty()) { - throw new AnalysisException("'filter' must not be empty"); - } - if (filter.indexOf('\0') >= 0) { - throw new AnalysisException("'filter' must not contain an embedded NUL byte"); - } - return filter.getBytes(StandardCharsets.UTF_8); - } - - private static long parseLong(String value, String property, long min, long max) - throws AnalysisException { - try { - long parsed = Long.parseLong(value); - if (parsed < min || parsed > max) { - throw new AnalysisException("'" + property + "' must be between " - + min + " and " + max); - } - return parsed; - } catch (NumberFormatException e) { - throw new AnalysisException("'" + property + "' must be an integer", e); - } + return requireLanceFieldId(metadata, field, "vector"); } private static int parsePositiveInt(String value, String property) diff --git a/fe/fe-core/src/test/java/org/apache/doris/datasource/lance/source/LanceScanNodeTest.java b/fe/fe-core/src/test/java/org/apache/doris/datasource/lance/source/LanceScanNodeTest.java index 2fb3b7c390df6f..e38f0120270748 100644 --- a/fe/fe-core/src/test/java/org/apache/doris/datasource/lance/source/LanceScanNodeTest.java +++ b/fe/fe-core/src/test/java/org/apache/doris/datasource/lance/source/LanceScanNodeTest.java @@ -39,6 +39,7 @@ import org.apache.arrow.vector.types.pojo.Schema; import org.junit.Assert; import org.junit.Test; +import org.lance.index.IndexType; import java.nio.ByteBuffer; import java.util.Arrays; @@ -203,9 +204,11 @@ public void testExternalSearchUsesOneSplitPerIndexSegmentAndKeepsUnindexedFragme Collections.singletonMap("vector", 9), Arrays.asList( new LanceIndexSegmentInfo(firstSegment, "vector_idx", - Collections.singletonList(9), Arrays.asList(1L, 2L), "L2"), + Collections.singletonList(9), Arrays.asList(1L, 2L), + IndexType.VECTOR, "L2"), new LanceIndexSegmentInfo(secondSegment, "vector_idx", - Collections.singletonList(9), Arrays.asList(3L, 4L), "L2")), + Collections.singletonList(9), Arrays.asList(3L, 4L), + IndexType.VECTOR, "L2")), Collections.emptyMap()); LanceScanNode node = newSearchNode(metadata, vectorSearchRequest(5, 0)); @@ -238,7 +241,8 @@ public void testExternalSearchUseIndexFalseKeepsFragmentSplits() throws Exceptio Collections.emptyMap(), Collections.singletonList( new LanceIndexSegmentInfo(UUID.randomUUID(), "vector_idx", - Collections.singletonList(9), Arrays.asList(1L, 2L), "L2")), + Collections.singletonList(9), Arrays.asList(1L, 2L), + IndexType.VECTOR, "L2")), Collections.emptyMap()); TExternalSearchRequest request = vectorSearchRequest(5, 0); request.setVectorSearchOptions(new TVectorSearchOptions().setUseIndex(false)); @@ -263,7 +267,8 @@ public void testExternalSearchFallsBackToFragmentSplitsForMetricMismatch() throw Collections.singletonMap("vector", 9), Collections.singletonList( new LanceIndexSegmentInfo(UUID.randomUUID(), "vector_idx", - Collections.singletonList(9), Arrays.asList(1L, 2L), "L2")), + Collections.singletonList(9), Arrays.asList(1L, 2L), + IndexType.VECTOR, "L2")), Collections.emptyMap()); TExternalSearchRequest request = vectorSearchRequest(5, 0); request.getSearchQuery().getVectorSearch().setMetric(TVectorMetric.COSINE); @@ -286,7 +291,8 @@ public void testExternalSearchRejectsMissingFieldIdForIndexSegmentPlanning() { Collections.emptyMap(), Collections.singletonList( new LanceIndexSegmentInfo(UUID.randomUUID(), "vector_idx", - Collections.singletonList(9), Collections.singletonList(1L), "L2")), + Collections.singletonList(9), Collections.singletonList(1L), + IndexType.VECTOR, "L2")), Collections.emptyMap()); LanceScanNode node = newSearchNode(metadata, vectorSearchRequest(5, 0)); @@ -297,14 +303,16 @@ public void testExternalSearchRejectsMissingFieldIdForIndexSegmentPlanning() { } @Test - public void testFragmentSearchRetainsTopKPlusOffsetCandidates() { + public void testSplitSearchRetainsTopKPlusOffsetCandidates() { TExternalSearchRequest logicalRequest = vectorSearchRequest(5, 2); + LanceScanNode node = LanceScanNode.forExternalSearch( + new PlanNodeId(0), new TupleDescriptor(new TupleId(0)), null, + null, -1, logicalRequest, new SessionVariable()); - TExternalSearchRequest fragmentRequest = - LanceScanNode.createFragmentSearchRequest(logicalRequest); + TExternalSearchRequest splitRequest = node.createSplitSearchRequest(); - Assert.assertEquals(7, fragmentRequest.getSearchQuery().getVectorSearch().getTopK()); - Assert.assertEquals(0, fragmentRequest.getSearchQuery().getVectorSearch().getOffset()); + Assert.assertEquals(7, splitRequest.getSearchQuery().getVectorSearch().getTopK()); + Assert.assertEquals(0, splitRequest.getSearchQuery().getVectorSearch().getOffset()); Assert.assertEquals(5, logicalRequest.getSearchQuery().getVectorSearch().getTopK()); Assert.assertEquals(2, logicalRequest.getSearchQuery().getVectorSearch().getOffset()); } @@ -332,7 +340,7 @@ private static LanceScanNode newSearchNode( LanceTableMetadata metadata, TExternalSearchRequest request) { String vectorColumn = request.getSearchQuery().getVectorSearch().getColumn(); int vectorFieldId = metadata.getLanceFieldId(vectorColumn).orElse(-1); - return LanceScanNode.forVectorSearch( + return LanceScanNode.forExternalSearch( new PlanNodeId(0), new TupleDescriptor(new TupleId(0)), null, metadata, vectorFieldId, request, new SessionVariable()); } diff --git a/gensrc/thrift/PlanNodes.thrift b/gensrc/thrift/PlanNodes.thrift index b543afca315737..c5f31d842bd918 100644 --- a/gensrc/thrift/PlanNodes.thrift +++ b/gensrc/thrift/PlanNodes.thrift @@ -475,6 +475,11 @@ struct TVectorSearchParams { 5: optional TVectorMetric metric } +enum TFtsCoverageMode { + STRICT, + INDEX_ONLY +} + // Logical parameters for one full-text query. `query` initially carries the backend query string; // richer structured query forms can be added as new fields without changing this basic contract. struct TFullTextSearchParams { @@ -482,6 +487,13 @@ struct TFullTextSearchParams { 2: optional string query 3: optional i64 top_k 4: optional i64 offset + // STRICT requires the selected FTS index to cover the complete pinned snapshot. INDEX_ONLY + // searches and scores only fragments covered by committed FTS index segments. + 5: optional TFtsCoverageMode coverage_mode + // Opaque, versioned global BM25 statistics prepared by Lance for this exact snapshot and + // query. Unset while BE scanners prepare statistics locally; future FE versions may populate + // this field once the bundled lance-c exposes the corresponding consumer API. + 6: optional binary global_statistics } enum TSearchFilterFormat { @@ -531,8 +543,9 @@ struct TLanceFileDesc { // most this many rows; the upper LIMIT operator still enforces the global bound. // Only set for ordinary scans whose predicates are fully pushed into Lance. 4: optional i64 limit - // Physical vector-index segments assigned to this distributed search split. Each value is one - // UUID encoded as 16 bytes in RFC 4122 order. Unset for ordinary and unindexed-fragment scans. + // Physical vector or FTS index segments assigned to this distributed search split. Each value + // is one UUID encoded as 16 bytes in RFC 4122 order. Unset for ordinary and vector + // unindexed-fragment scans. 5: optional list index_segment_uuids } @@ -541,8 +554,8 @@ struct TLanceScanParams { // ScanNode level so it is not serialized once per fragment split. 1: optional binary lance_substrait_filter // Provider-independent search request. Set at ScanNode level so all ranges use the same logical - // query. Lance vector search uses one range per fragment and Doris merges the split-local - // candidates. + // query. Lance external search uses one range per fragment or physical index segment, and Doris + // merges the split-local candidates. 2: optional TExternalSearchRequest external_search_request // Lance-native storage options, handed to lance-c untranslated. The namespace protocol treats // storage_options as opaque configuration passed directly to Lance, so any key vocabulary the diff --git a/regression-test/data/external_table_p0/lance/test_lance_runtime_filter_pushdown.out b/regression-test/data/external_table_p0/lance/test_lance_runtime_filter_pushdown.out new file mode 100644 index 00000000000000..d22abdc0d07eda --- /dev/null +++ b/regression-test/data/external_table_p0/lance/test_lance_runtime_filter_pushdown.out @@ -0,0 +1,18 @@ +-- This file is automatically generated. You should know what you did if you want to edit this +-- !runtime_filter_and_substrait -- +7 10 +9 100 + +-- !two_phase_vector_search -- +1 even item-0001 0.0 +2 odd item-0002 16.0 +3 even item-0003 64.0 +4 odd item-0004 144.0 +5 even item-0005 256.0 + +-- !explicit_join_vector_search -- +1 even item-0001 0.0 +2 odd item-0002 16.0 +3 even item-0003 64.0 +4 odd item-0004 144.0 +5 even item-0005 256.0 diff --git a/regression-test/suites/external_table_p0/lance/test_lance_runtime_filter_pushdown.groovy b/regression-test/suites/external_table_p0/lance/test_lance_runtime_filter_pushdown.groovy new file mode 100644 index 00000000000000..870f51cd7b2b9a --- /dev/null +++ b/regression-test/suites/external_table_p0/lance/test_lance_runtime_filter_pushdown.groovy @@ -0,0 +1,158 @@ +// 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_lance_runtime_filter_pushdown", "p0,external") { + String enabled = context.config.otherConfigs.get("enableIcebergTest") + if (enabled == null || !enabled.equalsIgnoreCase("true")) { + logger.info("disable Lance runtime-filter test because the Iceberg MinIO environment is disabled.") + return + } + + String externalEnvIp = context.config.otherConfigs.get("externalEnvIp") + String minioPort = context.config.otherConfigs.get("iceberg_minio_port") + String catalogName = "test_lance_runtime_filter_pushdown" + String buildTable = "test_lance_runtime_filter_build" + String internalDb = context.dbName + String lanceTable = "`${catalogName}`.`doris`.`predicate_pushdown`" + + sql "SWITCH internal" + sql "USE `${internalDb}`" + sql "DROP TABLE IF EXISTS `${buildTable}`" + sql "DROP CATALOG IF EXISTS `${catalogName}`" + + try { + sql """ + CREATE TABLE `${buildTable}` ( + id BIGINT NOT NULL + ) ENGINE=OLAP + DUPLICATE KEY(id) + DISTRIBUTED BY HASH(id) BUCKETS 1 + PROPERTIES ( + "replication_allocation" = "tag.location.default: 1" + ) + """ + sql "INSERT INTO `${buildTable}` VALUES (4), (7), (9), (11)" + + sql """ + CREATE CATALOG `${catalogName}` PROPERTIES ( + "type" = "lance", + "lance.catalog.type" = "filesystem", + "warehouse" = "s3://warehouse/lance", + "s3.endpoint" = "http://${externalEnvIp}:${minioPort}", + "s3.access_key" = "admin", + "s3.secret_key" = "password", + "s3.region" = "us-east-1", + "use_path_style" = "true" + ) + """ + + sql "SET enable_file_scanner_v2 = true" + sql "SET enable_sql_cache = false" + sql "SET enable_query_cache = false" + sql "SET runtime_filter_mode = 'GLOBAL'" + sql "SET runtime_filter_type = 'IN'" + sql "SET runtime_filter_wait_infinitely = true" + sql "SET enable_runtime_filter_prune = false" + + String query = """ + SELECT /*+ leading(l broadcast b) */ l.row_id, l.int64_value + FROM ${lanceTable} l + INNER JOIN `internal`.`${internalDb}`.`${buildTable}` b + ON l.row_id = b.id + WHERE l.int64_value >= 10 + ORDER BY l.row_id + """ + + // The regular WHERE predicate is converted to the primary Substrait filter. The join + // produces an IN runtime filter on row_id, which must be attached to the Lance scan. + explain { + sql "verbose ${query}" + check { explainString -> + assertTrue(explainString.contains("VLANCE_SCAN_NODE")) + assertTrue(explainString.contains("lancePushdownPredicate=")) + assertTrue(explainString.contains("int64_value")) + assertTrue(explainString.contains("runtime filters: RF")) + assertTrue(explainString.contains("-> row_id")) + return true + } + } + + // Build-side IDs are 4, 7, 9 and 11. The static Lance predicate keeps only rows whose + // int64_value is at least 10, so the intersection is exactly rows 7 and 9. In particular, + // this catches an implementation that replaces the Substrait filter with the later RF. + qt_runtime_filter_and_substrait "${query}" + + // Compare Doris's dedicated two-phase row-id fetch with the equivalent SQL formulation: + // first produce narrow ANN candidates, then broadcast them and fetch payload columns from + // a normal Lance scan. Both queries must return the same ordered rows and distances. + String vectorTable = "${catalogName}.doris.vs_ivf_pq_f32" + String qualifiedVectorTable = "`${catalogName}`.`doris`.`vs_ivf_pq_f32`" + String headQuery = "[0,1,2,3,4,5,6,7,8,9,10,11,12,13,14,15]" + String vectorSearch = """vector_search( + "table"="${vectorTable}", + "column"="embedding", + "query_vector"="${headQuery}", + "top_k"="5", + "metric"="l2", + "nprobes"="4", + "refine_factor"="10", + "use_index"="true")""" + String twoPhaseQuery = """ + SELECT row_id, category, label, _distance + FROM ${vectorSearch} + ORDER BY _distance, row_id + """ + + sql "SET topn_lazy_materialization_threshold = 1024" + explain { + sql "verbose ${twoPhaseQuery}" + contains "VMaterializeNode" + contains "__DORIS_GLOBAL_ROWID_COL__vector_search" + } + qt_two_phase_vector_search "${twoPhaseQuery}" + + String explicitJoinQuery = """ + WITH candidates AS ( + SELECT row_id, _distance + FROM ${vectorSearch} + ORDER BY _distance, row_id + LIMIT 5 + ) + SELECT w.row_id, w.category, w.label, c._distance + FROM ${qualifiedVectorTable} w + JOIN [broadcast] candidates c ON w.row_id = c.row_id + ORDER BY c._distance, c.row_id + """ + + sql "SET topn_lazy_materialization_threshold = -1" + sql "SET disable_join_reorder = true" + explain { + sql "verbose ${explicitJoinQuery}" + contains "VHASH JOIN" + contains "JOIN(BROADCAST)" + contains "runtime filters: RF" + contains "-> row_id" + notContains "VMaterializeNode" + } + qt_explicit_join_vector_search "${explicitJoinQuery}" + } finally { + // sql "SWITCH internal" + // sql "USE `${internalDb}`" + // sql "DROP TABLE IF EXISTS `${buildTable}`" + // sql "DROP CATALOG IF EXISTS `${catalogName}`" + } +} diff --git a/thirdparty/download-thirdparty.sh b/thirdparty/download-thirdparty.sh index bb9c5928809170..85631eaadda648 100755 --- a/thirdparty/download-thirdparty.sh +++ b/thirdparty/download-thirdparty.sh @@ -774,7 +774,7 @@ if [[ " ${TP_ARCHIVES[*]} " =~ " PAIMON_CPP " ]]; then echo "Finished patching ${PAIMON_CPP_SOURCE}" fi -# Patch lance-c with the scan execution statistics API from upstream PR #64. +# Apply Doris lance-c patches in dependency order. if [[ " ${TP_ARCHIVES[*]} " =~ " LANCE_C " ]]; then if [[ "${LANCE_C_SOURCE}" == "lance-c-0.1.7" ]]; then cd "${TP_SOURCE_DIR}/${LANCE_C_SOURCE}" @@ -782,6 +782,11 @@ if [[ " ${TP_ARCHIVES[*]} " =~ " LANCE_C " ]]; then patch -p1 <"${TP_PATCH_DIR}/lance-c-0.1.7-pr-64.patch" touch "${PATCHED_MARK}" fi + lance_runtime_filter_mark="patched_mark_runtime_filter" + if [[ ! -f "${lance_runtime_filter_mark}" ]]; then + patch -p1 <"${TP_PATCH_DIR}/lance-c-0.1.7-runtime-filter.patch" + touch "${lance_runtime_filter_mark}" + fi cd - fi echo "Finished patching ${LANCE_C_SOURCE}" diff --git a/thirdparty/patches/lance-c-0.1.7-runtime-filter.patch b/thirdparty/patches/lance-c-0.1.7-runtime-filter.patch new file mode 100644 index 00000000000000..dd12ca19ab7df6 --- /dev/null +++ b/thirdparty/patches/lance-c-0.1.7-runtime-filter.patch @@ -0,0 +1,4904 @@ +diff --git a/AGENTS.md b/AGENTS.md +index a1c0041..f236494 100644 +--- a/AGENTS.md ++++ b/AGENTS.md +@@ -2,8 +2,8 @@ + + ## Structure + Rust FFI source: `src/` +-C header (stable ABI): `include/lance.h` +-C++ RAII wrappers (header-only): `include/lance.hpp` ++C header (stable ABI): `include/lance/lance.h` ++C++ RAII wrappers (header-only): `include/lance/lance.hpp` + Tests (Rust): `tests/c_api_test.rs` + Tests (C/C++): `tests/cpp/` + Historical test data: `test_data/` +@@ -19,7 +19,8 @@ test C/C++ compilation: `cargo test --test compile_and_run_test -- --ignored` + - Opaque handles with `lance_*_open`/`lance_*_close` lifecycle. + - Thread-local error handling via `ffi_try!` macro. + - Arrow C Data Interface for zero-copy data exchange. +-- `panic = "abort"` in release to prevent unwinding across FFI. ++- `panic = "unwind"` is required so guarded FFI boundaries can translate ++ panics to `LANCE_ERR_PANIC`; `panic = "abort"` builds are rejected. + + ## Coding Standards + +diff --git a/Cargo.toml b/Cargo.toml +index 5bc30bd..b5ba59a 100644 +--- a/Cargo.toml ++++ b/Cargo.toml +@@ -25,6 +25,8 @@ lance-index = { git = "https://github.com/lance-format/lance.git", rev = "e934cc + lance-io = { git = "https://github.com/lance-format/lance.git", rev = "e934cc2c" } + lance-linalg = { git = "https://github.com/lance-format/lance.git", rev = "e934cc2c" } + lance-table = { git = "https://github.com/lance-format/lance.git", rev = "e934cc2c" } ++lance-datafusion = { git = "https://github.com/lance-format/lance.git", rev = "e934cc2c", features = ["substrait"] } ++datafusion = { version = "54.0.0", default-features = false } + arrow = { version = "58.0.0", features = ["prettyprint", "ffi"] } + arrow-array = "58.0.0" + arrow-schema = "58.0.0" +@@ -44,10 +46,8 @@ uuid = { version = "1", features = ["v4"] } + + [dev-dependencies] + lance = { git = "https://github.com/lance-format/lance.git", rev = "e934cc2c", features = ["substrait"] } +-lance-datafusion = { git = "https://github.com/lance-format/lance.git", rev = "e934cc2c", features = ["substrait"] } + lance-datagen = { git = "https://github.com/lance-format/lance.git", rev = "e934cc2c" } + lance-file = { git = "https://github.com/lance-format/lance.git", rev = "e934cc2c" } +-datafusion = { version = "54.0.0", default-features = false } + tokio = { version = "1", features = ["rt-multi-thread", "macros"] } + arrow-array = "58.0.0" + arrow-schema = "58.0.0" +diff --git a/README.md b/README.md +index 436c2f1..d6ae6c8 100644 +--- a/README.md ++++ b/README.md +@@ -67,7 +67,7 @@ Based on the [liblance RFC](https://github.com/lance-format/lance/discussions/60 + |--------|-----------|-------------| + | [x] | Async scan | Callback-based `lance_scanner_scan_async()` for non-blocking scans | + | [x] | Dataset metadata | `lance_dataset_version()`, `lance_dataset_count_rows()`, `lance_dataset_latest_version()` | +-| [x] | Substrait filter pushdown | `lance_scanner_set_substrait_filter()` accepts a serialized Substrait `ExtendedExpression` (preferred over SQL strings for query engines) | ++| [x] | Filter pushdown | `lance_scanner_set_substrait_filter()` accepts a serialized Substrait `ExtendedExpression`; `lance_scanner_additional_sql_filter()` adds SQL predicates with AND before scanning starts | + + ## Building + +diff --git a/include/lance/lance.h b/include/lance/lance.h +index 5b12f3d..ccd7e8c 100644 +--- a/include/lance/lance.h ++++ b/include/lance/lance.h +@@ -7,10 +7,15 @@ + * + * All data crosses this boundary via the Arrow C Data Interface + * (ArrowSchema, ArrowArray, ArrowArrayStream). ++ * For Arrow structures written to caller-provided output storage, the caller ++ * retains ownership of the outer structure and must invoke its non-NULL ++ * `release` callback exactly once to release the contents. APIs that allocate ++ * the outer structure as well document a separate matching free function. + * +- * Error handling uses thread-local storage: after any function returns +- * NULL (pointer) or -1 (int), call lance_last_error_code() and +- * lance_last_error_message() to get details. ++ * Error handling uses thread-local storage: after any function returns its ++ * documented error sentinel (for example NULL, -1, or 0 for selected scalar ++ * accessors), call lance_last_error_code() and lance_last_error_message() to ++ * get details. + */ + + #ifndef LANCE_H +@@ -100,16 +105,21 @@ typedef enum { + * Honest limits: a double panic, a panic in a destructor while unwinding, a + * stack overflow, or an allocation failure still aborts the process. A + * panic caught inside a close/free call (lance_*_close, lance_batch_free, +- * lance_free_string, or the release callback of an exported +- * ArrowArrayStream) is logged and the remainder of the value may leak — +- * close is best-effort by design. Post-panic process state is best-effort: +- * hosts should fail the in-flight query rather than retry a poisoned +- * handle. ++ * lance_free_string, lance_scanner_async_stream_free, or the release callback ++ * of an exported ArrowArrayStream) is logged and the remainder of the value ++ * may leak — close is best-effort by design. Post-panic process state is ++ * best-effort: hosts should fail the in-flight query rather than retry a ++ * poisoned handle. + * +- * Callbacks passed INTO the library (LanceCallback, LanceWaker) are the +- * reverse direction and are NOT covered by this contract: their ABI is +- * non-unwinding, so a panicking callback aborts the host process before +- * the library can contain it. Callbacks must not panic. ++ * Callbacks passed INTO the library (LanceCallback, LanceWaker, and ++ * LanceScanStatisticsCallback) are the reverse direction and are NOT covered ++ * by this contract: their ABI is non-unwinding, so a callback that throws or ++ * unwinds can abort the host process before the library can contain it. ++ * Callbacks must return normally. ++ * ++ * This contract requires Rust's `panic = "unwind"` strategy. The crate ++ * rejects `panic = "abort"` builds at compile time because catch_unwind ++ * cannot provide this API contract in such a build. + */ + + /* ─── Index types (Phase 2) ─── */ +@@ -175,6 +185,7 @@ typedef struct LanceVersions LanceVersions; + typedef struct LanceDataStatistics LanceDataStatistics; + typedef struct LanceIndexSegmentBuilder LanceIndexSegmentBuilder; + typedef struct LanceIndexSegmentMetadata LanceIndexSegmentMetadata; ++typedef struct LanceFtsQueryContext LanceFtsQueryContext; + + /* ─── Dataset lifecycle ─── */ + +@@ -202,13 +213,22 @@ void lance_dataset_close(LanceDataset* dataset); + + /* ─── Dataset metadata (sync, in-memory) ─── */ + +-/** Return the version number of this dataset snapshot. */ ++/** ++ * Return the version number of this dataset snapshot. ++ * @return version on success, or 0 on error (check lance_last_error_code()) ++ */ + uint64_t lance_dataset_version(const LanceDataset* dataset); + +-/** Return the number of rows. Returns 0 on error. */ ++/** ++ * Return the number of rows. Returns 0 on error; an empty dataset also returns ++ * 0, so check lance_last_error_code(). ++ */ + uint64_t lance_dataset_count_rows(const LanceDataset* dataset); + +-/** Return the latest version ID (I/O). Returns 0 on error. */ ++/** ++ * Return the latest version ID (I/O), or 0 on error (check ++ * lance_last_error_code()). ++ */ + uint64_t lance_dataset_latest_version(const LanceDataset* dataset); + + /* ─── Version history ─── */ +@@ -220,7 +240,10 @@ uint64_t lance_dataset_latest_version(const LanceDataset* dataset); + */ + LanceVersions* lance_dataset_versions(const LanceDataset* dataset); + +-/** Number of versions in the snapshot. Returns 0 on error. */ ++/** ++ * Number of versions in the snapshot, or 0 on error (check ++ * lance_last_error_code()). ++ */ + uint64_t lance_versions_count(const LanceVersions* versions); + + /** +@@ -755,7 +778,10 @@ int32_t lance_dataset_schema( + + /* ─── Fragment enumeration ─── */ + +-/** Return the number of fragments in the dataset. Returns 0 on error. */ ++/** ++ * Return the number of fragments in the dataset. Returns 0 on error; a ++ * dataset with no fragments also returns 0, so check lance_last_error_code(). ++ */ + uint64_t lance_dataset_fragment_count(const LanceDataset* dataset); + + /** +@@ -769,6 +795,15 @@ int32_t lance_dataset_fragment_ids(const LanceDataset* dataset, uint64_t* out_id + + /** + * Take rows by indices. ++ * ++ * On success, `out` is initialized in caller-owned storage; the caller must ++ * eventually invoke its non-NULL `release` callback exactly once. The schema ++ * is validated before the stream callbacks are exposed. A deferred iteration ++ * failure, including a caught panic in `get_next`, is reported through the ++ * Arrow C stream contract (nonzero `get_next` plus `get_last_error`). A panic ++ * during `release` cleanup is contained and logged; cleanup remains ++ * best-effort. ++ * + * @param indices Array of 0-based row offsets + * @param num_indices Length of indices array + * @param columns NULL-terminated column names, or NULL for all +@@ -791,6 +826,14 @@ int32_t lance_dataset_take( + * Missing or deleted row IDs may be omitted from the result. For found rows, + * input order and duplicates are preserved. + * ++ * On success, `out` is initialized in caller-owned storage; the caller must ++ * eventually invoke its non-NULL `release` callback exactly once. The schema ++ * is validated before the stream callbacks are exposed. A deferred iteration ++ * failure, including a caught panic in `get_next`, is reported through the ++ * Arrow C stream contract (nonzero `get_next` plus `get_last_error`). A panic ++ * during `release` cleanup is contained and logged; cleanup remains ++ * best-effort. ++ * + * @param dataset Open dataset snapshot. + * @param row_ids Array of dataset row IDs. May be NULL only when + * `num_row_ids` is zero. +@@ -863,6 +906,22 @@ int32_t lance_scanner_set_substrait_filter( + size_t len + ); + ++/** ++ * Add an SQL filter that is combined with the selected primary filter using ++ * AND. The primary filter is the Substrait filter when set, otherwise it is ++ * the SQL filter passed to `lance_scanner_new`. Multiple additional SQL ++ * filters are also combined using AND. ++ * ++ * Must be called before the scan starts. The filter string is copied. ++ * ++ * @param filter Non-NULL, non-empty SQL filter expression ++ * @return 0 on success, -1 on error ++ */ ++int32_t lance_scanner_additional_sql_filter( ++ LanceScanner* scanner, ++ const char* filter ++); ++ + /** Type of a dynamically named scan metric. */ + typedef enum { + LANCE_SCAN_METRIC_COUNT = 0, +@@ -967,7 +1026,16 @@ int32_t lance_scanner_set_statistics_callback( + void* callback_ctx + ); + +-/** Close and free a scanner handle. */ ++/** ++ * Close and free a scanner handle. Safe to call with NULL; a non-NULL handle ++ * must be closed exactly once. ++ * ++ * This is the retirement boundary for poll wakers registered by ++ * lance_scanner_poll_next(): it cancels callbacks that have not entered and ++ * waits for any callback already in progress to return before freeing the ++ * scanner. Do not call this function from one of the scanner's own waker ++ * callbacks, because close must wait for that callback to return. ++ */ + void lance_scanner_close(LanceScanner* scanner); + + /* ─── Sync scan: ArrowArrayStream ─── */ +@@ -975,6 +1043,10 @@ void lance_scanner_close(LanceScanner* scanner); + /** + * Materialize the scan as an ArrowArrayStream (blocking). + * The scanner remains valid, and each call creates an independent stream. ++ * `out` points to caller-owned storage. On success, the caller must eventually ++ * invoke `out->release(out)` exactly once when `release` is non-NULL; that ++ * releases the stream contents but not the caller-owned outer structure. Do ++ * not pass this caller-allocated stream to lance_scanner_async_stream_free(). + * + * Reading the exported stream may surface a mid-iteration panic as one + * error through the Arrow C stream contract (nonzero get_next plus +@@ -1005,14 +1077,18 @@ int32_t lance_scanner_next( + /** + * Callback type for async operations. + * +- * The callback runs on the dispatcher thread; on failure the error code and +- * message are installed in that thread's thread-local storage immediately +- * before the callback runs, so lance_last_error_* called from inside the +- * callback observes this completion's failure. ++ * The callback normally runs on the dedicated dispatcher thread. During a ++ * rare dispatcher startup or delivery failure, completion falls back to the ++ * thread that detects the failure (for example the calling or producing ++ * thread), so the callback must be thread-safe. On failure the error code and ++ * message are installed on the actual callback thread immediately before the ++ * callback runs, so lance_last_error_* called from inside the callback ++ * observes this completion's failure. + * +- * Callbacks must not panic: the callback ABI is non-unwinding, so a +- * panicking callback aborts the host process before the dispatcher can +- * contain it. ++ * Callbacks must return normally: the callback ABI is non-unwinding, so a ++ * callback that throws or unwinds can abort the host process before the ++ * dispatcher can contain it. ++ * A callback passed to lance_scanner_scan_async() must not be NULL. + * + * @param ctx Opaque pointer passed back from the caller + * @param status 0 = success, -1 = error +@@ -1021,14 +1097,28 @@ int32_t lance_scanner_next( + typedef void (*LanceCallback)(void* ctx, int32_t status, void* result); + + /** +- * Start an async scan. The callback fires on a dedicated dispatcher thread +- * when the ArrowArrayStream is ready. ++ * Start an async scan. The callback normally fires on a dedicated dispatcher ++ * thread when the ArrowArrayStream is ready. During a rare dispatcher ++ * infrastructure failure it may instead run on the calling or producing ++ * thread, so it must be thread-safe. ++ * ++ * For a non-NULL callback, exactly one completion is delivered, including for ++ * validation, setup, task, and dispatcher failures. The fallback path may ++ * invoke it before lance_scanner_scan_async() returns. `callback` and a ++ * non-NULL `callback_ctx` must remain valid until that invocation returns. ++ * ++ * `callback` must not be NULL; `callback_ctx` may be NULL. On success, result ++ * is a library-allocated ArrowArrayStream owned by the caller. The caller must ++ * eventually pass it exactly once to lance_scanner_async_stream_free(), even ++ * if it has already invoked the stream's release callback directly. Do not ++ * free the returned outer structure with free(), delete, or a platform ++ * allocator. + * + * On failure the callback receives status -1 with result NULL, and the +- * error code/message are installed in the dispatcher thread's thread-local +- * storage immediately before the callback runs (per completion). A panic in +- * the scan task also yields status -1 with LANCE_ERR_PANIC and poisons the +- * scanner handle. ++ * error code/message are installed in the actual callback thread's ++ * thread-local storage immediately before the callback runs (per completion). ++ * A panic in the scan task also yields status -1 with LANCE_ERR_PANIC and ++ * poisons the scanner handle. + */ + void lance_scanner_scan_async( + const LanceScanner* scanner, +@@ -1036,6 +1126,19 @@ void lance_scanner_scan_async( + void* callback_ctx + ); + ++/** ++ * Release and free an ArrowArrayStream returned by a successful ++ * lance_scanner_scan_async() callback. ++ * ++ * If `stream->release` is non-NULL, this function invokes it before freeing ++ * the library-allocated outer structure. It is therefore valid both before ++ * and after a consumer has directly released the stream contents. `stream` ++ * may be NULL. A non-NULL pointer must be passed exactly once and must be the ++ * pointer delivered by lance_scanner_scan_async(); using this function for a ++ * caller-allocated ArrowArrayStream is invalid. ++ */ ++void lance_scanner_async_stream_free(struct ArrowArrayStream* stream); ++ + /* ─── Poll-based scan (for cooperative async runtimes) ─── */ + + typedef enum { +@@ -1045,12 +1148,26 @@ typedef enum { + LANCE_POLL_ERROR = -1, + } LancePollStatus; + +-/** Waker callback: called from a Tokio thread when data is ready. */ ++/** ++ * Waker callback: called from a Tokio thread when data is ready. A waker ++ * passed to lance_scanner_poll_next() must not be NULL. For one poll call ++ * that returns LANCE_POLL_PENDING, all internal RawWaker clones share a ++ * one-shot gate, so the callback fires at most once. ++ * ++ * The callback and `ctx` must be thread-safe and must remain valid until the ++ * callback returns or lance_scanner_close() returns. Close cancels a pending ++ * callback and waits for an active callback before returning, so the caller ++ * may destroy `ctx` afterwards. The callback must return normally and must ++ * not call lance_scanner_close() or otherwise re-enter its originating ++ * scanner. ++ */ + typedef void (*LanceWaker)(void* ctx); + + /** + * Poll for the next batch without blocking. +- * See RFC for usage pattern. ++ * `waker` must not be NULL; `waker_ctx` may be NULL. `out` is set to a ++ * LanceBatch only for LANCE_POLL_READY and is set to NULL for ++ * LANCE_POLL_PENDING, LANCE_POLL_FINISHED, and LANCE_POLL_ERROR. + */ + LancePollStatus lance_scanner_poll_next( + LanceScanner* scanner, +@@ -1329,7 +1446,10 @@ const char* lance_index_segment_metadata_name( + const LanceIndexSegmentMetadata* metadata + ); + +-/** Return the dataset version against which the segment was built. */ ++/** ++ * Return the dataset version against which the segment was built, or 0 on ++ * error (check lance_last_error_code()). ++ */ + uint64_t lance_index_segment_metadata_dataset_version( + const LanceIndexSegmentMetadata* metadata + ); +@@ -1355,7 +1475,10 @@ const char* lance_index_segment_metadata_index_details_type_url( + const LanceIndexSegmentMetadata* metadata + ); + +-/** Return the number of indexed field IDs. */ ++/** ++ * Return the number of indexed field IDs. Returns 0 on error; zero may also be ++ * a valid count, so check lance_last_error_code(). ++ */ + size_t lance_index_segment_metadata_field_count( + const LanceIndexSegmentMetadata* metadata + ); +@@ -1368,7 +1491,10 @@ int32_t lance_index_segment_metadata_field_ids( + size_t* out_count + ); + +-/** Return the number of fragment IDs covered by the segment. */ ++/** ++ * Return the number of fragment IDs covered by the segment. Returns 0 on ++ * error; zero may also be a valid count, so check lance_last_error_code(). ++ */ + size_t lance_index_segment_metadata_fragment_count( + const LanceIndexSegmentMetadata* metadata + ); +@@ -1393,7 +1519,11 @@ void lance_index_segment_metadata_free(LanceIndexSegmentMetadata* metadata); + /** Drop an index by name. Returns -1 (NOT_FOUND) if no such index. */ + int32_t lance_dataset_drop_index(LanceDataset* dataset, const char* name); + +-/** Number of user indexes (excludes system indexes). Returns 0 on error. */ ++/** ++ * Number of user indexes (excludes system indexes). Returns 0 on error; a ++ * dataset with no user indexes also returns 0, so check ++ * lance_last_error_code(). ++ */ + uint64_t lance_dataset_index_count(const LanceDataset* dataset); + + /** +@@ -1489,6 +1619,56 @@ int32_t lance_scanner_set_index_segments( + + /* ─── Full-text search (Phase 2) ─── */ + ++/** ++ * Required relationship between a pinned dataset snapshot and its committed ++ * FTS index segments. Values are ABI-stable; API parameters use int32_t. ++ */ ++typedef enum { ++ /** Fail prepare if any current fragment is not covered by the FTS index. */ ++ LANCE_FTS_COVERAGE_STRICT = 0, ++ /** Score and search only rows covered by committed FTS index segments. */ ++ LANCE_FTS_COVERAGE_INDEX_ONLY = 1, ++} LanceFtsCoverageMode; ++ ++/** ++ * Prepare an immutable, process-local FTS query context for one column. ++ * ++ * Preparation pins the dataset handle's current snapshot, enumerates all ++ * committed FTS segments for `column`, checks fragment coverage, opens those ++ * segments, and computes one query-specific global BM25 scorer across their ++ * indexed documents. The context can then be shared by any number of scanners ++ * created from the exact same process-local dataset snapshot. It has no ++ * serialization or cross-process transport format. Reopening the same URI and ++ * manifest version creates a different identity and cannot reuse the context, ++ * because storage options and object-store endpoints may differ. ++ * ++ * In LANCE_FTS_COVERAGE_INDEX_ONLY mode, unindexed fragments are allowed and ++ * excluded from both the scorer corpus and query results. In STRICT mode any ++ * unindexed fragment makes this call fail. ++ * ++ * Prepared contexts currently support exact Match queries only. ++ * `max_fuzzy_distance` must be zero because fuzzy execution requires its ++ * canonical expanded vocabulary to be prepared together with the scorer. ++ * This restriction does not apply to lance_scanner_full_text_search(). ++ * ++ * @param max_fuzzy_distance Must be zero for prepared query contexts. ++ * @param coverage_mode Fixed-width LanceFtsCoverageMode discriminant. ++ * @return Context handle on success, or NULL on error. ++ */ ++LanceFtsQueryContext* lance_dataset_prepare_fts_query( ++ const LanceDataset* dataset, ++ const char* column, ++ const char* query, ++ uint32_t max_fuzzy_distance, ++ int32_t coverage_mode ++); ++ ++/** ++ * Close a context handle. NULL-safe. Scanners that already attached this ++ * context retain shared ownership and remain valid. ++ */ ++void lance_fts_query_context_close(LanceFtsQueryContext* context); ++ + /** + * Set a BM25 full-text search query on the scanner. + * +@@ -1508,6 +1688,30 @@ int32_t lance_scanner_full_text_search( + uint32_t max_fuzzy_distance + ); + ++/** ++ * Attach a prepared process-local FTS query context. The scanner must have ++ * been created from the exact LanceDataset snapshot used to prepare the ++ * context; URI and manifest version equality is not sufficient. The scanner ++ * retains shared ownership, so the caller may close `context` after success. ++ * This is mutually exclusive with nearest and lance_scanner_full_text_search ++ * because the context already owns the FTS query. ++ */ ++int32_t lance_scanner_set_fts_query_context( ++ LanceScanner* scanner, ++ const LanceFtsQueryContext* context ++); ++ ++/** ++ * Restrict a context-backed FTS scan to `len` context segment UUIDs supplied ++ * by the caller's planner. Pass `len == 0` to clear the restriction and search ++ * all context segments. Duplicate or unknown UUIDs are rejected. ++ */ ++int32_t lance_scanner_set_fts_index_segments( ++ LanceScanner* scanner, ++ const uint8_t* segment_uuids, ++ size_t len ++); ++ + /* ─── Dataset writer ─── */ + + /** +diff --git a/include/lance/lance.hpp b/include/lance/lance.hpp +index 8aa97e2..330e419 100644 +--- a/include/lance/lance.hpp ++++ b/include/lance/lance.hpp +@@ -49,6 +49,12 @@ inline void check_error() { + } + } + ++/// Release and free a library-allocated ArrowArrayStream returned by ++/// Scanner::scan_async. NULL-safe; do not use for caller-allocated streams. ++inline void scanner_async_stream_free(ArrowArrayStream* stream) noexcept { ++ lance_scanner_async_stream_free(stream); ++} ++ + // ─── RAII Handle Template ──────────────────────────────────────────────────── + + template +@@ -89,6 +95,7 @@ class Scanner; + class IndexModel; + class IndexSegmentBuilder; + class IndexSegmentMetadata; ++class FtsQueryContext; + + // ─── Version history ───────────────────────────────────────────────────────── + +@@ -116,6 +123,11 @@ enum class WriteMode : int32_t { + Overwrite = LANCE_WRITE_OVERWRITE, + }; + ++enum class FtsCoverageMode : int32_t { ++ Strict = LANCE_FTS_COVERAGE_STRICT, ++ IndexOnly = LANCE_FTS_COVERAGE_INDEX_ONLY, ++}; ++ + /// Tunable parameters for Dataset::write. Numeric fields default-out via 0; + /// `data_storage_version` defaults out via `std::nullopt`. + /// +@@ -157,6 +169,24 @@ struct SqlColumn { + std::string expression; + }; + ++// ─── Process-local FTS query context ──────────────────────────────────────── ++ ++/// Immutable, query-specific global BM25 scorer plus pinned FTS segment list. ++/// This handle is process-local and intentionally has no serialization API. ++class FtsQueryContext { ++ Handle handle_; ++ ++public: ++ explicit FtsQueryContext(LanceFtsQueryContext* context) : handle_(context) {} ++ ++ FtsQueryContext(FtsQueryContext&&) noexcept = default; ++ FtsQueryContext& operator=(FtsQueryContext&&) noexcept = default; ++ FtsQueryContext(const FtsQueryContext&) = delete; ++ FtsQueryContext& operator=(const FtsQueryContext&) = delete; ++ ++ const LanceFtsQueryContext* c_handle() const { return handle_.get(); } ++}; ++ + // ─── Dataset ───────────────────────────────────────────────────────────────── + + class Dataset { +@@ -329,7 +359,9 @@ public: + + /// Version of this dataset snapshot. + uint64_t version() const { +- return lance_dataset_version(handle_.get()); ++ uint64_t v = lance_dataset_version(handle_.get()); ++ if (lance_last_error_code() != LANCE_OK) check_error(); ++ return v; + } + + /// Latest version ID (queries object store). +@@ -347,11 +379,13 @@ public: + Handle snap(raw); + + uint64_t n = lance_versions_count(snap.get()); ++ if (lance_last_error_code() != LANCE_OK) check_error(); + std::vector out; + out.reserve(static_cast(n)); + for (uint64_t i = 0; i < n; i++) { + VersionInfo info; + info.id = lance_versions_id_at(snap.get(), static_cast(i)); ++ if (lance_last_error_code() != LANCE_OK) check_error(); + info.timestamp_ms = + lance_versions_timestamp_ms_at(snap.get(), static_cast(i)); + if (lance_last_error_code() != LANCE_OK) check_error(); +@@ -369,11 +403,13 @@ public: + Handle snap(raw); + + uint64_t n = lance_data_statistics_count(snap.get()); ++ if (lance_last_error_code() != LANCE_OK) check_error(); + std::vector out; + out.reserve(static_cast(n)); + for (uint64_t i = 0; i < n; i++) { + FieldStatistics fs; + fs.id = lance_data_statistics_field_id_at(snap.get(), static_cast(i)); ++ if (lance_last_error_code() != LANCE_OK) check_error(); + fs.bytes_on_disk = + lance_data_statistics_bytes_on_disk_at(snap.get(), static_cast(i)); + if (lance_last_error_code() != LANCE_OK) check_error(); +@@ -631,7 +667,9 @@ public: + } + } + +- /// Take rows by indices. Results exported as ArrowArrayStream. ++ /// Take rows by indices. `out` is caller-owned and its non-null `release` ++ /// must be called exactly once. Deferred iteration/cleanup panics are ++ /// contained by the exported stream guard. + void take(const uint64_t* indices, size_t num_indices, + const std::vector& columns, + ArrowArrayStream* out) const { +@@ -645,7 +683,7 @@ public: + } + } + +- /// Take all columns. ++ /// Take all columns with the same stream ownership as the overload above. + void take(const uint64_t* indices, size_t num_indices, + ArrowArrayStream* out) const { + if (lance_dataset_take(handle_.get(), indices, num_indices, nullptr, out) != 0) { +@@ -653,7 +691,9 @@ public: + } + } + +- /// Take rows by dataset row IDs. Results exported as ArrowArrayStream. ++ /// Take rows by dataset row IDs. `out` is caller-owned and its non-null ++ /// `release` must be called exactly once. Deferred iteration/cleanup panics ++ /// are contained by the exported stream guard. + void take_rows(const uint64_t* row_ids, size_t num_row_ids, + const std::vector& columns, + ArrowArrayStream* out) const { +@@ -668,7 +708,8 @@ public: + } + } + +- /// Take all columns by dataset row IDs. ++ /// Take all columns by dataset row IDs with the same stream ownership as ++ /// the overload above. + void take_rows(const uint64_t* row_ids, size_t num_row_ids, + ArrowArrayStream* out) const { + if (lance_dataset_take_rows( +@@ -680,6 +721,23 @@ public: + /// Create a Scanner builder for this dataset. + Scanner scan() const; + ++ /// Prepare a query-specific global BM25 scorer over the committed FTS ++ /// segments of this pinned snapshot. IndexOnly permits unindexed fragments; ++ /// Strict rejects them. Prepared contexts currently require ++ /// `max_fuzzy_distance == 0`. The context can only be attached to scanners ++ /// created from this exact process-local dataset snapshot. ++ FtsQueryContext prepare_fts_query( ++ const std::string& column, ++ const std::string& query, ++ uint32_t max_fuzzy_distance = 0, ++ FtsCoverageMode coverage_mode = FtsCoverageMode::Strict) const { ++ auto* context = lance_dataset_prepare_fts_query( ++ handle_.get(), column.c_str(), query.c_str(), max_fuzzy_distance, ++ static_cast(coverage_mode)); ++ if (!context) check_error(); ++ return FtsQueryContext(context); ++ } ++ + /// Number of fragments in the dataset. + uint64_t fragment_count() const { + uint64_t n = lance_dataset_fragment_count(handle_.get()); +@@ -783,7 +841,7 @@ public: + /// Throws lance::Error with code NotFound if the index does not exist. + uint64_t index_segment_count(const std::string& index_name) const { + uint64_t n = lance_dataset_index_segment_count(handle_.get(), index_name.c_str()); +- if (n == 0 && lance_last_error_code() != LANCE_OK) check_error(); ++ if (lance_last_error_code() != LANCE_OK) check_error(); + return n; + } + +@@ -1127,7 +1185,14 @@ public: + return substrait_filter(bytes.data(), bytes.size()); + } + +- /// Register a callback for scan statistics after successful full exhaustion. ++ /// Add an SQL filter that is combined with the selected primary filter using AND. ++ Scanner& additional_sql_filter(const std::string& filter) { ++ if (lance_scanner_additional_sql_filter(handle_.get(), filter.c_str()) != 0) ++ check_error(); ++ return *this; ++ } ++ ++ /// Register a non-null callback for scan statistics after successful full exhaustion. + /// The registration applies to every stream derived from this scanner, including + /// concurrent streams and streams created after an earlier callback returns. The + /// callback is not guaranteed on error, cancellation, or early release. It may +@@ -1166,12 +1231,20 @@ public: + } + + /// Materialize an independent ArrowArrayStream (blocking). The scanner remains valid. ++ /// `out` is caller-owned; call its non-null `release` callback exactly once. + void to_arrow_stream(ArrowArrayStream* out) { + if (lance_scanner_to_arrow_stream(handle_.get(), out) != 0) + check_error(); + } + +- /// Start an async scan. Callback fires when ArrowArrayStream is ready. ++ /// Start an async scan with a non-null callback. On success, the callback's ++ /// ArrowArrayStream result is library-allocated and must be passed exactly ++ /// once to `lance::scanner_async_stream_free`, which also invokes `release` ++ /// when necessary. The callback normally runs on the dispatcher thread, ++ /// but a rare infrastructure fallback may invoke it on the calling or ++ /// producing thread, possibly before this method returns, so it must be ++ /// thread-safe. Exactly one completion is delivered; callback and non-null ++ /// context storage must remain valid until it returns. + void scan_async(LanceCallback callback, void* ctx) const { + lance_scanner_scan_async(handle_.get(), callback, ctx); + } +@@ -1235,6 +1308,29 @@ public: + return *this; + } + ++ /// Attach a process-local prepared FTS query context. The scanner retains ++ /// shared ownership, so the context object may be destroyed after success. ++ Scanner& fts_query_context(const FtsQueryContext& context) { ++ if (lance_scanner_set_fts_query_context(handle_.get(), context.c_handle()) != 0) ++ check_error(); ++ return *this; ++ } ++ ++ /// Restrict a context-backed FTS query to a segment UUID subset. ++ Scanner& fts_index_segments(const uint8_t* segment_uuids, size_t segment_count) { ++ if (lance_scanner_set_fts_index_segments( ++ handle_.get(), segment_uuids, segment_count) != 0) ++ check_error(); ++ return *this; ++ } ++ ++ Scanner& fts_index_segments( ++ const std::vector>& segment_uuids) { ++ return fts_index_segments( ++ reinterpret_cast(segment_uuids.data()), ++ segment_uuids.size()); ++ } ++ + /// Access the underlying C handle. + LanceScanner* c_handle() { return handle_.get(); } + }; +diff --git a/src/add_columns.rs b/src/add_columns.rs +index 5300c26..a6f30b8 100644 +--- a/src/add_columns.rs ++++ b/src/add_columns.rs +@@ -176,7 +176,7 @@ unsafe fn add_columns_nulls_inner( + let ffi_schema = unsafe { &*schema }; + // Reject an already-released or never-initialised schema before handing it + // to arrow-rs, which would otherwise `assert!` on the NULL `format` field +- // and abort the host process under our `panic = "abort"` profile. Both ++ // and turn predictable invalid input into LANCE_ERR_PANIC. Both + // checks are intentional — `release == NULL` is the canonical Arrow C Data + // Interface "released" sentinel, while `format == NULL` catches a + // zero-initialised or half-built struct that would slip past the release +@@ -187,8 +187,8 @@ unsafe fn add_columns_nulls_inner( + )); + } + // arrow-rs's `FFI_ArrowSchema::format()` does `to_str().expect(..)` on the +- // format pointer; a non-NULL but non-UTF-8 top-level format would abort the +- // process under `panic = "abort"`. Validate it here so a malformed format ++ // format pointer; a non-NULL but non-UTF-8 top-level format would panic in ++ // the guarded FFI boundary. Validate it here so a malformed format + // surfaces as INVALID_ARGUMENT instead. (Child fields are still the caller's + // responsibility — see the doc comment — as walking them would duplicate + // arrow-rs's recursive descent.) +@@ -273,8 +273,8 @@ unsafe fn add_columns_stream_inner( + // Reject a stream missing a mandatory C Data Interface callback *before* + // handing it to arrow-rs. `ArrowArrayStreamReader` only guards against a + // NULL `release`; a NULL `get_schema` or `get_next` would otherwise reach an +- // `unwrap()` deep inside arrow-rs and abort the host process under our +- // `panic = "abort"` profile. We do not require `get_last_error` (the spec ++ // `unwrap()` deep inside arrow-rs and turn predictable invalid input into ++ // LANCE_ERR_PANIC. We do not require `get_last_error` (the spec + // marks it optional): requiring it would not close the abort anyway, since a + // present callback that *returns* NULL at error time hits the same + // `last_error.unwrap()` on arrow-rs's `get_next` error path — a residual +diff --git a/src/alter_columns.rs b/src/alter_columns.rs +index 5f85fae..e2da121 100644 +--- a/src/alter_columns.rs ++++ b/src/alter_columns.rs +@@ -208,8 +208,8 @@ unsafe fn parse_alteration( + let ffi_schema = unsafe { &*entry.data_type }; + // Reject an already-released or never-initialised schema before + // handing it to arrow-rs, which would otherwise `assert!` on the +- // NULL `format` field and abort the host process under our +- // `panic = "abort"` profile. Both checks are intentional: ++ // NULL `format` field and turn predictable invalid input into ++ // LANCE_ERR_PANIC. Both checks are intentional: + // - `release == NULL`: the canonical Arrow CADI "released" sentinel. + // - `format == NULL`: catches a zero-initialised or otherwise + // half-built struct that would slip past the release check. +diff --git a/src/async_dispatcher.rs b/src/async_dispatcher.rs +index 91df74d..0112ed5 100644 +--- a/src/async_dispatcher.rs ++++ b/src/async_dispatcher.rs +@@ -42,7 +42,7 @@ struct Dispatcher { + } + + impl Dispatcher { +- fn new() -> Self { ++ fn new() -> std::io::Result { + let (tx, rx) = mpsc::channel::(); + + std::thread::Builder::new() +@@ -50,48 +50,62 @@ impl Dispatcher { + .spawn(move || { + log::debug!("Lance C dispatcher thread started"); + while let Ok(msg) = rx.recv() { +- // Install the carried error on THIS thread's TLS so the +- // callback's `lance_last_error_*` calls observe it. TLS +- // persists across callbacks on this thread, so a success +- // must explicitly clear: a stale error from an earlier +- // failed callback must never leak into a later one. +- match &msg.error { +- Some((code, message)) => set_last_error(*code, message), +- None => clear_last_error(), +- } +- // Invoke the C callback under catch_unwind, best-effort +- // only (issue #61). The declared callback ABI is +- // `extern "C"` and therefore NON-unwinding — `lance.h` +- // requires callbacks not to panic, and a panic in such a +- // callback aborts at its own boundary before this catch +- // could ever run. The catch exists solely for Rust hosts +- // that pass an `extern "C-unwind"` callback: for them it +- // keeps the dispatcher thread (and with it every later +- // async completion) alive. It is not part of the panic +- // contract and must never be relied on as one. +- let outcome = catch_unwind(AssertUnwindSafe(|| unsafe { +- (msg.callback)(msg.callback_ctx, msg.status, msg.result); +- })); +- if let Err(payload) = outcome { +- log::error!( +- "lance-c dispatcher: unwinding (C-unwind) host callback panicked; contained best-effort: {}", +- panic_payload_message(&*payload) +- ); +- } ++ deliver_message(msg); + } + log::debug!("Lance C dispatcher thread shutting down"); +- }) +- .expect("Failed to spawn lance-c dispatcher thread"); ++ })?; + +- Self { tx } ++ Ok(Self { tx }) + } + +- fn send(&self, msg: DispatcherMessage) { +- let _ = self.tx.send(msg); ++ fn send(&self, msg: DispatcherMessage) -> Result<(), DispatcherMessage> { ++ self.tx.send(msg).map_err(|err| err.0) + } + } + +-static DISPATCHER: LazyLock = LazyLock::new(Dispatcher::new); ++/// Install one completion's TLS state and invoke its callback on the current ++/// thread. Normally that thread is the dispatcher; this is also the fallback ++/// when dispatcher creation or channel delivery fails, preserving the ++/// exactly-once completion contract instead of silently dropping the message. ++fn deliver_message(msg: DispatcherMessage) { ++ match &msg.error { ++ Some((code, message)) => set_last_error(*code, message), ++ None => clear_last_error(), ++ } ++ ++ // Best-effort only (issue #61). A real `extern "C"` callback cannot ++ // unwind; a panic aborts at its own boundary before this catch runs. The ++ // catch only helps Rust hosts that deliberately supply a C-unwind shim. ++ let outcome = catch_unwind(AssertUnwindSafe(|| unsafe { ++ (msg.callback)(msg.callback_ctx, msg.status, msg.result); ++ })); ++ if let Err(payload) = outcome { ++ log::error!( ++ "lance-c dispatcher: unwinding host callback panicked; contained best-effort: {}", ++ panic_payload_message(&*payload) ++ ); ++ } ++} ++ ++fn dispatch_message(dispatcher: Option<&Dispatcher>, msg: DispatcherMessage) { ++ let undelivered = match dispatcher { ++ Some(dispatcher) => match dispatcher.send(msg) { ++ Ok(()) => return, ++ Err(msg) => msg, ++ }, ++ None => msg, ++ }; ++ log::error!("lance-c dispatcher unavailable; invoking async completion on the current thread"); ++ deliver_message(undelivered); ++} ++ ++static DISPATCHER: LazyLock> = LazyLock::new(|| match Dispatcher::new() { ++ Ok(dispatcher) => Some(dispatcher), ++ Err(err) => { ++ log::error!("failed to start lance-c dispatcher thread: {err}"); ++ None ++ } ++}); + + /// Send a completion message to the dispatcher thread. Before invoking the + /// callback, the dispatcher installs `error` on its own thread-local error +@@ -105,13 +119,16 @@ pub(crate) fn dispatch_callback( + result: *mut c_void, + error: Option<(LanceErrorCode, String)>, + ) { +- DISPATCHER.send(DispatcherMessage { +- callback, +- callback_ctx, +- status, +- result, +- error, +- }); ++ dispatch_message( ++ DISPATCHER.as_ref(), ++ DispatcherMessage { ++ callback, ++ callback_ctx, ++ status, ++ result, ++ error, ++ }, ++ ); + } + + #[cfg(test)] +@@ -234,4 +251,53 @@ mod tests { + + unsafe { reclaim(ctx) }; + } ++ ++ #[test] ++ fn unavailable_dispatcher_falls_back_without_dropping_completion() { ++ let (rx, ctx) = probe(); ++ dispatch_message( ++ None, ++ DispatcherMessage { ++ callback: observe, ++ callback_ctx: ctx, ++ status: -1, ++ result: ptr::null_mut(), ++ error: Some(( ++ LanceErrorCode::Internal, ++ "dispatcher unavailable".to_string(), ++ )), ++ }, ++ ); ++ ++ let obs = recv(&rx); ++ assert_eq!(obs.status, -1); ++ assert_eq!(obs.code, LanceErrorCode::Internal); ++ assert_eq!(obs.message.as_deref(), Some("dispatcher unavailable")); ++ unsafe { reclaim(ctx) }; ++ } ++ ++ #[test] ++ fn closed_dispatch_channel_falls_back_without_dropping_completion() { ++ let (tx, dead_rx) = mpsc::channel(); ++ drop(dead_rx); ++ let dispatcher = Dispatcher { tx }; ++ let (rx, ctx) = probe(); ++ ++ dispatch_message( ++ Some(&dispatcher), ++ DispatcherMessage { ++ callback: observe, ++ callback_ctx: ctx, ++ status: 0, ++ result: ptr::dangling_mut::(), ++ error: None, ++ }, ++ ); ++ ++ let obs = recv(&rx); ++ assert_eq!(obs.status, 0); ++ assert!(!obs.result_was_null); ++ assert_eq!(obs.code, LanceErrorCode::Ok); ++ unsafe { reclaim(ctx) }; ++ } + } +diff --git a/src/dataset.rs b/src/dataset.rs +index 9fe63e7..1397b74 100644 +--- a/src/dataset.rs ++++ b/src/dataset.rs +@@ -17,6 +17,7 @@ use lance_core::Result; + use crate::error::{ffi_try, swallow_unwind}; + use crate::helpers; + use crate::runtime::block_on; ++use crate::stream_guard::guarded_ffi_stream_from_reader; + + /// Opaque handle representing an opened Lance dataset. + pub struct LanceDataset { +@@ -151,7 +152,8 @@ unsafe fn open_dataset_inner( + } + + /// Close and free a dataset handle. +-/// Safe to call with NULL. Safe to call multiple times (subsequent calls are no-ops). ++/// Safe to call with NULL. A non-NULL handle must be closed exactly once and ++/// must not be used again afterwards. + /// + /// Best-effort (issue #61): a panic raised while dropping the handle is + /// caught and logged rather than unwinding into the caller, and the +@@ -268,6 +270,10 @@ unsafe fn dataset_schema_inner( + /// - `columns`: NULL-terminated column name array, or NULL for all columns + /// - `out`: pointer to a stack-allocated `ArrowArrayStream` + /// ++/// The already-materialized batch is exported through a guarded reader: ++/// schema conversion is validated before callbacks are exposed, and later ++/// `get_next` / `release` panics are contained at the Arrow C boundary. ++/// + /// Returns 0 on success, -1 on error. + #[unsafe(no_mangle)] + pub unsafe extern "C" fn lance_dataset_take( +@@ -307,7 +313,7 @@ unsafe fn dataset_take_inner( + // Wrap the single RecordBatch as a RecordBatchReader, then export as FFI stream. + let schema = batch.schema(); + let reader = arrow::record_batch::RecordBatchIterator::new(vec![Ok(batch)], schema); +- let ffi_stream = FFI_ArrowArrayStream::new(Box::new(reader)); ++ let ffi_stream = guarded_ffi_stream_from_reader(reader)?; + unsafe { + std::ptr::write_unaligned(out, ffi_stream); + } +@@ -326,6 +332,10 @@ unsafe fn dataset_take_inner( + /// to the same dataset snapshot used for this read. Missing or deleted row IDs + /// may be omitted from the result by the upstream Lance implementation. + /// ++/// The already-materialized batch is exported through the same guarded reader ++/// as [`lance_dataset_take`], including schema preflight and deferred callback ++/// panic containment. ++/// + /// Returns 0 on success, -1 on error. + #[unsafe(no_mangle)] + pub unsafe extern "C" fn lance_dataset_take_rows( +@@ -376,7 +386,7 @@ unsafe fn dataset_take_rows_inner( + // Match lance_dataset_take: export the single RecordBatch as an Arrow stream. + let schema = batch.schema(); + let reader = arrow::record_batch::RecordBatchIterator::new(vec![Ok(batch)], schema); +- let ffi_stream = FFI_ArrowArrayStream::new(Box::new(reader)); ++ let ffi_stream = guarded_ffi_stream_from_reader(reader)?; + unsafe { + std::ptr::write_unaligned(out, ffi_stream); + } +@@ -492,10 +502,19 @@ mod tests { + #[test] + fn with_mut_panic_rolls_back_and_handle_stays_usable() { + let (_tmp, handle) = create_test_handle(); ++ let (_replacement_tmp, replacement_handle) = create_test_handle(); ++ let replacement = Dataset::clone(&*replacement_handle.snapshot()); + let uri_before = handle.snapshot().uri().to_string(); ++ assert_ne!(replacement.uri(), uri_before); + + let result = catch_unwind(AssertUnwindSafe(|| { +- handle.with_mut(|_ds| panic!("simulated bug in mutation")) ++ handle.with_mut(|ds| { ++ // Make a visible in-memory mutation before panicking. This ++ // distinguishes clone-execute-swap from mutating the handle's ++ // stored Dataset in place and merely skipping the final swap. ++ *ds = replacement; ++ panic!("simulated bug in mutation") ++ }) + })); + let payload = result.expect_err("panic must escape with_mut unchanged"); + let msg = crate::error::panic_payload_message(&*payload); +diff --git a/src/error.rs b/src/error.rs +index f8158d9..0fe31c4 100644 +--- a/src/error.rs ++++ b/src/error.rs +@@ -88,6 +88,64 @@ pub fn set_lance_error(err: &lance_core::Error) { + set_last_error(error_code_from_lance(err), err.to_string()); + } + ++/// Why an [`ffi_guard_with`] invocation failed. ++pub(crate) enum FfiFailure { ++ /// The guarded body returned a regular `lance_core::Error`. ++ Lance, ++ /// Something panicked while executing the body or mapping its result. ++ Panic, ++} ++ ++/// Finish a caught FFI panic without leaving panic reporting unguarded. ++/// ++/// A failure while recording the panic or constructing the caller's error ++/// value is itself caught. There is no type-safe value we can manufacture if ++/// that recovery also panics, so the second payload is resumed; this is the ++/// documented double-panic limit of the FFI firewall. ++fn recover_from_ffi_panic( ++ payload: Box, ++ recover: impl FnOnce() -> T, ++) -> T { ++ match std::panic::catch_unwind(std::panic::AssertUnwindSafe(move || { ++ set_last_error( ++ LanceErrorCode::Panic, ++ format!("panic in FFI call: {}", panic_payload_message(&*payload)), ++ ); ++ recover() ++ })) { ++ Ok(value) => value, ++ Err(payload) => std::panic::resume_unwind(payload), ++ } ++} ++ ++/// Run a complete fallible FFI operation under the panic firewall and map any ++/// failure to the ABI-specific return value. ++/// ++/// The guard deliberately includes result mapping, not just `body()`: a ++/// wrapped external error may itself panic from `Display` while ++/// [`set_lance_error`] formats it. Keeping formatting, TLS mutation, and the ++/// error sentinel inside the unwind boundary prevents those secondary panics ++/// from escaping an `extern "C"` entry point. If the sentinel itself panics, ++/// the recovery path records `LanceErrorCode::Panic` and asks for it once more. ++pub(crate) fn ffi_guard_with( ++ body: impl FnOnce() -> lance_core::Result, ++ mut on_failure: impl FnMut(FfiFailure) -> T, ++) -> T { ++ match std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| match body() { ++ Ok(value) => { ++ clear_last_error(); ++ value ++ } ++ Err(err) => { ++ set_lance_error(&err); ++ on_failure(FfiFailure::Lance) ++ } ++ })) { ++ Ok(value) => value, ++ Err(payload) => recover_from_ffi_panic(payload, || on_failure(FfiFailure::Panic)), ++ } ++} ++ + /// Extract a human-readable message from a `catch_unwind` panic payload. + /// + /// `panic!` only ever produces `&str` or `String` payloads; anything else +@@ -182,89 +240,16 @@ pub unsafe extern "C" fn lance_free_string(s: *const c_char) { + /// captured by the `$errval:expr` catch-all. + macro_rules! ffi_try { + ($body:expr, null) => { +- match std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| $body)) { +- Ok(Ok(val)) => { +- $crate::error::clear_last_error(); +- val +- } +- Ok(Err(err)) => { +- $crate::error::set_lance_error(&err); +- std::ptr::null_mut() +- } +- Err(payload) => { +- $crate::error::set_last_error( +- $crate::error::LanceErrorCode::Panic, +- format!( +- "panic in FFI call: {}", +- $crate::error::panic_payload_message(&*payload) +- ), +- ); +- std::ptr::null_mut() +- } +- } ++ $crate::error::ffi_guard_with(|| $body, |_| std::ptr::null_mut()) + }; + ($body:expr, neg) => { +- match std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| $body)) { +- Ok(Ok(val)) => { +- $crate::error::clear_last_error(); +- val +- } +- Ok(Err(err)) => { +- $crate::error::set_lance_error(&err); +- -1 +- } +- Err(payload) => { +- $crate::error::set_last_error( +- $crate::error::LanceErrorCode::Panic, +- format!( +- "panic in FFI call: {}", +- $crate::error::panic_payload_message(&*payload) +- ), +- ); +- -1 +- } +- } ++ $crate::error::ffi_guard_with(|| $body, |_| -1) + }; + ($body:expr, void) => { +- match std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| $body)) { +- Ok(Ok(_)) => { +- $crate::error::clear_last_error(); +- } +- Ok(Err(err)) => { +- $crate::error::set_lance_error(&err); +- } +- Err(payload) => { +- $crate::error::set_last_error( +- $crate::error::LanceErrorCode::Panic, +- format!( +- "panic in FFI call: {}", +- $crate::error::panic_payload_message(&*payload) +- ), +- ); +- } +- } ++ $crate::error::ffi_guard_with(|| $body, |_| ()) + }; + ($body:expr, $errval:expr) => { +- match std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| $body)) { +- Ok(Ok(val)) => { +- $crate::error::clear_last_error(); +- val +- } +- Ok(Err(err)) => { +- $crate::error::set_lance_error(&err); +- $errval +- } +- Err(payload) => { +- $crate::error::set_last_error( +- $crate::error::LanceErrorCode::Panic, +- format!( +- "panic in FFI call: {}", +- $crate::error::panic_payload_message(&*payload) +- ), +- ); +- $errval +- } +- } ++ $crate::error::ffi_guard_with(|| $body, |_| $errval) + }; + } + +@@ -275,6 +260,17 @@ mod tests { + use super::*; + use std::ffi::CStr; + ++ #[derive(Debug)] ++ struct PanickingDisplay; ++ ++ impl std::fmt::Display for PanickingDisplay { ++ fn fmt(&self, _f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { ++ panic!("simulated panic while formatting an FFI error") ++ } ++ } ++ ++ impl std::error::Error for PanickingDisplay {} ++ + /// Yields a `lance_core::Result` by panicking — the panic is what the + /// `ffi_try!` shapes under test must catch. (The panic hook prints to + /// stderr during these tests; that is expected noise.) +@@ -402,6 +398,48 @@ mod tests { + assert!(msg.contains("bad arg"), "got: {msg}"); + } + ++ #[test] ++ fn ffi_try_catches_panic_while_formatting_lance_error() { ++ let v: u64 = ffi_try!( ++ Err(lance_core::Error::invalid_input_source(Box::new( ++ PanickingDisplay, ++ ))), ++ 0 ++ ); ++ assert_eq!(v, 0); ++ assert_eq!(lance_last_error_code(), LanceErrorCode::Panic); ++ let msg = take_last_error_message().expect("panic must set a message"); ++ assert!( ++ msg.contains("simulated panic while formatting an FFI error"), ++ "got: {msg}" ++ ); ++ } ++ ++ #[test] ++ fn ffi_try_catches_panic_while_building_error_sentinel() { ++ let attempts = std::cell::Cell::new(0); ++ let v: i64 = ffi_try!( ++ Err(lance_core::Error::invalid_input_source("bad arg".into())), ++ { ++ let attempt = attempts.get(); ++ attempts.set(attempt + 1); ++ if attempt == 0 { ++ panic!("simulated panic while building an FFI error sentinel"); ++ } ++ 7 ++ } ++ ); ++ ++ assert_eq!(v, 7); ++ assert_eq!(attempts.get(), 2); ++ assert_eq!(lance_last_error_code(), LanceErrorCode::Panic); ++ let msg = take_last_error_message().expect("panic must set a message"); ++ assert!( ++ msg.contains("simulated panic while building an FFI error sentinel"), ++ "got: {msg}" ++ ); ++ } ++ + #[test] + fn ffi_try_errval_maps_panic_to_errval_and_panic_code() { + // A non-zero sentinel proves the arm returns `$errval` verbatim. +diff --git a/src/fts_query.rs b/src/fts_query.rs +new file mode 100644 +index 0000000..cd194c7 +--- /dev/null ++++ b/src/fts_query.rs +@@ -0,0 +1,324 @@ ++// SPDX-License-Identifier: Apache-2.0 ++// SPDX-FileCopyrightText: Copyright The Lance Authors ++ ++//! Process-local, immutable FTS query context shared by segment-scoped scans. ++ ++use std::collections::HashSet; ++use std::ffi::c_char; ++use std::ptr; ++use std::sync::Arc; ++ ++use futures::future::try_join_all; ++use lance::index::{DatasetIndexExt, DatasetIndexInternalExt}; ++use lance_core::{Error, Result}; ++use lance_index::IndexCriteria; ++use lance_index::metrics::NoOpMetricsCollector; ++use lance_index::scalar::FullTextSearchQuery; ++use lance_index::scalar::inverted::query::{FtsQuery, collect_query_tokens}; ++use lance_index::scalar::inverted::{InvertedIndex, MemBM25Scorer, build_global_bm25_scorer}; ++use lance_table::format::IndexMetadata; ++use uuid::Uuid; ++ ++use crate::dataset::LanceDataset; ++use crate::error::{ffi_try, swallow_unwind}; ++use crate::helpers; ++use crate::runtime::block_on; ++ ++/// Required relationship between the pinned dataset snapshot and its FTS index. ++#[repr(i32)] ++#[derive(Clone, Copy, Debug, PartialEq, Eq)] ++pub enum LanceFtsCoverageMode { ++ /// Every current fragment must be covered by a committed FTS segment. ++ Strict = 0, ++ /// Search and score only documents covered by committed FTS segments. ++ IndexOnly = 1, ++} ++ ++impl TryFrom for LanceFtsCoverageMode { ++ type Error = Error; ++ ++ fn try_from(value: i32) -> Result { ++ match value { ++ 0 => Ok(Self::Strict), ++ 1 => Ok(Self::IndexOnly), ++ _ => Err(Error::invalid_input(format!( ++ "invalid coverage_mode {value}; expected 0 (STRICT) or 1 (INDEX_ONLY)" ++ ))), ++ } ++ } ++} ++ ++/// Rust-owned immutable state behind [`LanceFtsQueryContext`]. ++pub(crate) struct FtsQueryContextInner { ++ pub(crate) dataset: Arc, ++ pub(crate) query: FullTextSearchQuery, ++ pub(crate) segments: Vec, ++ pub(crate) scorer: Arc, ++} ++ ++impl FtsQueryContextInner { ++ pub(crate) fn validate_dataset_identity(&self, dataset: &Arc) -> Result<()> { ++ if !Arc::ptr_eq(&self.dataset, dataset) { ++ return Err(invalid_input(format!( ++ "FTS query context and scanner must originate from the same process-local dataset snapshot; context has uri '{}' version {}, scanner has uri '{}' version {}", ++ self.dataset.uri(), ++ self.dataset.version_id(), ++ dataset.uri(), ++ dataset.version_id() ++ ))); ++ } ++ Ok(()) ++ } ++} ++ ++/// Opaque process-local FTS query context. ++/// ++/// The handle owns an `Arc`, and scanners clone that `Arc` when the context is ++/// attached. It is therefore safe to close the public handle after all scanner ++/// attachments have completed. ++pub struct LanceFtsQueryContext { ++ pub(crate) inner: Arc, ++} ++ ++fn invalid_input(message: impl Into) -> Error { ++ Error::invalid_input(message.into()) ++} ++ ++async fn prepare_fts_query_context( ++ dataset: Arc, ++ column: String, ++ query_text: String, ++ coverage_mode: LanceFtsCoverageMode, ++) -> Result { ++ let logical_index = dataset ++ .load_scalar_index(IndexCriteria::default().for_column(&column).supports_fts()) ++ .await? ++ .ok_or_else(|| { ++ invalid_input(format!( ++ "no committed FTS index exists for column '{column}' in dataset version {}", ++ dataset.version_id() ++ )) ++ })?; ++ let segments = dataset.load_indices_by_name(&logical_index.name).await?; ++ if segments.is_empty() { ++ return Err(invalid_input(format!( ++ "FTS index for column '{column}' has no committed segments in dataset version {}", ++ dataset.version_id() ++ ))); ++ } ++ ++ let expected_fields = &segments[0].fields; ++ if let Some(segment) = segments ++ .iter() ++ .find(|segment| &segment.fields != expected_fields) ++ { ++ return Err(invalid_input(format!( ++ "FTS index '{}' has inconsistent fields across segments; segment {} has fields {:?}, expected {:?}", ++ logical_index.name, segment.uuid, segment.fields, expected_fields ++ ))); ++ } ++ ++ let current_fragment_ids: HashSet = dataset ++ .get_fragments() ++ .into_iter() ++ .map(|fragment| { ++ u32::try_from(fragment.id()).map_err(|_| { ++ invalid_input(format!( ++ "fragment id {} exceeds the u32 index metadata range", ++ fragment.id() ++ )) ++ }) ++ }) ++ .collect::>()?; ++ ++ let mut indexed_fragment_ids = HashSet::new(); ++ for segment in &segments { ++ let fragment_bitmap = segment.fragment_bitmap.as_ref().ok_or_else(|| { ++ invalid_input(format!( ++ "FTS segment {} for column '{column}' has unknown fragment coverage", ++ segment.uuid ++ )) ++ })?; ++ indexed_fragment_ids.extend( ++ fragment_bitmap ++ .iter() ++ .filter(|fragment_id| current_fragment_ids.contains(fragment_id)), ++ ); ++ } ++ let mut unindexed_fragment_ids: Vec = current_fragment_ids ++ .difference(&indexed_fragment_ids) ++ .copied() ++ .collect(); ++ unindexed_fragment_ids.sort_unstable(); ++ ++ if coverage_mode == LanceFtsCoverageMode::Strict && !unindexed_fragment_ids.is_empty() { ++ return Err(invalid_input(format!( ++ "coverage_mode=STRICT requires every fragment in dataset version {} to be indexed; column '{column}' has {} unindexed fragments: {:?}", ++ dataset.version_id(), ++ unindexed_fragment_ids.len(), ++ unindexed_fragment_ids ++ ))); ++ } ++ ++ let indices: Vec> = try_join_all(segments.iter().map(|segment| { ++ let dataset = Arc::clone(&dataset); ++ let column = column.clone(); ++ async move { ++ let index = dataset ++ .open_scalar_index(&column, &segment.uuid, &NoOpMetricsCollector) ++ .await?; ++ let inverted = index ++ .as_any() ++ .downcast_ref::() ++ .ok_or_else(|| { ++ invalid_input(format!( ++ "index segment {} for column '{column}' is not an inverted index", ++ segment.uuid ++ )) ++ })?; ++ Ok::<_, Error>(Arc::new(inverted.clone())) ++ } ++ })) ++ .await?; ++ ++ let expected_params = indices[0].params(); ++ if let Some((position, _)) = indices ++ .iter() ++ .enumerate() ++ .find(|(_, index)| index.params() != expected_params) ++ { ++ return Err(invalid_input(format!( ++ "FTS index '{}' has inconsistent inverted index parameters; segment {} differs from segment {}", ++ logical_index.name, segments[position].uuid, segments[0].uuid ++ ))); ++ } ++ ++ let query = FullTextSearchQuery::new(query_text).with_column(column.clone())?; ++ let match_query = match &query.query { ++ FtsQuery::Match(query) => query, ++ _ => { ++ return Err(Error::internal( ++ "prepared FTS query unexpectedly produced a non-Match query".to_string(), ++ )); ++ } ++ }; ++ let mut tokenizer = indices[0].tokenizer(); ++ let query_tokens = collect_query_tokens(&match_query.terms, &mut tokenizer); ++ let params = query ++ .params() ++ .with_fuzziness(match_query.fuzziness) ++ .with_max_expansions(match_query.max_expansions) ++ .with_prefix_length(match_query.prefix_length); ++ let scorer = Arc::new(build_global_bm25_scorer(&indices, &query_tokens, ¶ms).await?); ++ ++ Ok(FtsQueryContextInner { ++ dataset, ++ query, ++ segments, ++ scorer, ++ }) ++} ++ ++/// Prepare a process-local global BM25 scorer and the committed segment list ++/// for one single-column Match query against the dataset's pinned snapshot. ++#[unsafe(no_mangle)] ++pub unsafe extern "C" fn lance_dataset_prepare_fts_query( ++ dataset: *const LanceDataset, ++ column: *const c_char, ++ query: *const c_char, ++ max_fuzzy_distance: u32, ++ coverage_mode: i32, ++) -> *mut LanceFtsQueryContext { ++ ffi_try!( ++ unsafe { ++ prepare_fts_query_inner(dataset, column, query, max_fuzzy_distance, coverage_mode) ++ }, ++ null ++ ) ++} ++ ++unsafe fn prepare_fts_query_inner( ++ dataset: *const LanceDataset, ++ column: *const c_char, ++ query: *const c_char, ++ max_fuzzy_distance: u32, ++ coverage_mode: i32, ++) -> Result<*mut LanceFtsQueryContext> { ++ if dataset.is_null() || column.is_null() || query.is_null() { ++ return Err(invalid_input("dataset, column, and query must not be NULL")); ++ } ++ let column = unsafe { helpers::parse_c_string(column)? } ++ .filter(|value| !value.is_empty()) ++ .ok_or_else(|| invalid_input("column must not be empty"))? ++ .to_string(); ++ let query = unsafe { helpers::parse_c_string(query)? } ++ .filter(|value| !value.is_empty()) ++ .ok_or_else(|| invalid_input("query must not be empty"))? ++ .to_string(); ++ let coverage_mode = LanceFtsCoverageMode::try_from(coverage_mode)?; ++ if max_fuzzy_distance != 0 { ++ return Err(invalid_input(format!( ++ "max_fuzzy_distance must be 0 for prepared FTS query contexts, got {max_fuzzy_distance}; fuzzy queries require a canonical prepared BM25 vocabulary" ++ ))); ++ } ++ let snapshot = unsafe { &*dataset }.snapshot(); ++ let inner = block_on(prepare_fts_query_context( ++ snapshot, ++ column, ++ query, ++ coverage_mode, ++ ))?; ++ Ok(Box::into_raw(Box::new(LanceFtsQueryContext { ++ inner: Arc::new(inner), ++ }))) ++} ++ ++/// Close a context handle. NULL-safe. Scanners that already attached the ++/// context retain their own shared reference. ++#[unsafe(no_mangle)] ++pub unsafe extern "C" fn lance_fts_query_context_close(context: *mut LanceFtsQueryContext) { ++ if !context.is_null() { ++ swallow_unwind("lance_fts_query_context_close", || unsafe { ++ drop(Box::from_raw(context)); ++ }); ++ } ++} ++ ++pub(crate) unsafe fn clone_context( ++ context: *const LanceFtsQueryContext, ++) -> Result> { ++ if context.is_null() { ++ return Err(invalid_input("context must not be NULL")); ++ } ++ Ok(Arc::clone(&unsafe { &*context }.inner)) ++} ++ ++pub(crate) fn parse_segment_uuids(segment_uuids: *const u8, len: usize) -> Result> { ++ if segment_uuids.is_null() && len > 0 { ++ return Err(invalid_input( ++ "segment_uuids is NULL but len is greater than 0", ++ )); ++ } ++ if len > isize::MAX as usize / 16 { ++ return Err(invalid_input(format!( ++ "segment UUID count {len} exceeds the maximum addressable byte slice length" ++ ))); ++ } ++ let mut uuids = Vec::with_capacity(len); ++ for position in 0..len { ++ let mut bytes = [0_u8; 16]; ++ unsafe { ++ ptr::copy_nonoverlapping(segment_uuids.add(position * 16), bytes.as_mut_ptr(), 16); ++ } ++ uuids.push(Uuid::from_bytes(bytes)); ++ } ++ let unique: HashSet = uuids.iter().copied().collect(); ++ if unique.len() != uuids.len() { ++ return Err(invalid_input(format!( ++ "segment_uuids contains duplicate UUIDs; len={}, unique={}", ++ uuids.len(), ++ unique.len() ++ ))); ++ } ++ Ok(uuids) ++} +diff --git a/src/index_segment.rs b/src/index_segment.rs +index 9ffc8d0..a46c4f3 100644 +--- a/src/index_segment.rs ++++ b/src/index_segment.rs +@@ -22,7 +22,7 @@ use prost::Message; + use uuid::Uuid; + + use crate::dataset::LanceDataset; +-use crate::error::{LanceErrorCode, clear_last_error, ffi_try, set_last_error}; ++use crate::error::{ffi_try, swallow_unwind}; + use crate::helpers; + use crate::index::{ + LanceMetricType, LanceScalarIndexType, LanceVectorIndexParams, LanceVectorIndexType, +@@ -1042,7 +1042,9 @@ pub unsafe extern "C" fn lance_free_bytes(bytes: *mut u8) { + #[unsafe(no_mangle)] + pub unsafe extern "C" fn lance_index_segment_builder_free(builder: *mut LanceIndexSegmentBuilder) { + if !builder.is_null() { +- unsafe { drop(Box::from_raw(builder)) }; ++ swallow_unwind("lance_index_segment_builder_free", || unsafe { ++ drop(Box::from_raw(builder)); ++ }); + } + } + +@@ -1155,12 +1157,15 @@ unsafe fn metadata_uuid_inner( + pub unsafe extern "C" fn lance_index_segment_metadata_name( + metadata: *const LanceIndexSegmentMetadata, + ) -> *const c_char { +- if metadata.is_null() { +- set_last_error(LanceErrorCode::InvalidArgument, "metadata is NULL"); +- return ptr::null(); +- } +- clear_last_error(); +- unsafe { (*metadata).name.as_ptr() } ++ ffi_try!( ++ (|| -> Result<*const c_char> { ++ if metadata.is_null() { ++ return Err(invalid_input("metadata is NULL")); ++ } ++ Ok(unsafe { (*metadata).name.as_ptr() }) ++ })(), ++ ptr::null() ++ ) + } + + /// Return the dataset version recorded in the metadata. +@@ -1168,12 +1173,15 @@ pub unsafe extern "C" fn lance_index_segment_metadata_name( + pub unsafe extern "C" fn lance_index_segment_metadata_dataset_version( + metadata: *const LanceIndexSegmentMetadata, + ) -> u64 { +- if metadata.is_null() { +- set_last_error(LanceErrorCode::InvalidArgument, "metadata is NULL"); +- return 0; +- } +- clear_last_error(); +- unsafe { (*metadata).metadata.dataset_version } ++ ffi_try!( ++ (|| -> Result { ++ if metadata.is_null() { ++ return Err(invalid_input("metadata is NULL")); ++ } ++ Ok(unsafe { (*metadata).metadata.dataset_version }) ++ })(), ++ 0 ++ ) + } + + /// Return the physical index version recorded in the metadata. +@@ -1181,12 +1189,15 @@ pub unsafe extern "C" fn lance_index_segment_metadata_dataset_version( + pub unsafe extern "C" fn lance_index_segment_metadata_index_version( + metadata: *const LanceIndexSegmentMetadata, + ) -> i32 { +- if metadata.is_null() { +- set_last_error(LanceErrorCode::InvalidArgument, "metadata is NULL"); +- return -1; +- } +- clear_last_error(); +- unsafe { (*metadata).metadata.index_version } ++ ffi_try!( ++ (|| -> Result { ++ if metadata.is_null() { ++ return Err(invalid_input("metadata is NULL")); ++ } ++ Ok(unsafe { (*metadata).metadata.index_version }) ++ })(), ++ neg ++ ) + } + + /// Return the concrete scalar/vector index enum value, or -1 on error. +@@ -1194,16 +1205,7 @@ pub unsafe extern "C" fn lance_index_segment_metadata_index_version( + pub unsafe extern "C" fn lance_index_segment_metadata_index_type( + metadata: *const LanceIndexSegmentMetadata, + ) -> i32 { +- match unsafe { metadata_index_type_inner(metadata) } { +- Ok(index_type) => { +- clear_last_error(); +- index_type +- } +- Err(error) => { +- crate::error::set_lance_error(&error); +- -1 +- } +- } ++ ffi_try!(unsafe { metadata_index_type_inner(metadata) }, neg) + } + + unsafe fn metadata_index_type_inner(metadata: *const LanceIndexSegmentMetadata) -> Result { +@@ -1258,19 +1260,20 @@ unsafe fn metadata_index_type_inner(metadata: *const LanceIndexSegmentMetadata) + pub unsafe extern "C" fn lance_index_segment_metadata_index_details_type_url( + metadata: *const LanceIndexSegmentMetadata, + ) -> *const c_char { +- if metadata.is_null() { +- set_last_error(LanceErrorCode::InvalidArgument, "metadata is NULL"); +- return ptr::null(); +- } +- let Some(type_url) = (unsafe { &(*metadata).index_details_type_url }) else { +- set_last_error( +- LanceErrorCode::NotFound, +- "index metadata does not contain index_details", +- ); +- return ptr::null(); +- }; +- clear_last_error(); +- type_url.as_ptr() ++ ffi_try!( ++ (|| -> Result<*const c_char> { ++ if metadata.is_null() { ++ return Err(invalid_input("metadata is NULL")); ++ } ++ let type_url = unsafe { &(*metadata).index_details_type_url } ++ .as_ref() ++ .ok_or_else(|| { ++ Error::index_not_found("index metadata does not contain index_details") ++ })?; ++ Ok(type_url.as_ptr()) ++ })(), ++ ptr::null() ++ ) + } + + /// Return the number of indexed field IDs. +@@ -1278,12 +1281,15 @@ pub unsafe extern "C" fn lance_index_segment_metadata_index_details_type_url( + pub unsafe extern "C" fn lance_index_segment_metadata_field_count( + metadata: *const LanceIndexSegmentMetadata, + ) -> usize { +- if metadata.is_null() { +- set_last_error(LanceErrorCode::InvalidArgument, "metadata is NULL"); +- return 0; +- } +- clear_last_error(); +- unsafe { (*metadata).metadata.fields.len() } ++ ffi_try!( ++ (|| -> Result { ++ if metadata.is_null() { ++ return Err(invalid_input("metadata is NULL")); ++ } ++ Ok(unsafe { (*metadata).metadata.fields.len() }) ++ })(), ++ 0 ++ ) + } + + /// Copy indexed field IDs in metadata order. +@@ -1334,12 +1340,15 @@ unsafe fn metadata_field_ids_inner( + pub unsafe extern "C" fn lance_index_segment_metadata_fragment_count( + metadata: *const LanceIndexSegmentMetadata, + ) -> usize { +- if metadata.is_null() { +- set_last_error(LanceErrorCode::InvalidArgument, "metadata is NULL"); +- return 0; +- } +- clear_last_error(); +- unsafe { (*metadata).fragment_ids.len() } ++ ffi_try!( ++ (|| -> Result { ++ if metadata.is_null() { ++ return Err(invalid_input("metadata is NULL")); ++ } ++ Ok(unsafe { (*metadata).fragment_ids.len() }) ++ })(), ++ 0 ++ ) + } + + /// Copy covered fragment IDs in ascending order. +@@ -1393,6 +1402,8 @@ pub unsafe extern "C" fn lance_index_segment_metadata_free( + metadata: *mut LanceIndexSegmentMetadata, + ) { + if !metadata.is_null() { +- unsafe { drop(Box::from_raw(metadata)) }; ++ swallow_unwind("lance_index_segment_metadata_free", || unsafe { ++ drop(Box::from_raw(metadata)); ++ }); + } + } +diff --git a/src/lib.rs b/src/lib.rs +index ed9cfe1..4d54641 100644 +--- a/src/lib.rs ++++ b/src/lib.rs +@@ -15,6 +15,11 @@ + //! - The caller is responsible for freeing returned strings with `lance_free_string()`. + #![allow(clippy::missing_safety_doc)] + ++#[cfg(not(panic = "unwind"))] ++compile_error!( ++ "lance-c requires panic=\"unwind\" so its C ABI panic firewall can honor LANCE_ERR_PANIC" ++); ++ + mod add_columns; + mod alter_columns; + mod async_dispatcher; +@@ -26,6 +31,7 @@ mod delete; + mod drop_columns; + mod error; + mod fragment_writer; ++mod fts_query; + mod helpers; + mod index; + mod index_model; +@@ -52,6 +58,7 @@ pub use error::{ + LanceErrorCode, lance_free_string, lance_last_error_code, lance_last_error_message, + }; + pub use fragment_writer::*; ++pub use fts_query::*; + pub use index::*; + pub use index_model::*; + pub use index_segment::*; +diff --git a/src/scanner.rs b/src/scanner.rs +index ef9d290..7110111 100644 +--- a/src/scanner.rs ++++ b/src/scanner.rs +@@ -6,28 +6,37 @@ + use std::ffi::{c_char, c_void}; + use std::pin::Pin; + use std::ptr; +-use std::sync::Arc; + use std::sync::atomic::{AtomicBool, Ordering}; ++use std::sync::{Arc, Condvar, Mutex, Weak}; + use std::task::{Context, Poll, RawWaker, RawWakerVTable, Waker}; + + use arrow::ffi_stream::FFI_ArrowArrayStream; +-use arrow_schema::SchemaRef; ++use arrow_schema::{Schema as ArrowSchema, SchemaRef}; ++use datafusion::physical_plan::ExecutionPlan; + use futures::{FutureExt, Stream, StreamExt}; + use lance::Dataset; + use lance::dataset::scanner::{ + DatasetRecordBatchStream, ExecutionStatsCallback, ExecutionSummaryCounts, + }; ++use lance::io::exec::fts::MatchQueryExec; + use lance_core::Result; ++use lance_datafusion::exec::{LanceExecutionOptions, get_session_context}; ++use lance_datafusion::planner::Planner; ++use lance_datafusion::substrait::parse_substrait; + use lance_index::scalar::FullTextSearchQuery; + use lance_io::stream::RecordBatchStream; ++use lance_table::format::IndexMetadata; + use uuid::Uuid; + + use crate::async_dispatcher::{self, LanceCallback}; + use crate::batch::LanceBatch; + use crate::dataset::LanceDataset; + use crate::error::{ +- LanceErrorCode, clear_last_error, error_code_from_lance, ffi_try, panic_payload_message, +- set_lance_error, set_last_error, swallow_unwind, ++ FfiFailure, LanceErrorCode, clear_last_error, error_code_from_lance, ffi_guard_with, ffi_try, ++ panic_payload_message, set_lance_error, set_last_error, swallow_unwind, ++}; ++use crate::fts_query::{ ++ FtsQueryContextInner, LanceFtsQueryContext, clone_context, parse_segment_uuids, + }; + use crate::helpers; + use crate::runtime::{RT, block_on}; +@@ -50,6 +59,7 @@ pub struct LanceScanner { + columns: Option>, + filter: Option, + substrait_filter: Option>, ++ additional_sql_filters: Vec, + limit: Option, + offset: Option, + batch_size: Option, +@@ -64,13 +74,20 @@ pub struct LanceScanner { + use_index: Option, + prefilter: bool, + fts_query: Option, +- // Set when a panic is caught in a stateful stream operation (issue #61): ++ fts_context: Option>, ++ fts_index_segments: Option>, ++ // Set when a panic is caught in any operation on this scanner (issue #61): + // once poisoned, every later `lance_scanner_*` call on this handle (except + // `lance_scanner_close`, which must always free memory) fails with + // `LANCE_ERR_PANIC`. Behind an `Arc` so the exported-stream wrapper and + // the spawned async task can poison the handle from outside this call + // frame via `poison_flag()`. + poisoned: Arc, ++ // Every RawWaker handed to the poll stream registers here. Close retires ++ // the registry before dropping the stream: pending callbacks are ++ // cancelled and callbacks already in progress are allowed to quiesce ++ // before the caller may destroy callback_ctx. ++ poll_wakers: PollWakerRegistry, + scan_statistics_callback: Option, + scan_started: AtomicBool, + // Materialized on first iteration call +@@ -111,6 +128,7 @@ impl LanceScanner { + columns: None, + filter: None, + substrait_filter: None, ++ additional_sql_filters: Vec::new(), + limit: None, + offset: None, + batch_size: None, +@@ -125,7 +143,10 @@ impl LanceScanner { + use_index: None, + prefilter: false, + fts_query: None, ++ fts_context: None, ++ fts_index_segments: None, + poisoned: Arc::new(AtomicBool::new(false)), ++ poll_wakers: PollWakerRegistry::default(), + scan_statistics_callback: None, + scan_started: AtomicBool::new(false), + stream: None, +@@ -161,86 +182,58 @@ impl LanceScanner { + Ok(()) + } + +- /// Build the underlying Scanner and open a stream. +- fn materialize_stream(&mut self) -> Result<()> { +- self.scan_started.store(true, Ordering::Release); +- let mut scanner = self.dataset.scan(); +- if let Some(cols) = &self.columns { +- scanner.project(cols)?; +- } +- // Substrait filter takes precedence over SQL filter when both are set. +- if let Some(bytes) = &self.substrait_filter { +- scanner.filter_substrait(bytes)?; +- } else if let Some(filter) = &self.filter { +- scanner.filter(filter)?; +- } +- if self.limit.is_some() || self.offset.is_some() { +- scanner.limit(self.limit, self.offset)?; +- } +- if let Some(bs) = self.batch_size { +- scanner.batch_size(bs); +- } +- if self.with_row_id { +- scanner.with_row_id(); +- } +- self.apply_fragment_filter(&mut scanner)?; +- if self.index_segments.is_some() && self.nearest.is_none() { +- return Err(lance_core::Error::invalid_input_source( +- "index_segments requires nearest() to be configured".into(), +- )); +- } +- // Lance validates fragment-scoped nearest searches when nearest() is +- // configured. Such searches are supported when the fragment scan is +- // the input to a prefilter, so this flag must be set first. +- if self.prefilter { +- scanner.prefilter(true); +- } +- if let Some(n) = &self.nearest { +- scanner.nearest(&n.column, n.query.as_ref(), n.k as usize)?; +- if let Some(np) = self.nprobes { +- scanner.nprobes(np as usize); +- } +- if let Some(rf) = self.refine_factor { +- scanner.refine(rf); +- } +- if let Some(ef) = self.ef { +- scanner.ef(ef as usize); +- } +- if let Some(m) = self.metric_override { +- scanner.distance_metric(m.to_distance()); +- } +- if let Some(ui) = self.use_index { +- scanner.use_index(ui); +- } +- if let Some(segments) = &self.index_segments { +- scanner.with_index_segments(segments.clone())?; ++ fn apply_filter(&self, scanner: &mut lance::dataset::scanner::Scanner) -> Result<()> { ++ if self.additional_sql_filters.is_empty() { ++ if let Some(substrait) = &self.substrait_filter { ++ scanner.filter_substrait(substrait)?; ++ } else if let Some(sql) = &self.filter { ++ scanner.filter(sql)?; + } ++ return Ok(()); + } +- if let Some(fts) = &self.fts_query { +- scanner.full_text_search(fts.clone())?; +- } +- if let Some(callback) = &self.scan_statistics_callback { +- scanner.scan_stats_callback(callback.clone()); ++ ++ let schema = Arc::new(ArrowSchema::from(self.dataset.schema())); ++ let planner = Planner::new(Arc::clone(&schema)); ++ let mut combined = if let Some(substrait) = &self.substrait_filter { ++ let context = get_session_context(&LanceExecutionOptions::default()); ++ Some( ++ parse_substrait(substrait, schema, &context.state()) ++ .now_or_never() ++ .expect("Substrait filter parsing must complete synchronously")?, ++ ) ++ } else if let Some(sql) = &self.filter { ++ Some(planner.parse_filter(sql)?) ++ } else { ++ None ++ }; ++ for sql in &self.additional_sql_filters { ++ let sql = planner.parse_filter(sql)?; ++ combined = Some(match combined { ++ Some(existing) => existing.and(sql), ++ None => sql, ++ }); + } +- let stream = block_on(scanner.try_into_stream())?; ++ scanner.filter_expr(planner.optimize_expr(combined.expect("additional filter exists"))?); ++ Ok(()) ++ } ++ ++ /// Build the underlying Scanner and open a stream. ++ fn materialize_stream(&mut self) -> Result<()> { ++ let prepared_scanner = self.build_scanner()?; ++ let stream = block_on(prepared_scanner.try_into_stream())?; + self.schema = Some(stream.schema()); + self.stream = Some(Box::pin(stream)); + Ok(()) + } + + /// Build a Scanner (without materializing) and return it. +- fn build_scanner(&self) -> Result { ++ fn build_scanner(&self) -> Result { + self.scan_started.store(true, Ordering::Release); + let mut scanner = self.dataset.scan(); + if let Some(cols) = &self.columns { + scanner.project(cols)?; + } +- // Substrait filter takes precedence over SQL filter when both are set. +- if let Some(bytes) = &self.substrait_filter { +- scanner.filter_substrait(bytes)?; +- } else if let Some(filter) = &self.filter { +- scanner.filter(filter)?; +- } ++ self.apply_filter(&mut scanner)?; + if self.limit.is_some() || self.offset.is_some() { + scanner.limit(self.limit, self.offset)?; + } +@@ -256,6 +249,16 @@ impl LanceScanner { + "index_segments requires nearest() to be configured".into(), + )); + } ++ if self.fts_index_segments.is_some() && self.fts_context.is_none() { ++ return Err(lance_core::Error::invalid_input_source( ++ "fts_index_segments requires an FTS query context".into(), ++ )); ++ } ++ if self.fts_context.is_some() && self.fragment_ids.is_some() { ++ return Err(lance_core::Error::invalid_input_source( ++ "fragment_ids cannot be combined with an FTS query context; split the query by FTS index segment UUID instead".into(), ++ )); ++ } + // nearest() checks the current prefilter setting before accepting a + // fragment-scoped search. Enable it before installing the query. + if self.prefilter { +@@ -285,11 +288,141 @@ impl LanceScanner { + if let Some(fts) = &self.fts_query { + scanner.full_text_search(fts.clone())?; + } ++ let distributed_fts = if let Some(context) = &self.fts_context { ++ context.validate_dataset_identity(&self.dataset)?; ++ let segments = select_fts_segments(context, self.fts_index_segments.as_deref())?; ++ scanner.full_text_search(context.query.clone())?; ++ // Both STRICT and INDEX_ONLY context scans must use only the ++ // committed segments pinned in the context. In STRICT mode all ++ // current fragments were already proven covered during prepare. ++ scanner.fast_search(); ++ Some(PreparedFtsExecution { ++ context: Arc::clone(context), ++ segments, ++ batch_size: self.batch_size, ++ scan_statistics_callback: self.scan_statistics_callback.clone(), ++ }) ++ } else { ++ None ++ }; + if let Some(callback) = &self.scan_statistics_callback { + scanner.scan_stats_callback(callback.clone()); + } +- Ok(scanner) ++ Ok(PreparedScanner { ++ scanner, ++ distributed_fts, ++ }) ++ } ++} ++ ++struct PreparedFtsExecution { ++ context: Arc, ++ segments: Vec, ++ batch_size: Option, ++ scan_statistics_callback: Option, ++} ++ ++struct PreparedScanner { ++ scanner: lance::dataset::scanner::Scanner, ++ distributed_fts: Option, ++} ++ ++impl PreparedScanner { ++ async fn try_into_stream(self) -> Result { ++ let Some(distributed_fts) = self.distributed_fts else { ++ return self.scanner.try_into_stream().await; ++ }; ++ let plan = self.scanner.create_plan().await?; ++ let (plan, replaced) = replace_match_query_exec( ++ plan, ++ &distributed_fts.segments, ++ &distributed_fts.context.scorer, ++ )?; ++ if replaced != 1 { ++ return Err(lance_core::Error::internal(format!( ++ "expected exactly one MatchQueryExec in prepared FTS plan, replaced {replaced}" ++ ))); ++ } ++ let stream = lance_datafusion::exec::execute_plan( ++ plan, ++ lance_datafusion::exec::LanceExecutionOptions { ++ batch_size: distributed_fts.batch_size, ++ execution_stats_callback: distributed_fts.scan_statistics_callback, ++ ..Default::default() ++ }, ++ )?; ++ Ok(DatasetRecordBatchStream::new(stream)) ++ } ++} ++ ++fn select_fts_segments( ++ context: &FtsQueryContextInner, ++ selected_uuids: Option<&[Uuid]>, ++) -> Result> { ++ let Some(selected_uuids) = selected_uuids else { ++ return Ok(context.segments.clone()); ++ }; ++ let mut selected = Vec::with_capacity(selected_uuids.len()); ++ for uuid in selected_uuids { ++ let segment = context ++ .segments ++ .iter() ++ .find(|segment| segment.uuid == *uuid) ++ .ok_or_else(|| { ++ lance_core::Error::invalid_input_source( ++ format!( ++ "FTS segment UUID {uuid} is not present in the attached query context for dataset version {}", ++ context.dataset.version_id() ++ ) ++ .into(), ++ ) ++ })?; ++ selected.push(segment.clone()); ++ } ++ if selected.is_empty() { ++ return Err(lance_core::Error::invalid_input_source( ++ "FTS segment subset must contain at least one UUID".into(), ++ )); ++ } ++ Ok(selected) ++} ++ ++fn replace_match_query_exec( ++ plan: Arc, ++ segments: &[IndexMetadata], ++ scorer: &Arc, ++) -> Result<(Arc, usize)> { ++ let children = plan.children(); ++ let mut replaced = 0; ++ let rebuilt = if children.is_empty() { ++ plan ++ } else { ++ let mut new_children = Vec::with_capacity(children.len()); ++ for child in children { ++ let (new_child, child_replaced) = ++ replace_match_query_exec(Arc::clone(child), segments, scorer)?; ++ new_children.push(new_child); ++ replaced += child_replaced; ++ } ++ plan.with_new_children(new_children).map_err(|error| { ++ lance_core::Error::internal(format!( ++ "failed to rebuild FTS execution plan children: {error}" ++ )) ++ })? ++ }; ++ ++ if let Some(exec) = rebuilt.downcast_ref::() { ++ let replacement = MatchQueryExec::new_with_segments( ++ Arc::clone(exec.dataset()), ++ exec.query().clone(), ++ exec.params().clone(), ++ exec.prefilter_source().clone(), ++ segments.to_vec(), ++ ) ++ .with_base_scorer(Arc::clone(scorer)); ++ return Ok((Arc::new(replacement), replaced + 1)); + } ++ Ok((rebuilt, replaced)) + } + + /// Type of a dynamically named scan metric. +@@ -442,6 +575,26 @@ macro_rules! scanner_poison_check { + }; + } + ++/// Run a scanner configuration call through the common FFI guard and poison ++/// the handle if that call catches a panic. Regular Lance errors remain ++/// recoverable and do not poison the builder. ++macro_rules! scanner_ffi_try { ++ ($scanner:expr, $body:expr $(,)?) => {{ ++ let scanner_ptr = $scanner; ++ ffi_guard_with( ++ || $body, ++ |failure| { ++ if matches!(failure, FfiFailure::Panic) && !scanner_ptr.is_null() { ++ unsafe { &*scanner_ptr } ++ .poison_flag() ++ .store(true, Ordering::SeqCst); ++ } ++ -1 ++ }, ++ ) ++ }}; ++} ++ + // --------------------------------------------------------------------------- + // Scanner lifecycle + builder + // --------------------------------------------------------------------------- +@@ -486,7 +639,7 @@ unsafe fn scanner_new_inner( + #[unsafe(no_mangle)] + pub unsafe extern "C" fn lance_scanner_set_limit(scanner: *mut LanceScanner, limit: i64) -> i32 { + scanner_poison_check!(scanner, -1); +- ffi_try!(unsafe { scanner_set_limit_inner(scanner, limit) }, neg) ++ scanner_ffi_try!(scanner, unsafe { scanner_set_limit_inner(scanner, limit) }) + } + + unsafe fn scanner_set_limit_inner(scanner: *mut LanceScanner, limit: i64) -> Result { +@@ -504,7 +657,9 @@ unsafe fn scanner_set_limit_inner(scanner: *mut LanceScanner, limit: i64) -> Res + #[unsafe(no_mangle)] + pub unsafe extern "C" fn lance_scanner_set_offset(scanner: *mut LanceScanner, offset: i64) -> i32 { + scanner_poison_check!(scanner, -1); +- ffi_try!(unsafe { scanner_set_offset_inner(scanner, offset) }, neg) ++ scanner_ffi_try!(scanner, unsafe { ++ scanner_set_offset_inner(scanner, offset) ++ }) + } + + unsafe fn scanner_set_offset_inner(scanner: *mut LanceScanner, offset: i64) -> Result { +@@ -525,10 +680,9 @@ pub unsafe extern "C" fn lance_scanner_set_batch_size( + batch_size: i64, + ) -> i32 { + scanner_poison_check!(scanner, -1); +- ffi_try!( +- unsafe { scanner_set_batch_size_inner(scanner, batch_size) }, +- neg +- ) ++ scanner_ffi_try!(scanner, unsafe { ++ scanner_set_batch_size_inner(scanner, batch_size) ++ }) + } + + unsafe fn scanner_set_batch_size_inner(scanner: *mut LanceScanner, batch_size: i64) -> Result { +@@ -549,7 +703,9 @@ pub unsafe extern "C" fn lance_scanner_with_row_id( + enable: bool, + ) -> i32 { + scanner_poison_check!(scanner, -1); +- ffi_try!(unsafe { scanner_with_row_id_inner(scanner, enable) }, neg) ++ scanner_ffi_try!(scanner, unsafe { ++ scanner_with_row_id_inner(scanner, enable) ++ }) + } + + unsafe fn scanner_with_row_id_inner(scanner: *mut LanceScanner, enable: bool) -> Result { +@@ -574,10 +730,9 @@ pub unsafe extern "C" fn lance_scanner_set_fragment_ids( + len: usize, + ) -> i32 { + scanner_poison_check!(scanner, -1); +- ffi_try!( +- unsafe { scanner_set_fragment_ids_inner(scanner, ids, len) }, +- neg +- ) ++ scanner_ffi_try!(scanner, unsafe { ++ scanner_set_fragment_ids_inner(scanner, ids, len) ++ }) + } + + unsafe fn scanner_set_fragment_ids_inner( +@@ -631,10 +786,9 @@ pub unsafe extern "C" fn lance_scanner_set_substrait_filter( + len: usize, + ) -> i32 { + scanner_poison_check!(scanner, -1); +- ffi_try!( +- unsafe { scanner_set_substrait_filter_inner(scanner, bytes, len) }, +- neg +- ) ++ scanner_ffi_try!(scanner, unsafe { ++ scanner_set_substrait_filter_inner(scanner, bytes, len) ++ }) + } + + unsafe fn scanner_set_substrait_filter_inner( +@@ -663,6 +817,50 @@ unsafe fn scanner_set_substrait_filter_inner( + Ok(0) + } + ++/// Add an SQL filter that is combined with the scanner's selected primary filter using AND. ++/// ++/// The primary filter is the Substrait filter when one is set, otherwise it is the SQL filter ++/// passed to `lance_scanner_new`. Multiple additional SQL filters are also combined using AND. ++/// This must be called before the scan starts. The string is copied into the scanner. ++/// ++/// Returns 0 on success, -1 on error (check `lance_last_error_*`). ++#[unsafe(no_mangle)] ++pub unsafe extern "C" fn lance_scanner_additional_sql_filter( ++ scanner: *mut LanceScanner, ++ filter: *const c_char, ++) -> i32 { ++ scanner_poison_check!(scanner, -1); ++ scanner_ffi_try!(scanner, unsafe { ++ scanner_additional_sql_filter_inner(scanner, filter) ++ }) ++} ++ ++unsafe fn scanner_additional_sql_filter_inner( ++ scanner: *mut LanceScanner, ++ filter: *const c_char, ++) -> Result { ++ if scanner.is_null() { ++ return Err(lance_core::Error::invalid_input_source( ++ "scanner is NULL".into(), ++ )); ++ } ++ let filter = unsafe { helpers::parse_c_string(filter)? } ++ .ok_or_else(|| lance_core::Error::invalid_input_source("filter must not be NULL".into()))?; ++ if filter.is_empty() { ++ return Err(lance_core::Error::invalid_input_source( ++ "additional SQL filter must be non-empty".into(), ++ )); ++ } ++ let scanner = unsafe { &mut *scanner }; ++ if scanner.scan_started.load(Ordering::Acquire) { ++ return Err(lance_core::Error::invalid_input_source( ++ "additional SQL filter must be set before the scan starts".into(), ++ )); ++ } ++ scanner.additional_sql_filters.push(filter.to_string()); ++ Ok(0) ++} ++ + /// Register a callback that receives execution statistics after the scan stream + /// is fully consumed to EOF. + /// +@@ -693,10 +891,9 @@ pub unsafe extern "C" fn lance_scanner_set_statistics_callback( + callback_ctx: *mut c_void, + ) -> i32 { + scanner_poison_check!(scanner, -1); +- ffi_try!( +- unsafe { scanner_set_statistics_callback_inner(scanner, callback, callback_ctx) }, +- neg +- ) ++ scanner_ffi_try!(scanner, unsafe { ++ scanner_set_statistics_callback_inner(scanner, callback, callback_ctx) ++ }) + } + + unsafe fn scanner_set_statistics_callback_inner( +@@ -734,6 +931,12 @@ unsafe fn scanner_set_statistics_callback_inner( + + /// Close and free a scanner handle. + /// ++/// Pending poll wakers are cancelled before the stream is dropped. If a poll ++/// waker callback is already running on another thread, close waits for that ++/// callback to return, making this function the retirement boundary for its ++/// callback context. A waker callback must therefore never close or otherwise ++/// re-enter its originating scanner. ++/// + /// Best-effort (issue #61): this drops a possibly-live + /// `DatasetRecordBatchStream`, the highest-risk `Drop` in this crate. A + /// panic raised while dropping the handle is caught and logged rather than +@@ -743,7 +946,9 @@ unsafe fn scanner_set_statistics_callback_inner( + pub unsafe extern "C" fn lance_scanner_close(scanner: *mut LanceScanner) { + if !scanner.is_null() { + swallow_unwind("lance_scanner_close", || unsafe { +- let _ = Box::from_raw(scanner); ++ let scanner = Box::from_raw(scanner); ++ scanner.poll_wakers.retire_and_wait(); ++ drop(scanner); + }); + } + } +@@ -770,27 +975,30 @@ pub unsafe extern "C" fn lance_scanner_to_arrow_stream( + scanner: *mut LanceScanner, + out: *mut FFI_ArrowArrayStream, + ) -> i32 { +- if scanner.is_null() || out.is_null() { +- set_last_error( +- LanceErrorCode::InvalidArgument, +- "scanner and out must not be NULL", +- ); ++ if scanner.is_null() { ++ set_last_error(LanceErrorCode::InvalidArgument, "scanner must not be NULL"); + return -1; + } + scanner_poison_check!(scanner, -1); ++ if out.is_null() { ++ set_last_error(LanceErrorCode::InvalidArgument, "out must not be NULL"); ++ return -1; ++ } + let s = unsafe { &*scanner }; + let poisoned = s.poison_flag(); +- match std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| unsafe { +- scanner_to_arrow_stream_inner(s, out) +- })) { +- Ok(Ok(rc)) => { +- clear_last_error(); +- rc +- } +- Ok(Err(err)) => { +- set_lance_error(&err); +- -1 ++ match std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| { ++ match unsafe { scanner_to_arrow_stream_inner(s, out) } { ++ Ok(rc) => { ++ clear_last_error(); ++ rc ++ } ++ Err(err) => { ++ set_lance_error(&err); ++ -1 ++ } + } ++ })) { ++ Ok(rc) => rc, + Err(payload) => { + poisoned.store(true, Ordering::SeqCst); + set_last_error( +@@ -848,14 +1056,18 @@ pub unsafe extern "C" fn lance_scanner_next( + scanner: *mut LanceScanner, + out: *mut *mut LanceBatch, + ) -> i32 { +- if scanner.is_null() || out.is_null() { +- set_last_error( +- LanceErrorCode::InvalidArgument, +- "scanner and out must not be NULL", +- ); ++ if !out.is_null() { ++ unsafe { *out = ptr::null_mut() }; ++ } ++ if scanner.is_null() { ++ set_last_error(LanceErrorCode::InvalidArgument, "scanner must not be NULL"); + return -1; + } + scanner_poison_check!(scanner, -1); ++ if out.is_null() { ++ set_last_error(LanceErrorCode::InvalidArgument, "out must not be NULL"); ++ return -1; ++ } + let s = unsafe { &mut *scanner }; + let poisoned = s.poison_flag(); + match std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| unsafe { +@@ -919,16 +1131,22 @@ unsafe fn scanner_next_inner(s: &mut LanceScanner, out: *mut *mut LanceBatch) -> + /// Start an async scan. The callback is invoked on a dedicated dispatcher thread + /// when the ArrowArrayStream is ready. + /// +-/// - `callback`: Called with `(ctx, 0, *mut ArrowArrayStream)` on success, +-/// or `(ctx, -1, NULL)` on error. On error, the dispatcher installs the +-/// error on the callback thread's TLS first, so `lance_last_error_*` +-/// called from inside the callback observes the failure. ++/// - `callback`: Must not be NULL. Called with ++/// `(ctx, 0, *mut ArrowArrayStream)` on success or `(ctx, -1, NULL)` on ++/// error. The successful result is a Rust-allocated outer stream container ++/// and must eventually be passed to [`lance_scanner_async_stream_free`]. On ++/// error, the dispatcher installs the error on the callback thread's TLS ++/// first, so `lance_last_error_*` called from inside the callback observes ++/// the failure. + /// - `callback_ctx`: Opaque pointer passed back to the callback. + /// + /// The scanner configuration is captured at call time. The scanner handle + /// can be closed immediately after this call. + /// +-/// The promised contract is exactly one callback completion, even on panic. ++/// With a non-NULL callback, the promised contract is exactly one completion, ++/// even on panic. Completions normally run on the dispatcher thread; if that ++/// thread cannot be created or its channel has failed, delivery falls back to ++/// the thread producing the completion rather than dropping it. + /// A panic anywhere in call-time setup (validation, scanner building, + /// runtime access, task spawn) is caught by the entry guard below and still + /// reported through the callback: `(ctx, -1, NULL)` with `LANCE_ERR_PANIC`, +@@ -939,9 +1157,13 @@ unsafe fn scanner_next_inner(s: &mut LanceScanner, out: *mut *mut LanceBatch) -> + #[unsafe(no_mangle)] + pub unsafe extern "C" fn lance_scanner_scan_async( + scanner: *const LanceScanner, +- callback: LanceCallback, ++ callback: Option, + callback_ctx: *mut c_void, + ) { ++ let Some(callback) = callback else { ++ set_last_error(LanceErrorCode::InvalidArgument, "callback must not be NULL"); ++ return; ++ }; + unsafe { + scan_async_guarded(scanner, callback, callback_ctx, |s, cb, ctx| { + scan_async_setup(s, cb, ctx) +@@ -1148,6 +1370,22 @@ unsafe fn scan_async_setup( + }); + } + ++/// Release the heap-allocated Arrow stream container returned through a ++/// successful [`lance_scanner_scan_async`] callback. ++/// ++/// The Arrow stream's own `release` callback is invoked first when it is still ++/// present, then the outer Rust allocation is freed. Passing NULL is a no-op. ++/// This function must not be used for stack-allocated streams returned by ++/// [`lance_scanner_to_arrow_stream`]. ++#[unsafe(no_mangle)] ++pub unsafe extern "C" fn lance_scanner_async_stream_free(stream: *mut FFI_ArrowArrayStream) { ++ if !stream.is_null() { ++ swallow_unwind("lance_scanner_async_stream_free", || unsafe { ++ drop(Box::from_raw(stream)); ++ }); ++ } ++} ++ + // --------------------------------------------------------------------------- + // Poll-based iteration (for cooperative async runtimes) + // --------------------------------------------------------------------------- +@@ -1155,9 +1393,13 @@ unsafe fn scan_async_setup( + /// Poll for the next batch without blocking. + /// + /// - If data is already buffered, returns `LANCE_POLL_READY` immediately. +-/// - If I/O is needed, returns `LANCE_POLL_PENDING` and schedules the waker callback. ++/// - If I/O is needed, returns `LANCE_POLL_PENDING` and schedules the non-NULL ++/// waker callback. + /// The caller should yield the thread and re-poll after the waker fires. + /// - The waker is single-use: it fires at most once per poll call that returns PENDING. ++/// Its context must remain valid until the callback returns or ++/// `lance_scanner_close` returns. Close cancels callbacks that have not ++/// entered and waits for callbacks already in progress. + /// + /// The stream is lazily materialized on the first poll call (which will typically + /// return PENDING while the stream opens). +@@ -1168,18 +1410,26 @@ unsafe fn scan_async_setup( + #[unsafe(no_mangle)] + pub unsafe extern "C" fn lance_scanner_poll_next( + scanner: *mut LanceScanner, +- waker: LanceWaker, ++ waker: Option, + waker_ctx: *mut c_void, + out: *mut *mut LanceBatch, + ) -> LancePollStatus { +- if scanner.is_null() || out.is_null() { +- set_last_error( +- LanceErrorCode::InvalidArgument, +- "scanner and out must not be NULL", +- ); ++ if !out.is_null() { ++ unsafe { *out = ptr::null_mut() }; ++ } ++ if scanner.is_null() { ++ set_last_error(LanceErrorCode::InvalidArgument, "scanner must not be NULL"); + return LancePollStatus::Error; + } + scanner_poison_check!(scanner, LancePollStatus::Error); ++ if out.is_null() { ++ set_last_error(LanceErrorCode::InvalidArgument, "out must not be NULL"); ++ return LancePollStatus::Error; ++ } ++ let Some(waker) = waker else { ++ set_last_error(LanceErrorCode::InvalidArgument, "waker must not be NULL"); ++ return LancePollStatus::Error; ++ }; + let s = unsafe { &mut *scanner }; + let poisoned = s.poison_flag(); + match std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| unsafe { +@@ -1220,13 +1470,13 @@ unsafe fn scanner_poll_next_inner( + return LancePollStatus::Error; + } + +- let stream = s.stream.as_mut().unwrap(); +- + // Construct a std::task::Waker from the C function pointer. +- let raw_waker = make_raw_waker(waker, waker_ctx); ++ let raw_waker = make_raw_waker(&s.poll_wakers, waker, waker_ctx); + let waker_obj = unsafe { Waker::from_raw(raw_waker) }; + let mut cx = Context::from_waker(&waker_obj); + ++ let stream = s.stream.as_mut().unwrap(); ++ + // Enter the Tokio runtime context so internal I/O futures can access + // the reactor. Without this, polling from a non-Tokio thread panics. + let _guard = RT.enter(); +@@ -1264,39 +1514,66 @@ unsafe fn scanner_poll_next_inner( + struct CWakerContext { + waker_fn: LanceWaker, + ctx: *mut c_void, ++ state: Mutex, ++ quiesced: Condvar, ++} ++ ++#[derive(Default)] ++struct CWakerState { ++ fired: bool, ++ cancelled: bool, ++ active: bool, ++} ++ ++#[derive(Default)] ++struct PollWakerRegistry { ++ state: Mutex, ++} ++ ++#[derive(Default)] ++struct PollWakerRegistryState { ++ retired: bool, ++ registrations: Vec>, + } + + // C function pointers + void* are Send by convention for FFI. + unsafe impl Send for CWakerContext {} + unsafe impl Sync for CWakerContext {} + +-fn make_raw_waker(waker_fn: LanceWaker, ctx: *mut c_void) -> RawWaker { +- let data = Box::into_raw(Box::new(CWakerContext { waker_fn, ctx })) as *const (); ++fn make_raw_waker( ++ registry: &PollWakerRegistry, ++ waker_fn: LanceWaker, ++ ctx: *mut c_void, ++) -> RawWaker { ++ let context = Arc::new(CWakerContext { ++ waker_fn, ++ ctx, ++ state: Mutex::new(CWakerState::default()), ++ quiesced: Condvar::new(), ++ }); ++ registry.register(&context); ++ let data = Arc::into_raw(context) as *const (); + + const VTABLE: RawWakerVTable = RawWakerVTable::new( + // clone + |data| { +- let orig = unsafe { &*(data as *const CWakerContext) }; +- let cloned = Box::new(CWakerContext { +- waker_fn: orig.waker_fn, +- ctx: orig.ctx, +- }); +- RawWaker::new(Box::into_raw(cloned) as *const (), &VTABLE) ++ unsafe { Arc::::increment_strong_count(data.cast()) }; ++ RawWaker::new(data, &VTABLE) + }, + // wake (consumes) + |data| { +- let ctx = unsafe { Box::from_raw(data as *mut CWakerContext) }; +- unsafe { (ctx.waker_fn)(ctx.ctx) }; ++ let ctx = unsafe { Arc::from_raw(data as *const CWakerContext) }; ++ ctx.wake_once(); + }, + // wake_by_ref + |data| { + let ctx = unsafe { &*(data as *const CWakerContext) }; +- unsafe { (ctx.waker_fn)(ctx.ctx) }; ++ ctx.wake_once(); + }, + // drop + |data| { + unsafe { +- let _ = Box::from_raw(data as *mut CWakerContext); ++ drop(Arc::from_raw(data as *const CWakerContext)); + }; + }, + ); +@@ -1304,6 +1581,99 @@ fn make_raw_waker(waker_fn: LanceWaker, ctx: *mut c_void) -> RawWaker { + RawWaker::new(data, &VTABLE) + } + ++impl CWakerContext { ++ fn wake_once(&self) { ++ { ++ let mut state = self ++ .state ++ .lock() ++ .unwrap_or_else(|poisoned| poisoned.into_inner()); ++ if state.cancelled || state.fired { ++ return; ++ } ++ state.fired = true; ++ state.active = true; ++ } ++ ++ unsafe { (self.waker_fn)(self.ctx) }; ++ ++ let mut state = self ++ .state ++ .lock() ++ .unwrap_or_else(|poisoned| poisoned.into_inner()); ++ state.active = false; ++ self.quiesced.notify_all(); ++ } ++ ++ fn cancel(&self) { ++ self.state ++ .lock() ++ .unwrap_or_else(|poisoned| poisoned.into_inner()) ++ .cancelled = true; ++ } ++ ++ fn wait_until_quiescent(&self) { ++ let mut state = self ++ .state ++ .lock() ++ .unwrap_or_else(|poisoned| poisoned.into_inner()); ++ while state.active { ++ state = self ++ .quiesced ++ .wait(state) ++ .unwrap_or_else(|poisoned| poisoned.into_inner()); ++ } ++ } ++} ++ ++impl PollWakerRegistry { ++ fn register(&self, registration: &Arc) { ++ let retired = { ++ let mut state = self ++ .state ++ .lock() ++ .unwrap_or_else(|poisoned| poisoned.into_inner()); ++ state ++ .registrations ++ .retain(|candidate| candidate.strong_count() > 0); ++ if state.retired { ++ true ++ } else { ++ state.registrations.push(Arc::downgrade(registration)); ++ false ++ } ++ }; ++ if retired { ++ registration.cancel(); ++ } ++ } ++ ++ fn retire_and_wait(&self) { ++ let registrations = { ++ let mut state = self ++ .state ++ .lock() ++ .unwrap_or_else(|poisoned| poisoned.into_inner()); ++ state.retired = true; ++ state ++ .registrations ++ .drain(..) ++ .filter_map(|registration| registration.upgrade()) ++ .collect::>() ++ }; ++ ++ // Cancel every registration before waiting for any one callback, so ++ // no later registration can enter while close is quiescing an earlier ++ // one. ++ for registration in ®istrations { ++ registration.cancel(); ++ } ++ for registration in registrations { ++ registration.wait_until_quiescent(); ++ } ++ } ++} ++ + // --------------------------------------------------------------------------- + // Vector search (Phase 2): setter knobs + // --------------------------------------------------------------------------- +@@ -1313,7 +1683,8 @@ macro_rules! scanner_set_u32 { + #[unsafe(no_mangle)] + pub unsafe extern "C" fn $name(scanner: *mut LanceScanner, value: u32) -> i32 { + scanner_poison_check!(scanner, -1); +- ffi_try!( ++ scanner_ffi_try!( ++ scanner, + (|| -> Result { + if scanner.is_null() { + return Err(lance_core::Error::invalid_input_source( +@@ -1324,8 +1695,7 @@ macro_rules! scanner_set_u32 { + (*scanner).$field = Some(value); + } + Ok(0) +- })(), +- neg ++ })() + ) + } + }; +@@ -1338,7 +1708,9 @@ scanner_set_u32!(lance_scanner_set_ef, ef); + #[unsafe(no_mangle)] + pub unsafe extern "C" fn lance_scanner_set_metric(scanner: *mut LanceScanner, metric: i32) -> i32 { + scanner_poison_check!(scanner, -1); +- ffi_try!(unsafe { scanner_set_metric_inner(scanner, metric) }, neg) ++ scanner_ffi_try!(scanner, unsafe { ++ scanner_set_metric_inner(scanner, metric) ++ }) + } + + unsafe fn scanner_set_metric_inner(scanner: *mut LanceScanner, metric: i32) -> Result { +@@ -1370,7 +1742,9 @@ pub unsafe extern "C" fn lance_scanner_set_use_index( + enable: bool, + ) -> i32 { + scanner_poison_check!(scanner, -1); +- ffi_try!(unsafe { scanner_set_use_index_inner(scanner, enable) }, neg) ++ scanner_ffi_try!(scanner, unsafe { ++ scanner_set_use_index_inner(scanner, enable) ++ }) + } + + unsafe fn scanner_set_use_index_inner(scanner: *mut LanceScanner, enable: bool) -> Result { +@@ -1391,7 +1765,9 @@ pub unsafe extern "C" fn lance_scanner_set_prefilter( + enable: bool, + ) -> i32 { + scanner_poison_check!(scanner, -1); +- ffi_try!(unsafe { scanner_set_prefilter_inner(scanner, enable) }, neg) ++ scanner_ffi_try!(scanner, unsafe { ++ scanner_set_prefilter_inner(scanner, enable) ++ }) + } + + unsafe fn scanner_set_prefilter_inner(scanner: *mut LanceScanner, enable: bool) -> Result { +@@ -1424,10 +1800,9 @@ pub unsafe extern "C" fn lance_scanner_set_index_segments( + len: usize, + ) -> i32 { + scanner_poison_check!(scanner, -1); +- ffi_try!( +- unsafe { scanner_set_index_segments_inner(scanner, segment_uuids, len) }, +- neg +- ) ++ scanner_ffi_try!(scanner, unsafe { ++ scanner_set_index_segments_inner(scanner, segment_uuids, len) ++ }) + } + + unsafe fn scanner_set_index_segments_inner( +@@ -1485,10 +1860,9 @@ pub unsafe extern "C" fn lance_scanner_nearest( + k: u32, + ) -> i32 { + scanner_poison_check!(scanner, -1); +- ffi_try!( +- unsafe { scanner_nearest_inner(scanner, column, query_data, query_len, element_type, k) }, +- neg +- ) ++ scanner_ffi_try!(scanner, unsafe { ++ scanner_nearest_inner(scanner, column, query_data, query_len, element_type, k) ++ },) + } + + unsafe fn scanner_nearest_inner( +@@ -1510,9 +1884,9 @@ unsafe fn scanner_nearest_inner( + )); + } + let s = unsafe { &mut *scanner }; +- if s.fts_query.is_some() { ++ if s.fts_query.is_some() || s.fts_context.is_some() { + return Err(lance_core::Error::invalid_input_source( +- "cannot call nearest after full_text_search; they are mutually exclusive".into(), ++ "cannot call nearest after full_text_search or attaching an FTS query context; they are mutually exclusive".into(), + )); + } + let column_str = unsafe { helpers::parse_c_string(column)? }.unwrap(); +@@ -1586,10 +1960,9 @@ pub unsafe extern "C" fn lance_scanner_full_text_search( + max_fuzzy_distance: u32, + ) -> i32 { + scanner_poison_check!(scanner, -1); +- ffi_try!( +- unsafe { fts_inner(scanner, query, columns, max_fuzzy_distance) }, +- neg +- ) ++ scanner_ffi_try!(scanner, unsafe { ++ fts_inner(scanner, query, columns, max_fuzzy_distance) ++ },) + } + + unsafe fn fts_inner( +@@ -1611,6 +1984,11 @@ unsafe fn fts_inner( + "cannot call full_text_search after nearest; they are mutually exclusive".into(), + )); + } ++ if s.fts_context.is_some() { ++ return Err(lance_core::Error::invalid_input_source( ++ "cannot call full_text_search after attaching an FTS query context; the context already owns the query".into(), ++ )); ++ } + + let query_str = unsafe { helpers::parse_c_string(query)? } + .unwrap() +@@ -1633,13 +2011,89 @@ unsafe fn fts_inner( + Ok(0) + } + ++/// Attach an immutable, process-local FTS query context to this scanner. ++/// The scanner clones the context's shared ownership; the caller may close ++/// the public context handle after this function returns successfully. ++#[unsafe(no_mangle)] ++pub unsafe extern "C" fn lance_scanner_set_fts_query_context( ++ scanner: *mut LanceScanner, ++ context: *const LanceFtsQueryContext, ++) -> i32 { ++ scanner_poison_check!(scanner, -1); ++ scanner_ffi_try!(scanner, unsafe { ++ scanner_set_fts_query_context_inner(scanner, context) ++ }) ++} ++ ++unsafe fn scanner_set_fts_query_context_inner( ++ scanner: *mut LanceScanner, ++ context: *const LanceFtsQueryContext, ++) -> Result { ++ if scanner.is_null() { ++ return Err(lance_core::Error::invalid_input_source( ++ "scanner must not be NULL".into(), ++ )); ++ } ++ let context = unsafe { clone_context(context)? }; ++ let scanner = unsafe { &mut *scanner }; ++ if scanner.nearest.is_some() { ++ return Err(lance_core::Error::invalid_input_source( ++ "cannot attach an FTS query context after nearest; they are mutually exclusive".into(), ++ )); ++ } ++ if scanner.fts_query.is_some() { ++ return Err(lance_core::Error::invalid_input_source( ++ "cannot attach an FTS query context after full_text_search; the context already owns the query" ++ .into(), ++ )); ++ } ++ context.validate_dataset_identity(&scanner.dataset)?; ++ scanner.fts_context = Some(context); ++ Ok(0) ++} ++ ++/// Restrict a context-backed FTS scan to a subset of segment UUIDs. ++/// Passing `len == 0` clears the restriction so all context segments are used. ++#[unsafe(no_mangle)] ++pub unsafe extern "C" fn lance_scanner_set_fts_index_segments( ++ scanner: *mut LanceScanner, ++ segment_uuids: *const u8, ++ len: usize, ++) -> i32 { ++ scanner_poison_check!(scanner, -1); ++ scanner_ffi_try!(scanner, unsafe { ++ scanner_set_fts_index_segments_inner(scanner, segment_uuids, len) ++ }) ++} ++ ++unsafe fn scanner_set_fts_index_segments_inner( ++ scanner: *mut LanceScanner, ++ segment_uuids: *const u8, ++ len: usize, ++) -> Result { ++ if scanner.is_null() { ++ return Err(lance_core::Error::invalid_input_source( ++ "scanner must not be NULL".into(), ++ )); ++ } ++ let segments = if len == 0 { ++ None ++ } else { ++ Some(parse_segment_uuids(segment_uuids, len)?) ++ }; ++ unsafe { &mut *scanner }.fts_index_segments = segments; ++ Ok(0) ++} ++ + #[cfg(test)] + mod tests { + use super::*; + use crate::dataset::{lance_dataset_close, lance_dataset_open}; + use crate::error::{lance_last_error_code, lance_last_error_message}; + use std::ffi::{CStr, CString}; +- use std::sync::atomic::AtomicI32; ++ use std::sync::atomic::{AtomicI32, AtomicUsize}; ++ use std::sync::{Barrier, mpsc}; ++ use std::time::Duration; + + use arrow_array::{Int32Array, RecordBatch, StringArray}; + use arrow_schema::{DataType, Field, Schema}; +@@ -1706,6 +2160,198 @@ mod tests { + + unsafe extern "C" fn noop_waker(_ctx: *mut c_void) {} + ++ #[test] ++ fn null_async_callback_is_rejected_without_poisoning_scanner() { ++ let (_tmp, uri) = create_test_dataset(); ++ let (dataset, scanner) = open_dataset_and_scanner(&uri); ++ ++ unsafe { lance_scanner_scan_async(scanner, None, ptr::null_mut()) }; ++ assert_eq!(lance_last_error_code(), LanceErrorCode::InvalidArgument); ++ let msg_ptr = lance_last_error_message(); ++ assert!(!msg_ptr.is_null()); ++ let msg = unsafe { CStr::from_ptr(msg_ptr) }.to_string_lossy(); ++ assert!(msg.contains("callback must not be NULL"), "got: {msg}"); ++ unsafe { crate::error::lance_free_string(msg_ptr) }; ++ assert!(!unsafe { &*scanner }.is_poisoned()); ++ ++ unsafe { ++ lance_scanner_close(scanner); ++ lance_dataset_close(dataset); ++ } ++ } ++ ++ #[test] ++ fn null_poll_waker_is_rejected_and_clears_out() { ++ let (_tmp, uri) = create_test_dataset(); ++ let (dataset, scanner) = open_dataset_and_scanner(&uri); ++ let mut batch = std::ptr::NonNull::::dangling().as_ptr(); ++ ++ let status = unsafe { lance_scanner_poll_next(scanner, None, ptr::null_mut(), &mut batch) }; ++ assert_eq!(status, LancePollStatus::Error); ++ assert!(batch.is_null()); ++ assert_eq!(lance_last_error_code(), LanceErrorCode::InvalidArgument); ++ assert!(!unsafe { &*scanner }.is_poisoned()); ++ ++ unsafe { ++ lance_scanner_close(scanner); ++ lance_dataset_close(dataset); ++ } ++ } ++ ++ #[test] ++ fn raw_waker_clones_share_one_shot_gate() { ++ static WAKES: AtomicUsize = AtomicUsize::new(0); ++ unsafe extern "C" fn count_wake(_ctx: *mut c_void) { ++ WAKES.fetch_add(1, Ordering::SeqCst); ++ } ++ ++ WAKES.store(0, Ordering::SeqCst); ++ let registry = PollWakerRegistry::default(); ++ let waker = ++ unsafe { Waker::from_raw(make_raw_waker(®istry, count_wake, ptr::null_mut())) }; ++ let cloned = waker.clone(); ++ waker.wake_by_ref(); ++ cloned.wake_by_ref(); ++ drop(cloned); ++ drop(waker); ++ assert_eq!(WAKES.load(Ordering::SeqCst), 1); ++ } ++ ++ #[test] ++ fn scanner_close_cancels_a_retained_poll_waker() { ++ let (_tmp, uri) = create_test_dataset(); ++ let (dataset, scanner) = open_dataset_and_scanner(&uri); ++ let calls = Box::into_raw(Box::new(AtomicUsize::new(0))); ++ ++ unsafe extern "C" fn count_wake(ctx: *mut c_void) { ++ let calls = unsafe { &*(ctx.cast::()) }; ++ calls.fetch_add(1, Ordering::SeqCst); ++ } ++ ++ // Model a future retaining the RawWaker clone returned from a PENDING ++ // poll. Closing the scanner is the documented retirement boundary, so ++ // waking that retained clone afterwards must not touch callback_ctx. ++ let waker = unsafe { ++ Waker::from_raw(make_raw_waker( ++ &(*scanner).poll_wakers, ++ count_wake, ++ calls.cast(), ++ )) ++ }; ++ unsafe { lance_scanner_close(scanner) }; ++ waker.wake(); ++ ++ let calls = unsafe { Box::from_raw(calls) }; ++ assert_eq!( ++ calls.load(Ordering::SeqCst), ++ 0, ++ "a retained RawWaker invoked callback_ctx after scanner close" ++ ); ++ unsafe { lance_dataset_close(dataset) }; ++ } ++ ++ struct BlockingWakeProbe { ++ calls: AtomicUsize, ++ entered: Arc, ++ release: Arc, ++ } ++ ++ unsafe extern "C" fn blocking_waker(ctx: *mut c_void) { ++ let probe = unsafe { &*(ctx.cast::()) }; ++ probe.calls.fetch_add(1, Ordering::SeqCst); ++ probe.entered.wait(); ++ probe.release.wait(); ++ } ++ ++ #[test] ++ fn scanner_close_waits_for_an_active_poll_waker() { ++ let (_tmp, uri) = create_test_dataset(); ++ let (dataset, scanner) = open_dataset_and_scanner(&uri); ++ let entered = Arc::new(Barrier::new(2)); ++ let release = Arc::new(Barrier::new(2)); ++ let probe = Box::into_raw(Box::new(BlockingWakeProbe { ++ calls: AtomicUsize::new(0), ++ entered: Arc::clone(&entered), ++ release: Arc::clone(&release), ++ })); ++ let waker = unsafe { ++ Waker::from_raw(make_raw_waker( ++ &(*scanner).poll_wakers, ++ blocking_waker, ++ probe.cast(), ++ )) ++ }; ++ ++ let wake_thread = std::thread::spawn(move || waker.wake()); ++ entered.wait(); ++ ++ let close_started = Arc::new(Barrier::new(2)); ++ let close_started_in_thread = Arc::clone(&close_started); ++ let (closed_tx, closed_rx) = mpsc::channel(); ++ let scanner_address = scanner as usize; ++ let close_thread = std::thread::spawn(move || { ++ close_started_in_thread.wait(); ++ unsafe { lance_scanner_close(scanner_address as *mut LanceScanner) }; ++ closed_tx.send(()).unwrap(); ++ }); ++ close_started.wait(); ++ ++ let closed_while_callback_was_active = ++ closed_rx.recv_timeout(Duration::from_millis(500)).is_ok(); ++ release.wait(); ++ wake_thread.join().unwrap(); ++ close_thread.join().unwrap(); ++ ++ let probe = unsafe { Box::from_raw(probe) }; ++ assert_eq!(probe.calls.load(Ordering::SeqCst), 1); ++ assert!( ++ !closed_while_callback_was_active, ++ "scanner close returned before an active poll waker callback completed" ++ ); ++ unsafe { lance_dataset_close(dataset) }; ++ } ++ ++ fn panicking_setter_body() -> Result { ++ panic!("simulated panic in scanner setter") ++ } ++ ++ unsafe fn panicking_scanner_setter(scanner: *mut LanceScanner) -> i32 { ++ scanner_poison_check!(scanner, -1); ++ scanner_ffi_try!(scanner, panicking_setter_body()) ++ } ++ ++ #[test] ++ fn scanner_setter_panic_poisons_handle() { ++ let (_tmp, uri) = create_test_dataset(); ++ let (dataset, scanner) = open_dataset_and_scanner(&uri); ++ ++ let rc = unsafe { panicking_scanner_setter(scanner) }; ++ assert_eq!(rc, -1); ++ assert!(unsafe { &*scanner }.is_poisoned()); ++ assert_eq!(lance_last_error_code(), LanceErrorCode::Panic); ++ let msg_ptr = lance_last_error_message(); ++ assert!(!msg_ptr.is_null()); ++ let msg = unsafe { CStr::from_ptr(msg_ptr) } ++ .to_string_lossy() ++ .into_owned(); ++ unsafe { crate::error::lance_free_string(msg_ptr) }; ++ assert!( ++ msg.contains("simulated panic in scanner setter"), ++ "got: {msg}" ++ ); ++ ++ // The original panic message is reported once; later calls use the ++ // stable poison error and never touch scanner state again. ++ let rc = unsafe { lance_scanner_set_limit(scanner, 10) }; ++ assert_eq!(rc, -1); ++ assert_poison_error_pending(); ++ ++ unsafe { ++ lance_scanner_close(scanner); ++ lance_dataset_close(dataset); ++ } ++ } ++ + #[test] + fn poisoned_scanner_rejects_setters_with_panic_code() { + let (_tmp, uri) = create_test_dataset(); +@@ -1745,12 +2391,17 @@ mod tests { + let (dataset, scanner) = open_dataset_and_scanner(&uri); + poison(scanner); + +- let mut batch: *mut LanceBatch = ptr::null_mut(); ++ let mut batch = std::ptr::NonNull::::dangling().as_ptr(); + let rc = unsafe { lance_scanner_next(scanner, &mut batch) }; + assert_eq!(rc, -1); + assert!(batch.is_null(), "error path must leave *out NULL"); + assert_poison_error_pending(); + ++ // Poison has precedence over validation of secondary arguments. ++ let rc = unsafe { lance_scanner_next(scanner, ptr::null_mut()) }; ++ assert_eq!(rc, -1); ++ assert_poison_error_pending(); ++ + unsafe { + lance_scanner_close(scanner); + lance_dataset_close(dataset); +@@ -1763,13 +2414,20 @@ mod tests { + let (dataset, scanner) = open_dataset_and_scanner(&uri); + poison(scanner); + +- let mut batch: *mut LanceBatch = ptr::null_mut(); +- let status = +- unsafe { lance_scanner_poll_next(scanner, noop_waker, ptr::null_mut(), &mut batch) }; ++ let mut batch = std::ptr::NonNull::::dangling().as_ptr(); ++ let status = unsafe { ++ lance_scanner_poll_next(scanner, Some(noop_waker), ptr::null_mut(), &mut batch) ++ }; + assert_eq!(status, LancePollStatus::Error); + assert!(batch.is_null(), "error path must leave *out NULL"); + assert_poison_error_pending(); + ++ let status = unsafe { ++ lance_scanner_poll_next(scanner, Some(noop_waker), ptr::null_mut(), ptr::null_mut()) ++ }; ++ assert_eq!(status, LancePollStatus::Error); ++ assert_poison_error_pending(); ++ + unsafe { + lance_scanner_close(scanner); + lance_dataset_close(dataset); +@@ -1809,7 +2467,7 @@ mod tests { + let (dataset, scanner) = open_dataset_and_scanner(&uri); + poison(scanner); + +- unsafe { lance_scanner_scan_async(scanner, record_status, ptr::null_mut()) }; ++ unsafe { lance_scanner_scan_async(scanner, Some(record_status), ptr::null_mut()) }; + // The poison error is also visible on the calling thread. + assert_poison_error_pending(); + +diff --git a/src/stream_guard.rs b/src/stream_guard.rs +index f4f418e..d785082 100644 +--- a/src/stream_guard.rs ++++ b/src/stream_guard.rs +@@ -43,6 +43,8 @@ use std::panic::{AssertUnwindSafe, catch_unwind}; + use std::sync::Arc; + use std::sync::atomic::{AtomicBool, Ordering}; + ++use arrow::ffi::FFI_ArrowSchema; ++use arrow::ffi_stream::FFI_ArrowArrayStream; + use arrow::record_batch::RecordBatchReader; + use arrow_array::RecordBatch; + use arrow_schema::{ArrowError, SchemaRef}; +@@ -50,6 +52,44 @@ use futures::{Stream, StreamExt}; + + use crate::error::{panic_payload_message, swallow_unwind}; + ++/// An owned, NUL-free error whose `Display` implementation cannot call back ++/// into an arbitrary external error source. Arrow formats this value from ++/// inside its non-unwinding `get_next` callback. ++#[derive(Debug)] ++struct FfiSafeStreamError(String); ++ ++impl std::fmt::Display for FfiSafeStreamError { ++ fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { ++ f.write_str(&self.0) ++ } ++} ++ ++impl std::error::Error for FfiSafeStreamError {} ++ ++fn ffi_safe_stream_error(message: String) -> ArrowError { ++ ArrowError::ExternalError(Box::new(FfiSafeStreamError(message.replace('\0', "\\0")))) ++} ++ ++/// Exercise arrow-rs's exact schema conversion before its non-unwinding ++/// `get_schema` callback is exposed to C. ++fn preflight_schema(schema: &SchemaRef) -> std::result::Result<(), ArrowError> { ++ match catch_unwind(AssertUnwindSafe(|| { ++ let ffi_schema = FFI_ArrowSchema::try_from(schema.as_ref()).map_err(|err| { ++ // Detach the error under the guard for the same reason `next` ++ // does: this value may ultimately be formatted by an FFI caller. ++ ffi_safe_stream_error(err.to_string()) ++ })?; ++ drop(ffi_schema); ++ Ok(()) ++ })) { ++ Ok(result) => result, ++ Err(payload) => Err(ffi_safe_stream_error(format!( ++ "panic exporting Arrow schema: {}", ++ panic_payload_message(&*payload) ++ ))), ++ } ++} ++ + /// A [`RecordBatchReader`] that owns the exported Lance stream, drives it + /// with a Tokio runtime handle, and contains panics at both C-reachable + /// edges (`next` and `drop`) — see the module docs for why the guard lives +@@ -74,12 +114,26 @@ impl GuardedReader { + /// Wrap `inner`, driving it with `handle` and wiring the shared + /// `scanner_poison` flag that a caught panic sets (from + /// `LanceScanner::poison_flag()` at the export sites). ++ /// ++ /// # Panics ++ /// ++ /// Panics if `schema` cannot be converted to the Arrow C Data Interface. ++ /// Production callers construct this reader inside their outer FFI panic ++ /// guard, before arrow-rs's non-unwinding `get_schema` callback is exposed. + pub fn new( + inner: S, + schema: SchemaRef, + handle: tokio::runtime::Handle, + scanner_poison: Arc, + ) -> Self { ++ // arrow-rs converts this schema later from inside its non-unwinding ++ // `get_schema` callback. Perform the exact conversion once while the ++ // scanner export's outer catch_unwind is still active, so a malformed ++ // schema (for example, a field name containing NUL) cannot first ++ // panic after control has crossed into that callback. ++ preflight_schema(&schema) ++ .unwrap_or_else(|err| panic!("Arrow schema cannot be exported: {err}")); ++ + Self { + inner: Some(inner), + schema, +@@ -115,19 +169,27 @@ where + // stream's `poll_next` lands here, one frame below arrow-rs's + // `extern "C"` callback, so neither can unwind across the FFI + // boundary. +- let polled = catch_unwind(AssertUnwindSafe(|| handle.block_on(inner.next()))); ++ let polled = catch_unwind(AssertUnwindSafe(|| { ++ match handle.block_on(inner.next()) { ++ Some(Ok(batch)) => Some(Ok(batch)), ++ Some(Err(err)) => { ++ // Format and detach the arbitrary Lance error while still ++ // inside the guard. arrow-rs later calls Display and ++ // CString::new from a non-unwinding callback, so neither a ++ // panicking source nor an embedded NUL may reach it. ++ Some(Err(ffi_safe_stream_error(err.to_string()))) ++ } ++ None => None, ++ } ++ })); + match polled { +- Ok(Some(Ok(batch))) => Some(Ok(batch)), +- Ok(Some(Err(err))) => Some(Err(ArrowError::ExternalError(Box::new(err)))), +- Ok(None) => None, ++ Ok(item) => item, + Err(payload) => { + *poisoned = true; + scanner_poison.store(true, Ordering::SeqCst); +- Some(Err(ArrowError::ExternalError(Box::new( +- lance_core::Error::internal(format!( +- "panic in stream: {}", +- panic_payload_message(&*payload) +- )), ++ Some(Err(ffi_safe_stream_error(format!( ++ "panic in stream: {}", ++ panic_payload_message(&*payload) + )))) + } + } +@@ -158,11 +220,101 @@ impl Drop for GuardedReader { + } + } + ++/// A panic-safe owner for an already-materialized [`RecordBatchReader`]. ++/// ++/// Dataset `take` operations use readers whose batches are already in memory, ++/// so no Tokio handle is needed. Arrow still invokes `schema`, `next`, and ++/// `drop` later from non-unwinding C callbacks, however, which requires the ++/// same error-detachment and cleanup containment as [`GuardedReader`]. ++struct GuardedRecordBatchReader { ++ inner: Option, ++ schema: SchemaRef, ++ poisoned: bool, ++} ++ ++impl Iterator for GuardedRecordBatchReader ++where ++ R: RecordBatchReader, ++{ ++ type Item = std::result::Result; ++ ++ fn next(&mut self) -> Option { ++ if self.poisoned { ++ return None; ++ } ++ ++ let inner = self.inner.as_mut()?; ++ let next = catch_unwind(AssertUnwindSafe(|| match inner.next() { ++ Some(Ok(batch)) => Some(Ok(batch)), ++ Some(Err(err)) => Some(Err(ffi_safe_stream_error(err.to_string()))), ++ None => None, ++ })); ++ ++ match next { ++ Ok(item) => item, ++ Err(payload) => { ++ self.poisoned = true; ++ Some(Err(ffi_safe_stream_error(format!( ++ "panic in record batch reader: {}", ++ panic_payload_message(&*payload) ++ )))) ++ } ++ } ++ } ++} ++ ++impl RecordBatchReader for GuardedRecordBatchReader ++where ++ R: RecordBatchReader + Send, ++{ ++ fn schema(&self) -> SchemaRef { ++ Arc::clone(&self.schema) ++ } ++} ++ ++impl Drop for GuardedRecordBatchReader { ++ fn drop(&mut self) { ++ let Some(inner) = self.inner.take() else { ++ return; ++ }; ++ swallow_unwind( ++ "GuardedRecordBatchReader::drop (ArrowArrayStream release)", ++ || drop(inner), ++ ); ++ } ++} ++ ++/// Export an already-materialized reader through panic-safe Arrow C stream ++/// callbacks. ++/// ++/// The schema is converted once before the callback table is returned. This ++/// turns deterministic schema conversion failures into an ordinary export ++/// failure (or lets the caller's outer FFI guard catch an arrow-rs conversion ++/// panic) instead of deferring them to `get_schema`. ++pub(crate) fn guarded_ffi_stream_from_reader( ++ reader: R, ++) -> std::result::Result ++where ++ R: RecordBatchReader + Send + 'static, ++{ ++ let schema = reader.schema(); ++ preflight_schema(&schema)?; ++ let reader = GuardedRecordBatchReader { ++ inner: Some(reader), ++ schema, ++ poisoned: false, ++ }; ++ Ok(FFI_ArrowArrayStream::new(Box::new(reader))) ++} ++ + #[cfg(test)] + mod tests { + use super::*; ++ use arrow::ffi::FFI_ArrowArray; ++ use arrow::ffi_stream::FFI_ArrowArrayStream; + use arrow_array::Int32Array; + use arrow_schema::{DataType, Field, Schema}; ++ use std::ffi::CStr; + use std::pin::Pin; + use std::task::{Context, Poll}; + +@@ -187,6 +339,17 @@ mod tests { + message: &'static str, + } + ++ #[derive(Debug)] ++ struct PanickingDisplay; ++ ++ impl std::fmt::Display for PanickingDisplay { ++ fn fmt(&self, _f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { ++ panic!("simulated panic while formatting a reader error") ++ } ++ } ++ ++ impl std::error::Error for PanickingDisplay {} ++ + impl Stream for PanicOnSecondPoll { + type Item = lance_core::Result; + +@@ -212,6 +375,76 @@ mod tests { + (rt, reader) + } + ++ fn guarded_export( ++ stream: S, ++ schema: SchemaRef, ++ scanner_poison: Arc, ++ ) -> (tokio::runtime::Runtime, FFI_ArrowArrayStream) ++ where ++ S: Stream> + Unpin + Send + 'static, ++ { ++ let rt = tokio::runtime::Runtime::new().unwrap(); ++ let reader = GuardedReader::new(stream, schema, rt.handle().clone(), scanner_poison); ++ (rt, FFI_ArrowArrayStream::new(Box::new(reader))) ++ } ++ ++ unsafe fn c_get_next(stream: *mut FFI_ArrowArrayStream, array: *mut FFI_ArrowArray) -> i32 { ++ let get_next = unsafe { (*stream).get_next }.expect("get_next callback is NULL"); ++ unsafe { get_next(stream, array) } ++ } ++ ++ unsafe fn c_get_schema(stream: *mut FFI_ArrowArrayStream, schema: *mut FFI_ArrowSchema) -> i32 { ++ let get_schema = unsafe { (*stream).get_schema }.expect("get_schema callback is NULL"); ++ unsafe { get_schema(stream, schema) } ++ } ++ ++ unsafe fn c_get_last_error(stream: *mut FFI_ArrowArrayStream) -> Option { ++ let get_last_error = ++ unsafe { (*stream).get_last_error }.expect("get_last_error callback is NULL"); ++ let message = unsafe { get_last_error(stream) }; ++ if message.is_null() { ++ None ++ } else { ++ Some( ++ unsafe { CStr::from_ptr(message) } ++ .to_string_lossy() ++ .into_owned(), ++ ) ++ } ++ } ++ ++ fn run_child(test_name: &str, environment_variable: &str) -> std::process::Output { ++ let exact_name = format!("stream_guard::tests::{test_name}"); ++ std::process::Command::new(std::env::current_exe().unwrap()) ++ .args([&exact_name, "--exact", "--nocapture", "--test-threads=1"]) ++ .env(environment_variable, "1") ++ .output() ++ .unwrap() ++ } ++ ++ fn assert_child_succeeds(test_name: &str, environment_variable: &str) -> String { ++ let output = run_child(test_name, environment_variable); ++ let stderr = String::from_utf8_lossy(&output.stderr).into_owned(); ++ assert!( ++ output.status.success(), ++ "guarded child must exit cleanly, got status {:?}\nstderr:\n{stderr}", ++ output.status ++ ); ++ stderr ++ } ++ ++ fn raw_error_then_eos(stream: &mut FFI_ArrowArrayStream) -> String { ++ let mut array = FFI_ArrowArray::empty(); ++ let status = unsafe { c_get_next(stream, &mut array) }; ++ assert_ne!(status, 0, "expected an Arrow C stream error"); ++ let message = unsafe { c_get_last_error(stream) }.expect("get_last_error returned NULL"); ++ ++ let mut eos = FFI_ArrowArray::empty(); ++ assert_eq!(unsafe { c_get_next(stream, &mut eos) }, 0); ++ assert!(eos.release.is_none(), "error must be followed by EOS"); ++ message ++ } ++ + #[test] + fn panic_yields_one_error_then_fuses_and_flips_flag() { + let scanner_poison = Arc::new(AtomicBool::new(false)); +@@ -363,6 +596,48 @@ mod tests { + } + } + ++ struct PanicOnReaderNext { ++ schema: SchemaRef, ++ } ++ ++ impl Iterator for PanicOnReaderNext { ++ type Item = std::result::Result; ++ ++ fn next(&mut self) -> Option { ++ panic!("simulated panic in materialized reader next"); ++ } ++ } ++ ++ impl RecordBatchReader for PanicOnReaderNext { ++ fn schema(&self) -> SchemaRef { ++ Arc::clone(&self.schema) ++ } ++ } ++ ++ struct PanicOnReaderDrop { ++ schema: SchemaRef, ++ } ++ ++ impl Iterator for PanicOnReaderDrop { ++ type Item = std::result::Result; ++ ++ fn next(&mut self) -> Option { ++ None ++ } ++ } ++ ++ impl RecordBatchReader for PanicOnReaderDrop { ++ fn schema(&self) -> SchemaRef { ++ Arc::clone(&self.schema) ++ } ++ } ++ ++ impl Drop for PanicOnReaderDrop { ++ fn drop(&mut self) { ++ panic!("simulated panic in materialized reader drop"); ++ } ++ } ++ + /// Regression for the review finding that the release path was unguarded: + /// arrow-rs's `release_stream` drops this reader inside its `extern "C"` + /// callback, so a cleanup panic must be contained here (best-effort: +@@ -381,4 +656,311 @@ mod tests { + "cleanup panic is best-effort and must not poison the handle" + ); + } ++ ++ #[test] ++ fn raw_stream_get_next_contains_poll_panic() { ++ const CHILD: &str = "LANCE_C_CHILD_STREAM_POLL_PANIC"; ++ if std::env::var(CHILD).is_err() { ++ let stderr = assert_child_succeeds("raw_stream_get_next_contains_poll_panic", CHILD); ++ assert!(stderr.contains("simulated raw poll panic")); ++ return; ++ } ++ ++ let scanner_poison = Arc::new(AtomicBool::new(false)); ++ let (_runtime, mut stream) = guarded_export( ++ PanicOnSecondPoll { ++ yielded: false, ++ message: "simulated raw poll panic", ++ }, ++ test_schema(), ++ Arc::clone(&scanner_poison), ++ ); ++ ++ let mut first = FFI_ArrowArray::empty(); ++ assert_eq!(unsafe { c_get_next(&mut stream, &mut first) }, 0); ++ assert!(first.release.is_some()); ++ unsafe { first.release.unwrap()(&mut first) }; ++ ++ let message = raw_error_then_eos(&mut stream); ++ assert!( ++ message.contains("simulated raw poll panic"), ++ "got: {message}" ++ ); ++ assert!(scanner_poison.load(Ordering::SeqCst)); ++ } ++ ++ #[test] ++ fn raw_stream_get_next_sanitizes_regular_error() { ++ const CHILD: &str = "LANCE_C_CHILD_STREAM_NUL_ERROR"; ++ if std::env::var(CHILD).is_err() { ++ assert_child_succeeds("raw_stream_get_next_sanitizes_regular_error", CHILD); ++ return; ++ } ++ ++ let scanner_poison = Arc::new(AtomicBool::new(false)); ++ let stream = futures::stream::iter(vec![Err(lance_core::Error::invalid_input_source( ++ "ordinary error with a NUL: bo\0om".into(), ++ ))]); ++ let (_runtime, mut stream) = ++ guarded_export(stream, test_schema(), Arc::clone(&scanner_poison)); ++ ++ let message = raw_error_then_eos(&mut stream); ++ assert!(message.contains("bo\\0om"), "got: {message:?}"); ++ assert!(!message.contains('\0')); ++ assert!(!scanner_poison.load(Ordering::SeqCst)); ++ } ++ ++ #[test] ++ fn raw_stream_get_next_contains_error_display_panic() { ++ const CHILD: &str = "LANCE_C_CHILD_STREAM_DISPLAY_PANIC"; ++ if std::env::var(CHILD).is_err() { ++ let stderr = ++ assert_child_succeeds("raw_stream_get_next_contains_error_display_panic", CHILD); ++ assert!(stderr.contains("simulated panic while formatting a reader error")); ++ return; ++ } ++ ++ let scanner_poison = Arc::new(AtomicBool::new(false)); ++ let stream = futures::stream::iter(vec![Err(lance_core::Error::invalid_input_source( ++ Box::new(PanickingDisplay), ++ ))]); ++ let (_runtime, mut stream) = ++ guarded_export(stream, test_schema(), Arc::clone(&scanner_poison)); ++ ++ let message = raw_error_then_eos(&mut stream); ++ assert!( ++ message.contains("simulated panic while formatting a reader error"), ++ "got: {message}" ++ ); ++ assert!(scanner_poison.load(Ordering::SeqCst)); ++ } ++ ++ #[test] ++ fn stream_schema_is_rejected_before_raw_get_schema_is_exposed() { ++ const CHILD: &str = "LANCE_C_CHILD_STREAM_NUL_SCHEMA"; ++ if std::env::var(CHILD).is_err() { ++ assert_child_succeeds( ++ "stream_schema_is_rejected_before_raw_get_schema_is_exposed", ++ CHILD, ++ ); ++ return; ++ } ++ ++ let schema = Arc::new(Schema::new(vec![Field::new( ++ "field\0name", ++ DataType::Int32, ++ false, ++ )])); ++ let runtime = tokio::runtime::Runtime::new().unwrap(); ++ let callback_was_exposed = std::cell::Cell::new(false); ++ let outcome = catch_unwind(AssertUnwindSafe(|| { ++ let reader = GuardedReader::new( ++ futures::stream::empty::>(), ++ schema, ++ runtime.handle().clone(), ++ Arc::new(AtomicBool::new(false)), ++ ); ++ callback_was_exposed.set(true); ++ let mut stream = FFI_ArrowArrayStream::new(Box::new(reader)); ++ let mut ffi_schema = FFI_ArrowSchema::empty(); ++ unsafe { c_get_schema(&mut stream, &mut ffi_schema) } ++ })); ++ ++ assert!(outcome.is_err(), "invalid schema must fail during export"); ++ assert!( ++ !callback_was_exposed.get(), ++ "invalid schema reached raw get_schema" ++ ); ++ } ++ ++ #[test] ++ fn raw_stream_get_next_inside_tokio_runtime_is_contained() { ++ const CHILD: &str = "LANCE_C_CHILD_STREAM_NESTED_RUNTIME"; ++ if std::env::var(CHILD).is_err() { ++ let stderr = assert_child_succeeds( ++ "raw_stream_get_next_inside_tokio_runtime_is_contained", ++ CHILD, ++ ); ++ assert!(stderr.contains("Cannot start a runtime from within a runtime")); ++ return; ++ } ++ ++ let scanner_poison = Arc::new(AtomicBool::new(false)); ++ let (runtime, mut stream) = guarded_export( ++ futures::stream::iter(vec![Ok(test_batch())]), ++ test_schema(), ++ Arc::clone(&scanner_poison), ++ ); ++ let mut array = FFI_ArrowArray::empty(); ++ let status = runtime.block_on(async { unsafe { c_get_next(&mut stream, &mut array) } }); ++ assert_ne!(status, 0); ++ let message = unsafe { c_get_last_error(&mut stream) }.unwrap(); ++ assert!(message.contains("runtime"), "got: {message}"); ++ assert!(scanner_poison.load(Ordering::SeqCst)); ++ ++ let mut eos = FFI_ArrowArray::empty(); ++ assert_eq!(unsafe { c_get_next(&mut stream, &mut eos) }, 0); ++ assert!(eos.release.is_none()); ++ } ++ ++ #[test] ++ fn raw_stream_release_contains_drop_panic() { ++ const CHILD: &str = "LANCE_C_CHILD_STREAM_DROP_PANIC"; ++ if std::env::var(CHILD).is_err() { ++ let stderr = assert_child_succeeds("raw_stream_release_contains_drop_panic", CHILD); ++ assert!(stderr.contains("simulated drop bug in stream cleanup")); ++ return; ++ } ++ ++ let scanner_poison = Arc::new(AtomicBool::new(false)); ++ let (_runtime, mut stream) = ++ guarded_export(PanicOnDrop, test_schema(), Arc::clone(&scanner_poison)); ++ let release = stream.release.expect("release callback is NULL"); ++ unsafe { release(&mut stream) }; ++ assert!(stream.release.is_none()); ++ assert!(!scanner_poison.load(Ordering::SeqCst)); ++ } ++ ++ #[test] ++ fn guarded_in_memory_export_supports_raw_arrow_callbacks() { ++ let reader = ++ arrow::record_batch::RecordBatchIterator::new(vec![Ok(test_batch())], test_schema()); ++ let mut stream = guarded_ffi_stream_from_reader(reader).unwrap(); ++ ++ let mut schema = FFI_ArrowSchema::empty(); ++ let get_schema = stream.get_schema.expect("get_schema callback is NULL"); ++ assert_eq!(unsafe { get_schema(&mut stream, &mut schema) }, 0); ++ assert!(schema.release.is_some()); ++ unsafe { schema.release.unwrap()(&mut schema) }; ++ ++ let get_next = stream.get_next.expect("get_next callback is NULL"); ++ let mut array = FFI_ArrowArray::empty(); ++ assert_eq!(unsafe { get_next(&mut stream, &mut array) }, 0); ++ assert!(array.release.is_some()); ++ unsafe { array.release.unwrap()(&mut array) }; ++ ++ let mut eos = FFI_ArrowArray::empty(); ++ assert_eq!(unsafe { get_next(&mut stream, &mut eos) }, 0); ++ assert!(eos.release.is_none()); ++ ++ let release = stream.release.expect("release callback is NULL"); ++ unsafe { release(&mut stream) }; ++ assert!(stream.release.is_none()); ++ } ++ ++ #[test] ++ fn guarded_in_memory_get_next_sanitizes_nul_error() { ++ const CHILD: &str = "LANCE_C_CHILD_READER_NUL_ERROR"; ++ if std::env::var(CHILD).is_err() { ++ assert_child_succeeds("guarded_in_memory_get_next_sanitizes_nul_error", CHILD); ++ return; ++ } ++ ++ let reader = arrow::record_batch::RecordBatchIterator::new( ++ vec![Err(ArrowError::ComputeError( ++ "ordinary reader error: bo\0om".into(), ++ ))], ++ test_schema(), ++ ); ++ let mut stream = guarded_ffi_stream_from_reader(reader).unwrap(); ++ let message = raw_error_then_eos(&mut stream); ++ assert!(message.contains("bo\\0om"), "got: {message:?}"); ++ assert!(!message.contains('\0')); ++ } ++ ++ #[test] ++ fn guarded_in_memory_get_next_contains_error_display_panic() { ++ const CHILD: &str = "LANCE_C_CHILD_READER_DISPLAY_PANIC"; ++ if std::env::var(CHILD).is_err() { ++ let stderr = assert_child_succeeds( ++ "guarded_in_memory_get_next_contains_error_display_panic", ++ CHILD, ++ ); ++ assert!(stderr.contains("simulated panic while formatting a reader error")); ++ return; ++ } ++ ++ let reader = arrow::record_batch::RecordBatchIterator::new( ++ vec![Err(ArrowError::ExternalError(Box::new(PanickingDisplay)))], ++ test_schema(), ++ ); ++ let mut stream = guarded_ffi_stream_from_reader(reader).unwrap(); ++ let message = raw_error_then_eos(&mut stream); ++ assert!( ++ message.contains("simulated panic while formatting a reader error"), ++ "got: {message}" ++ ); ++ } ++ ++ #[test] ++ fn guarded_in_memory_get_next_contains_reader_panic() { ++ const CHILD: &str = "LANCE_C_CHILD_READER_NEXT_PANIC"; ++ if std::env::var(CHILD).is_err() { ++ let stderr = ++ assert_child_succeeds("guarded_in_memory_get_next_contains_reader_panic", CHILD); ++ assert!(stderr.contains("simulated panic in materialized reader next")); ++ return; ++ } ++ ++ let reader = PanicOnReaderNext { ++ schema: test_schema(), ++ }; ++ let mut stream = guarded_ffi_stream_from_reader(reader).unwrap(); ++ let message = raw_error_then_eos(&mut stream); ++ assert!( ++ message.contains("simulated panic in materialized reader next"), ++ "got: {message}" ++ ); ++ } ++ ++ #[test] ++ fn guarded_in_memory_rejects_nul_schema_before_callback_exposure() { ++ const CHILD: &str = "LANCE_C_CHILD_READER_NUL_SCHEMA"; ++ if std::env::var(CHILD).is_err() { ++ let stderr = assert_child_succeeds( ++ "guarded_in_memory_rejects_nul_schema_before_callback_exposure", ++ CHILD, ++ ); ++ assert!(stderr.contains("NulError")); ++ return; ++ } ++ ++ let schema = Arc::new(Schema::new(vec![Field::new( ++ "field\0name", ++ DataType::Int32, ++ false, ++ )])); ++ let reader = arrow::record_batch::RecordBatchIterator::new( ++ Vec::>::new(), ++ schema, ++ ); ++ ++ let result = guarded_ffi_stream_from_reader(reader); ++ let error = result.expect_err("invalid schema must fail before export"); ++ assert!( ++ error.to_string().contains("panic exporting Arrow schema"), ++ "got: {error}" ++ ); ++ } ++ ++ #[test] ++ fn guarded_in_memory_release_contains_reader_drop_panic() { ++ const CHILD: &str = "LANCE_C_CHILD_READER_DROP_PANIC"; ++ if std::env::var(CHILD).is_err() { ++ let stderr = assert_child_succeeds( ++ "guarded_in_memory_release_contains_reader_drop_panic", ++ CHILD, ++ ); ++ assert!(stderr.contains("simulated panic in materialized reader drop")); ++ return; ++ } ++ ++ let reader = PanicOnReaderDrop { ++ schema: test_schema(), ++ }; ++ let mut stream = guarded_ffi_stream_from_reader(reader).unwrap(); ++ let release = stream.release.expect("release callback is NULL"); ++ unsafe { release(&mut stream) }; ++ assert!(stream.release.is_none()); ++ } + } +diff --git a/tests/c_api_test.rs b/tests/c_api_test.rs +index daa7425..17dc6c6 100644 +--- a/tests/c_api_test.rs ++++ b/tests/c_api_test.rs +@@ -993,7 +993,7 @@ fn test_scanner_scan_async() { + unsafe { + lance_scanner_scan_async( + scanner, +- on_complete, ++ Some(on_complete), + Arc::as_ptr(&pair_clone) as *mut std::ffi::c_void, + ); + lance_scanner_close(scanner); +@@ -1015,10 +1015,25 @@ fn test_scanner_scan_async() { + assert_eq!(total_rows, 5); + assert_eq!(captured.calls.load(AtomicOrdering::SeqCst), 1); + assert!(!captured.invalid_statistics.load(AtomicOrdering::SeqCst)); ++ unsafe { ++ lance_scanner_async_stream_free(result.stream_ptr.cast::()); ++ } + + unsafe { lance_dataset_close(ds) }; + } + ++#[test] ++fn test_scanner_async_stream_free_releases_stream_and_accepts_null() { ++ let (stream, drop_count) = make_counted_column_stream("value", vec![1]); ++ let stream = Box::into_raw(Box::new(stream)); ++ ++ unsafe { lance_scanner_async_stream_free(stream) }; ++ assert_eq!(drop_count.load(AtomicOrdering::SeqCst), 1); ++ ++ // Match the other close/free APIs: NULL is a no-op. ++ unsafe { lance_scanner_async_stream_free(ptr::null_mut()) }; ++} ++ + // =========================================================================== + // Additional tests + // =========================================================================== +@@ -1594,7 +1609,7 @@ fn test_async_scan_with_filter() { + unsafe { + lance_scanner_scan_async( + scanner, +- on_complete, ++ Some(on_complete), + Arc::as_ptr(&pair_clone) as *mut std::ffi::c_void, + ); + } +@@ -1609,6 +1624,9 @@ fn test_async_scan_with_filter() { + let ffi_stream = unsafe { &mut *(result.stream_ptr as *mut FFI_ArrowArrayStream) }; + let reader = unsafe { ArrowArrayStreamReader::from_raw(ffi_stream) }.unwrap(); + assert_eq!(reader.map(|r| r.unwrap().num_rows()).sum::(), 2); ++ unsafe { ++ lance_scanner_async_stream_free(result.stream_ptr.cast::()); ++ } + + unsafe { lance_scanner_close(scanner) }; + unsafe { lance_dataset_close(ds) }; +@@ -1641,7 +1659,7 @@ fn test_poll_next_basic() { + loop { + let mut batch: *mut LanceBatch = ptr::null_mut(); + let status = unsafe { +- lance_scanner_poll_next(scanner, test_waker, ptr::null_mut(), &mut batch) ++ lance_scanner_poll_next(scanner, Some(test_waker), ptr::null_mut(), &mut batch) + }; + match status { + LancePollStatus::Ready => { +@@ -3042,6 +3060,45 @@ fn test_index_segment_builder_owns_snapshot_and_is_single_use() { + } + } + ++#[test] ++fn test_index_segment_metadata_accessors_reject_null_handles() { ++ assert!(unsafe { lance_index_segment_metadata_name(ptr::null()) }.is_null()); ++ assert_eq!(lance_last_error_code(), LanceErrorCode::InvalidArgument); ++ ++ assert_eq!( ++ unsafe { lance_index_segment_metadata_dataset_version(ptr::null()) }, ++ 0 ++ ); ++ assert_eq!(lance_last_error_code(), LanceErrorCode::InvalidArgument); ++ ++ assert_eq!( ++ unsafe { lance_index_segment_metadata_index_version(ptr::null()) }, ++ -1 ++ ); ++ assert_eq!(lance_last_error_code(), LanceErrorCode::InvalidArgument); ++ ++ assert_eq!( ++ unsafe { lance_index_segment_metadata_index_type(ptr::null()) }, ++ -1 ++ ); ++ assert_eq!(lance_last_error_code(), LanceErrorCode::InvalidArgument); ++ ++ assert!(unsafe { lance_index_segment_metadata_index_details_type_url(ptr::null()) }.is_null()); ++ assert_eq!(lance_last_error_code(), LanceErrorCode::InvalidArgument); ++ ++ assert_eq!( ++ unsafe { lance_index_segment_metadata_field_count(ptr::null()) }, ++ 0 ++ ); ++ assert_eq!(lance_last_error_code(), LanceErrorCode::InvalidArgument); ++ ++ assert_eq!( ++ unsafe { lance_index_segment_metadata_fragment_count(ptr::null()) }, ++ 0 ++ ); ++ assert_eq!(lance_last_error_code(), LanceErrorCode::InvalidArgument); ++} ++ + #[test] + fn test_index_segment_metadata_parse_rejects_malformed_and_dangerous_input() { + use prost::Message; +@@ -5543,6 +5600,431 @@ fn test_fts_fuzzy() { + unsafe { lance_dataset_close(ds) }; + } + ++fn collect_context_fts_scores( ++ dataset: *const LanceDataset, ++ context: *const LanceFtsQueryContext, ++ segment_uuids: Option<&[[u8; 16]]>, ++) -> std::collections::HashMap { ++ let scanner = unsafe { lance_scanner_new(dataset, ptr::null(), ptr::null()) }; ++ assert!(!scanner.is_null()); ++ assert_eq!( ++ unsafe { lance_scanner_set_fts_query_context(scanner, context) }, ++ 0, ++ "{}", ++ unsafe { std::ffi::CStr::from_ptr(lance_last_error_message()).to_string_lossy() } ++ ); ++ if let Some(segment_uuids) = segment_uuids { ++ assert_eq!( ++ unsafe { ++ lance_scanner_set_fts_index_segments( ++ scanner, ++ segment_uuids.as_ptr().cast::(), ++ segment_uuids.len(), ++ ) ++ }, ++ 0 ++ ); ++ } ++ ++ let mut stream = FFI_ArrowArrayStream::empty(); ++ assert_eq!( ++ unsafe { lance_scanner_to_arrow_stream(scanner, &mut stream) }, ++ 0, ++ "{}", ++ unsafe { std::ffi::CStr::from_ptr(lance_last_error_message()).to_string_lossy() } ++ ); ++ let reader = unsafe { ArrowArrayStreamReader::from_raw(&mut stream).unwrap() }; ++ let mut scores = std::collections::HashMap::new(); ++ for batch in reader { ++ let batch = batch.unwrap(); ++ let ids = batch ++ .column_by_name("id") ++ .unwrap() ++ .as_any() ++ .downcast_ref::() ++ .unwrap(); ++ let batch_scores = batch ++ .column_by_name("_score") ++ .unwrap() ++ .as_any() ++ .downcast_ref::() ++ .unwrap(); ++ for row in 0..batch.num_rows() { ++ assert!( ++ scores ++ .insert(ids.value(row), batch_scores.value(row)) ++ .is_none() ++ ); ++ } ++ } ++ unsafe { lance_scanner_close(scanner) }; ++ scores ++} ++ ++fn load_fts_segment_uuids(uri: &str, column: &str) -> Vec<[u8; 16]> { ++ use lance::index::DatasetIndexExt; ++ use lance_index::IndexCriteria; ++ ++ lance_c::runtime::block_on(async { ++ let dataset = Dataset::open(uri).await.unwrap(); ++ let logical_index = dataset ++ .load_scalar_index(IndexCriteria::default().for_column(column).supports_fts()) ++ .await ++ .unwrap() ++ .unwrap(); ++ dataset ++ .load_indices_by_name(&logical_index.name) ++ .await ++ .unwrap() ++ .into_iter() ++ .map(|segment| *segment.uuid.as_bytes()) ++ .collect() ++ }) ++} ++ ++#[test] ++fn test_prepare_fts_query_index_only_allows_unindexed_fragment() { ++ let (_tmp, uri) = create_test_dataset(); ++ let uri_c = c_str(&uri); ++ let column = c_str("name"); ++ let query = c_str("alice"); ++ let inverted_params = c_str(r#"{"base_tokenizer":"simple","language":"English"}"#); ++ ++ let indexed_snapshot = unsafe { lance_dataset_open(uri_c.as_ptr(), ptr::null(), 0) }; ++ assert_eq!( ++ unsafe { ++ lance_dataset_create_scalar_index( ++ indexed_snapshot, ++ column.as_ptr(), ++ ptr::null(), ++ LanceScalarIndexType::Inverted as i32, ++ inverted_params.as_ptr(), ++ false, ++ ) ++ }, ++ 0 ++ ); ++ unsafe { lance_dataset_close(indexed_snapshot) }; ++ ++ let schema = Arc::new(Schema::new(vec![ ++ Field::new("id", DataType::Int32, false), ++ Field::new("name", DataType::Utf8, true), ++ ])); ++ let batch = RecordBatch::try_new( ++ schema.clone(), ++ vec![ ++ Arc::new(Int32Array::from(vec![6, 7])), ++ Arc::new(StringArray::from(vec!["alice", "alice alice"])), ++ ], ++ ) ++ .unwrap(); ++ append_batch(&uri, schema, batch); ++ ++ let dataset = unsafe { lance_dataset_open(uri_c.as_ptr(), ptr::null(), 0) }; ++ let strict = unsafe { ++ lance_dataset_prepare_fts_query( ++ dataset, ++ column.as_ptr(), ++ query.as_ptr(), ++ 0, ++ LanceFtsCoverageMode::Strict as i32, ++ ) ++ }; ++ assert!(strict.is_null()); ++ assert_eq!(lance_last_error_code(), LanceErrorCode::InvalidArgument); ++ let message = unsafe { ++ std::ffi::CStr::from_ptr(lance_last_error_message()) ++ .to_string_lossy() ++ .into_owned() ++ }; ++ assert!(message.contains("unindexed fragments"), "{message}"); ++ ++ let context = unsafe { ++ lance_dataset_prepare_fts_query( ++ dataset, ++ column.as_ptr(), ++ query.as_ptr(), ++ 0, ++ LanceFtsCoverageMode::IndexOnly as i32, ++ ) ++ }; ++ assert!(!context.is_null()); ++ let segment_uuids = load_fts_segment_uuids(&uri, "name"); ++ assert_eq!(segment_uuids.len(), 1); ++ ++ let scanner = unsafe { lance_scanner_new(dataset, ptr::null(), ptr::null()) }; ++ assert_eq!( ++ unsafe { lance_scanner_set_fts_query_context(scanner, context) }, ++ 0 ++ ); ++ // Scanner retains an Arc; closing the public handle does not invalidate it. ++ unsafe { lance_fts_query_context_close(context) }; ++ assert_eq!( ++ unsafe { ++ lance_scanner_set_fts_index_segments( ++ scanner, ++ segment_uuids.as_ptr().cast::(), ++ segment_uuids.len(), ++ ) ++ }, ++ 0 ++ ); ++ let mut stream = FFI_ArrowArrayStream::empty(); ++ assert_eq!( ++ unsafe { lance_scanner_to_arrow_stream(scanner, &mut stream) }, ++ 0 ++ ); ++ let reader = unsafe { ArrowArrayStreamReader::from_raw(&mut stream).unwrap() }; ++ let total_rows: usize = reader.map(|batch| batch.unwrap().num_rows()).sum(); ++ assert_eq!( ++ total_rows, 1, ++ "INDEX_ONLY must exclude both matching rows in the unindexed fragment" ++ ); ++ ++ unsafe { lance_scanner_close(scanner) }; ++ unsafe { lance_dataset_close(dataset) }; ++} ++ ++#[test] ++fn test_prepared_fts_global_scorer_is_shared_across_segment_splits() { ++ use lance::index::DatasetIndexExt; ++ use lance_index::optimize::OptimizeOptions; ++ ++ let (_tmp, uri) = create_test_dataset(); ++ let uri_c = c_str(&uri); ++ let column = c_str("name"); ++ let query = c_str("alice"); ++ let inverted_params = c_str(r#"{"base_tokenizer":"simple","language":"English"}"#); ++ ++ let dataset = unsafe { lance_dataset_open(uri_c.as_ptr(), ptr::null(), 0) }; ++ assert_eq!( ++ unsafe { ++ lance_dataset_create_scalar_index( ++ dataset, ++ column.as_ptr(), ++ ptr::null(), ++ LanceScalarIndexType::Inverted as i32, ++ inverted_params.as_ptr(), ++ false, ++ ) ++ }, ++ 0 ++ ); ++ unsafe { lance_dataset_close(dataset) }; ++ ++ let schema = Arc::new(Schema::new(vec![ ++ Field::new("id", DataType::Int32, false), ++ Field::new("name", DataType::Utf8, true), ++ ])); ++ append_batch( ++ &uri, ++ schema.clone(), ++ RecordBatch::try_new( ++ schema, ++ vec![ ++ Arc::new(Int32Array::from(vec![6, 7])), ++ Arc::new(StringArray::from(vec!["alice", "alice alice"])), ++ ], ++ ) ++ .unwrap(), ++ ); ++ lance_c::runtime::block_on(async { ++ let mut dataset = Dataset::open(&uri).await.unwrap(); ++ dataset ++ .optimize_indices(&OptimizeOptions::append()) ++ .await ++ .unwrap(); ++ }); ++ ++ let dataset = unsafe { lance_dataset_open(uri_c.as_ptr(), ptr::null(), 0) }; ++ let context = unsafe { ++ lance_dataset_prepare_fts_query( ++ dataset, ++ column.as_ptr(), ++ query.as_ptr(), ++ 0, ++ LanceFtsCoverageMode::Strict as i32, ++ ) ++ }; ++ assert!(!context.is_null(), "{}", unsafe { ++ std::ffi::CStr::from_ptr(lance_last_error_message()).to_string_lossy() ++ }); ++ let segment_uuids = load_fts_segment_uuids(&uri, "name"); ++ assert_eq!(segment_uuids.len(), 2); ++ ++ let full_scores = collect_context_fts_scores(dataset, context, None); ++ assert_eq!(full_scores.len(), 3); ++ let mut split_scores = std::collections::HashMap::new(); ++ for segment_uuid in &segment_uuids { ++ for (id, score) in ++ collect_context_fts_scores(dataset, context, Some(std::slice::from_ref(segment_uuid))) ++ { ++ assert!(split_scores.insert(id, score).is_none()); ++ } ++ } ++ assert_eq!(split_scores.len(), full_scores.len()); ++ for (id, expected_score) in full_scores { ++ let actual_score = split_scores.get(&id).unwrap(); ++ assert!( ++ (actual_score - expected_score).abs() < 1e-6, ++ "id={id}, full={expected_score}, split={actual_score}" ++ ); ++ } ++ ++ let scanner = unsafe { lance_scanner_new(dataset, ptr::null(), ptr::null()) }; ++ let duplicate_segments = [segment_uuids[0], segment_uuids[0]]; ++ assert_eq!( ++ unsafe { ++ lance_scanner_set_fts_index_segments( ++ scanner, ++ duplicate_segments.as_ptr().cast::(), ++ duplicate_segments.len(), ++ ) ++ }, ++ -1 ++ ); ++ assert!(unsafe { lance_scanner_set_fts_query_context(scanner, ptr::null()) } < 0); ++ ++ unsafe { lance_scanner_close(scanner) }; ++ ++ let unknown_segment_scanner = unsafe { lance_scanner_new(dataset, ptr::null(), ptr::null()) }; ++ assert_eq!( ++ unsafe { lance_scanner_set_fts_query_context(unknown_segment_scanner, context) }, ++ 0 ++ ); ++ let unknown_uuid = [0_u8; 16]; ++ assert_eq!( ++ unsafe { ++ lance_scanner_set_fts_index_segments(unknown_segment_scanner, unknown_uuid.as_ptr(), 1) ++ }, ++ 0, ++ "membership is validated against the attached context at scan time" ++ ); ++ let mut stream = FFI_ArrowArrayStream::empty(); ++ assert_eq!( ++ unsafe { lance_scanner_to_arrow_stream(unknown_segment_scanner, &mut stream) }, ++ -1 ++ ); ++ unsafe { lance_scanner_close(unknown_segment_scanner) }; ++ ++ let independently_reopened = unsafe { lance_dataset_open(uri_c.as_ptr(), ptr::null(), 0) }; ++ assert!(!independently_reopened.is_null()); ++ assert_eq!( ++ unsafe { lance_dataset_version(independently_reopened) }, ++ unsafe { lance_dataset_version(dataset) }, ++ "the identity check must reject equal URI/version locator metadata" ++ ); ++ let reopened_scanner = ++ unsafe { lance_scanner_new(independently_reopened, ptr::null(), ptr::null()) }; ++ assert_eq!( ++ unsafe { lance_scanner_set_fts_query_context(reopened_scanner, context) }, ++ -1, ++ "an independently opened dataset must not reuse the prepared context" ++ ); ++ let message = unsafe { ++ std::ffi::CStr::from_ptr(lance_last_error_message()) ++ .to_string_lossy() ++ .into_owned() ++ }; ++ assert!( ++ message.contains("same process-local dataset snapshot"), ++ "{message}" ++ ); ++ unsafe { lance_scanner_close(reopened_scanner) }; ++ unsafe { lance_dataset_close(independently_reopened) }; ++ ++ let old_snapshot = unsafe { lance_dataset_open(uri_c.as_ptr(), ptr::null(), 2) }; ++ assert!(!old_snapshot.is_null()); ++ let old_snapshot_scanner = unsafe { lance_scanner_new(old_snapshot, ptr::null(), ptr::null()) }; ++ assert_eq!( ++ unsafe { lance_scanner_set_fts_query_context(old_snapshot_scanner, context) }, ++ -1, ++ "a context must not be attached to a different dataset version" ++ ); ++ unsafe { lance_scanner_close(old_snapshot_scanner) }; ++ unsafe { lance_dataset_close(old_snapshot) }; ++ ++ unsafe { lance_fts_query_context_close(context) }; ++ unsafe { lance_dataset_close(dataset) }; ++} ++ ++#[test] ++fn test_prepare_fts_query_rejects_null_empty_invalid_mode_and_fuzzy() { ++ let (_tmp, uri) = create_test_dataset(); ++ let uri_c = c_str(&uri); ++ let dataset = unsafe { lance_dataset_open(uri_c.as_ptr(), ptr::null(), 0) }; ++ let column = c_str("name"); ++ let query = c_str("alice"); ++ let empty = c_str(""); ++ ++ assert!( ++ unsafe { ++ lance_dataset_prepare_fts_query( ++ ptr::null(), ++ column.as_ptr(), ++ query.as_ptr(), ++ 0, ++ LanceFtsCoverageMode::Strict as i32, ++ ) ++ } ++ .is_null() ++ ); ++ assert!( ++ unsafe { ++ lance_dataset_prepare_fts_query( ++ dataset, ++ empty.as_ptr(), ++ query.as_ptr(), ++ 0, ++ LanceFtsCoverageMode::Strict as i32, ++ ) ++ } ++ .is_null() ++ ); ++ assert!( ++ unsafe { lance_dataset_prepare_fts_query(dataset, column.as_ptr(), empty.as_ptr(), 0, 0) } ++ .is_null() ++ ); ++ assert!( ++ unsafe { lance_dataset_prepare_fts_query(dataset, column.as_ptr(), query.as_ptr(), 0, 99) } ++ .is_null() ++ ); ++ assert!( ++ unsafe { ++ lance_dataset_prepare_fts_query( ++ dataset, ++ column.as_ptr(), ++ query.as_ptr(), ++ 1, ++ LanceFtsCoverageMode::Strict as i32, ++ ) ++ } ++ .is_null() ++ ); ++ assert_eq!(lance_last_error_code(), LanceErrorCode::InvalidArgument); ++ let message = unsafe { ++ std::ffi::CStr::from_ptr(lance_last_error_message()) ++ .to_string_lossy() ++ .into_owned() ++ }; ++ assert!( ++ message.contains("max_fuzzy_distance must be 0"), ++ "{message}" ++ ); ++ let scanner = unsafe { lance_scanner_new(dataset, ptr::null(), ptr::null()) }; ++ assert_eq!( ++ unsafe { lance_scanner_set_fts_index_segments(scanner, ptr::null(), 1) }, ++ -1 ++ ); ++ assert_eq!( ++ unsafe { lance_scanner_set_fts_index_segments(scanner, ptr::null(), 0) }, ++ 0 ++ ); ++ unsafe { lance_scanner_close(scanner) }; ++ unsafe { lance_fts_query_context_close(ptr::null_mut()) }; ++ unsafe { lance_dataset_close(dataset) }; ++} ++ + #[test] + fn test_nearest_after_fts_is_rejected() { + let (_tmp, uri) = create_vector_dataset(64, 8); +@@ -6418,18 +6900,17 @@ fn test_scanner_with_substrait_filter() { + + #[test] + fn test_scanner_substrait_filter_overrides_sql_filter() { +- // If both SQL and Substrait filters are set, Substrait wins (last write). ++ // If both primary filters are set, Substrait wins. + let (_tmp, uri) = create_test_dataset(); + let c_uri = c_str(&uri); + let ds = unsafe { lance_dataset_open(c_uri.as_ptr(), ptr::null(), 0) }; + assert!(!ds.is_null()); + +- // Start with SQL filter "id < 0" (matches 0 rows). + let sql = c_str("id < 0"); + let scanner = unsafe { lance_scanner_new(ds, ptr::null(), sql.as_ptr()) }; + assert!(!scanner.is_null()); + +- // Override with Substrait filter "id > 3" (matches 2 rows). ++ // Attach Substrait filter "id > 3" (matches id=4 and id=5). + let bytes = substrait_id_gt_3(); + let rc = unsafe { lance_scanner_set_substrait_filter(scanner, bytes.as_ptr(), bytes.len()) }; + assert_eq!(rc, 0); +@@ -6446,6 +6927,81 @@ fn test_scanner_substrait_filter_overrides_sql_filter() { + unsafe { lance_dataset_close(ds) }; + } + ++#[test] ++fn test_scanner_additional_sql_filters_are_anded_with_substrait() { ++ let (_tmp, uri) = create_test_dataset(); ++ let c_uri = c_str(&uri); ++ let ds = unsafe { lance_dataset_open(c_uri.as_ptr(), ptr::null(), 0) }; ++ assert!(!ds.is_null()); ++ ++ let scanner = unsafe { lance_scanner_new(ds, ptr::null(), ptr::null()) }; ++ assert!(!scanner.is_null()); ++ ++ let bytes = substrait_id_gt_3(); ++ assert_eq!( ++ unsafe { lance_scanner_set_substrait_filter(scanner, bytes.as_ptr(), bytes.len()) }, ++ 0 ++ ); ++ for sql in [c_str("id < 6"), c_str("id < 5")] { ++ assert_eq!( ++ unsafe { lance_scanner_additional_sql_filter(scanner, sql.as_ptr()) }, ++ 0 ++ ); ++ } ++ ++ let mut ffi_stream = FFI_ArrowArrayStream::empty(); ++ assert_eq!( ++ unsafe { lance_scanner_to_arrow_stream(scanner, &mut ffi_stream) }, ++ 0 ++ ); ++ ++ let reader = unsafe { ArrowArrayStreamReader::from_raw(&mut ffi_stream) }.unwrap(); ++ let total_rows: usize = reader.map(|r| r.unwrap().num_rows()).sum(); ++ assert_eq!(total_rows, 1, "id > 3 AND id < 6 AND id < 5 matches id=4"); ++ ++ unsafe { lance_scanner_close(scanner) }; ++ unsafe { lance_dataset_close(ds) }; ++} ++ ++#[test] ++fn test_scanner_additional_sql_filter_rejects_invalid_inputs() { ++ let (_tmp, uri) = create_test_dataset(); ++ let c_uri = c_str(&uri); ++ let ds = unsafe { lance_dataset_open(c_uri.as_ptr(), ptr::null(), 0) }; ++ let scanner = unsafe { lance_scanner_new(ds, ptr::null(), ptr::null()) }; ++ assert!(!scanner.is_null()); ++ ++ let filter = c_str("id > 3"); ++ assert_eq!( ++ unsafe { lance_scanner_additional_sql_filter(ptr::null_mut(), filter.as_ptr()) }, ++ -1 ++ ); ++ assert_eq!( ++ unsafe { lance_scanner_additional_sql_filter(scanner, ptr::null()) }, ++ -1 ++ ); ++ let empty = c_str(""); ++ assert_eq!( ++ unsafe { lance_scanner_additional_sql_filter(scanner, empty.as_ptr()) }, ++ -1 ++ ); ++ ++ let mut ffi_stream = FFI_ArrowArrayStream::empty(); ++ assert_eq!( ++ unsafe { lance_scanner_to_arrow_stream(scanner, &mut ffi_stream) }, ++ 0 ++ ); ++ assert_eq!( ++ unsafe { lance_scanner_additional_sql_filter(scanner, filter.as_ptr()) }, ++ -1, ++ "additional filters must be rejected after the scan starts" ++ ); ++ drop(unsafe { ArrowArrayStreamReader::from_raw(&mut ffi_stream) }.unwrap()); ++ ++ unsafe { lance_scanner_close(scanner) }; ++ unsafe { lance_dataset_close(ds) }; ++} ++ + #[test] + fn test_scanner_set_substrait_filter_invalid_inputs() { + let (_tmp, uri) = create_test_dataset(); +@@ -10388,8 +10944,8 @@ fn test_add_columns_nulls_released_schema_rejected() { + #[test] + fn test_add_columns_nulls_non_utf8_format_rejected() { + // A non-NULL but non-UTF-8 top-level `format` must be rejected at the FFI +- // boundary rather than aborting via arrow-rs's `format().to_str().expect()` +- // under `panic = "abort"`. ++ // boundary rather than reaching arrow-rs's `format().to_str().expect()` ++ // and being downgraded from a precise InvalidArgument to Panic. + let (_tmp, uri) = create_large_dataset(2); + let c_uri = c_str(&uri); + let ds = unsafe { lance_dataset_open(c_uri.as_ptr(), ptr::null(), 0) }; +diff --git a/tests/cpp/test_cpp_api.cpp b/tests/cpp/test_cpp_api.cpp +index 3293bfb..b5d090a 100644 +--- a/tests/cpp/test_cpp_api.cpp ++++ b/tests/cpp/test_cpp_api.cpp +@@ -12,8 +12,11 @@ + + #include "lance/lance.hpp" + #include ++#include ++#include + #include + #include ++#include + #include + #include + #include +@@ -44,6 +47,29 @@ static void capture_scan_statistics( + captured->bytes_read = statistics->bytes_read; + } + ++struct AsyncScanCapture { ++ std::mutex mutex; ++ std::condition_variable ready; ++ bool completed = false; ++ int32_t status = -1; ++ ArrowArrayStream* stream = nullptr; ++}; ++ ++static void capture_async_scan( ++ void* callback_ctx, ++ int32_t status, ++ void* result) noexcept { ++ if (!callback_ctx) return; ++ auto* captured = static_cast(callback_ctx); ++ { ++ std::lock_guard lock(captured->mutex); ++ captured->status = status; ++ captured->stream = static_cast(result); ++ captured->completed = true; ++ } ++ captured->ready.notify_one(); ++} ++ + static void test_dataset_open(const std::string& uri) { + TEST(test_dataset_open); + +@@ -121,6 +147,47 @@ static void test_scanner_fluent(const std::string& uri) { + PASS(); + } + ++static void test_scanner_async_stream_ownership(const std::string& uri) { ++ TEST(test_scanner_async_stream_ownership); ++ ++ auto ds = lance::Dataset::open(uri); ++ auto scanner = ds.scan(); ++ AsyncScanCapture captured; ++ scanner.scan_async(capture_async_scan, &captured); ++ ++ ArrowArrayStream* stream = nullptr; ++ { ++ std::unique_lock lock(captured.mutex); ++ bool completed = captured.ready.wait_for( ++ lock, std::chrono::seconds(30), [&captured] { ++ return captured.completed; ++ }); ++ assert(completed && "async scan callback timed out"); ++ assert(captured.status == 0); ++ assert(captured.stream != nullptr); ++ stream = captured.stream; ++ } ++ ++ uint64_t total = 0; ++ while (true) { ++ ArrowArray array; ++ memset(&array, 0, sizeof(array)); ++ int rc = stream->get_next(stream, &array); ++ assert(rc == 0); ++ if (!array.release) break; ++ total += static_cast(array.length); ++ array.release(&array); ++ } ++ assert(total > 0); ++ ++ // This releases the stream contents (if still live) and the separate ++ // library-allocated outer structure. It is also explicitly NULL-safe. ++ lance::scanner_async_stream_free(stream); ++ lance::scanner_async_stream_free(nullptr); ++ ++ PASS(); ++} ++ + static void test_dataset_take(const std::string& uri) { + TEST(test_dataset_take); + +@@ -194,6 +261,15 @@ static void test_raii_cleanup(const std::string& uri) { + auto ds1 = lance::Dataset::open(uri); + auto ds2 = std::move(ds1); + assert(ds2.count_rows() > 0); ++ ++ bool moved_from_version_threw = false; ++ try { ++ (void)ds1.version(); ++ } catch (const lance::Error& e) { ++ moved_from_version_threw = true; ++ assert(e.code == LANCE_ERR_INVALID_ARGUMENT); ++ } ++ assert(moved_from_version_threw); + } + + PASS(); +@@ -832,6 +908,7 @@ int main(int argc, char** argv) { + test_dataset_open(uri); + test_dataset_schema(uri); + test_scanner_fluent(uri); ++ test_scanner_async_stream_ownership(uri); + test_dataset_take(uri); + test_dataset_take_rows(uri); + test_raii_cleanup(uri); +diff --git a/tests/panic_stream_guard.rs b/tests/panic_stream_guard.rs +index f4690c5..12efd0d 100644 +--- a/tests/panic_stream_guard.rs ++++ b/tests/panic_stream_guard.rs +@@ -40,8 +40,16 @@ + //! host. The guard's `Drop` detaches the inner stream and contains + //! cleanup. Runs in a child process, asserting a clean exit AND that + //! the destructor panic really fired (caught). ++//! ++//! 5. Regular errors containing NUL are sanitized before arrow-rs formats ++//! them inside `get_next`. ++//! 6. A panic from an external error's `Display` is caught, reported as one ++//! terminal stream error, and poisons the scanner. ++//! 7. An unexportable schema is rejected while still inside the Rust guard, ++//! before arrow-rs's non-unwinding `get_schema` callback is exposed. + + use std::ffi::CStr; ++use std::panic::{AssertUnwindSafe, catch_unwind}; + use std::pin::Pin; + use std::sync::Arc; + use std::sync::atomic::{AtomicBool, Ordering}; +@@ -49,7 +57,7 @@ use std::task::{Context, Poll}; + + use arrow::array::{Int32Array, RecordBatch}; + use arrow::datatypes::{DataType, Field, Schema, SchemaRef}; +-use arrow::ffi::FFI_ArrowArray; ++use arrow::ffi::{FFI_ArrowArray, FFI_ArrowSchema}; + use arrow::ffi_stream::FFI_ArrowArrayStream; + use futures::Stream; + use lance_c::stream_guard::GuardedReader; +@@ -91,6 +99,17 @@ impl Stream for PanicOnSecondPoll { + /// destructor reached from the Arrow C `release` callback. + struct PanicOnDrop; + ++#[derive(Debug)] ++struct PanickingDisplay; ++ ++impl std::fmt::Display for PanickingDisplay { ++ fn fmt(&self, _f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { ++ panic!("simulated panic while formatting a stream error") ++ } ++} ++ ++impl std::error::Error for PanickingDisplay {} ++ + impl Stream for PanicOnDrop { + type Item = lance_core::Result; + +@@ -112,6 +131,11 @@ unsafe fn c_get_next(stream: *mut FFI_ArrowArrayStream, array: *mut FFI_ArrowArr + unsafe { get_next(stream, array) } + } + ++unsafe fn c_get_schema(stream: *mut FFI_ArrowArrayStream, schema: *mut FFI_ArrowSchema) -> i32 { ++ let get_schema = unsafe { (*stream).get_schema }.expect("get_schema callback is NULL"); ++ unsafe { get_schema(stream, schema) } ++} ++ + unsafe fn c_get_last_error(stream: *mut FFI_ArrowArrayStream) -> Option { + let get_last_error = + unsafe { (*stream).get_last_error }.expect("get_last_error callback is NULL"); +@@ -241,6 +265,120 @@ fn guarded_stream_maps_panic_to_c_stream_error() { + ); + } + ++#[test] ++fn guarded_stream_sanitizes_nul_in_regular_error() { ++ if std::env::var("POC_CHILD_NUL_ERROR").is_err() { ++ let output = run_child( ++ "guarded_stream_sanitizes_nul_in_regular_error", ++ "POC_CHILD_NUL_ERROR", ++ ); ++ let stderr = String::from_utf8_lossy(&output.stderr); ++ assert!( ++ output.status.success(), ++ "guarded child must exit cleanly, got status {:?}\nstderr:\n{stderr}", ++ output.status ++ ); ++ return; ++ } ++ ++ let scanner_poison = Arc::new(AtomicBool::new(false)); ++ let stream = futures::stream::iter(vec![Err(lance_core::Error::invalid_input_source( ++ "ordinary error with a NUL: bo\0om".into(), ++ ))]); ++ let (_rt, mut ffi) = guarded_export(stream, Arc::clone(&scanner_poison)); ++ let mut array = FFI_ArrowArray::empty(); ++ ++ let rc = unsafe { c_get_next(&mut ffi, &mut array) }; ++ assert_ne!(rc, 0, "the ordinary stream error must reach Arrow C"); ++ let msg = unsafe { c_get_last_error(&mut ffi) }.expect("get_last_error returned NULL"); ++ assert!(msg.contains("bo\\0om"), "NUL must be escaped, got: {msg:?}"); ++ assert!( ++ !scanner_poison.load(Ordering::SeqCst), ++ "an ordinary stream error must not poison the scanner" ++ ); ++} ++ ++#[test] ++fn guarded_stream_catches_panicking_error_display() { ++ if std::env::var("POC_CHILD_DISPLAY_ERROR").is_err() { ++ let output = run_child( ++ "guarded_stream_catches_panicking_error_display", ++ "POC_CHILD_DISPLAY_ERROR", ++ ); ++ let stderr = String::from_utf8_lossy(&output.stderr); ++ assert!( ++ output.status.success(), ++ "guarded child must exit cleanly, got status {:?}\nstderr:\n{stderr}", ++ output.status ++ ); ++ assert!( ++ stderr.contains("simulated panic while formatting a stream error"), ++ "the formatting panic must have fired and been caught\nstderr:\n{stderr}" ++ ); ++ return; ++ } ++ ++ let scanner_poison = Arc::new(AtomicBool::new(false)); ++ let stream = futures::stream::iter(vec![Err(lance_core::Error::invalid_input_source( ++ Box::new(PanickingDisplay), ++ ))]); ++ let (_rt, mut ffi) = guarded_export(stream, Arc::clone(&scanner_poison)); ++ let mut array = FFI_ArrowArray::empty(); ++ ++ let rc = unsafe { c_get_next(&mut ffi, &mut array) }; ++ assert_ne!(rc, 0, "the caught panic must reach Arrow C as an error"); ++ let msg = unsafe { c_get_last_error(&mut ffi) }.expect("get_last_error returned NULL"); ++ assert!( ++ msg.contains("simulated panic while formatting a stream error"), ++ "panic message should propagate to get_last_error, got: {msg}" ++ ); ++ assert!( ++ scanner_poison.load(Ordering::SeqCst), ++ "a formatting panic must poison the owning scanner" ++ ); ++} ++ ++#[test] ++fn guarded_stream_rejects_nul_schema_before_arrow_callback() { ++ if std::env::var("POC_CHILD_NUL_SCHEMA").is_err() { ++ let output = run_child( ++ "guarded_stream_rejects_nul_schema_before_arrow_callback", ++ "POC_CHILD_NUL_SCHEMA", ++ ); ++ let stderr = String::from_utf8_lossy(&output.stderr); ++ assert!( ++ output.status.success(), ++ "schema validation must fail before Arrow's callback can abort, got status {:?}\nstderr:\n{stderr}", ++ output.status ++ ); ++ return; ++ } ++ ++ let schema = Arc::new(Schema::new(vec![Field::new( ++ "field\0name", ++ DataType::Int32, ++ false, ++ )])); ++ let rt = tokio::runtime::Runtime::new().unwrap(); ++ let scanner_poison = Arc::new(AtomicBool::new(false)); ++ let outcome = catch_unwind(AssertUnwindSafe(|| { ++ let reader = GuardedReader::new( ++ futures::stream::empty::>(), ++ schema, ++ rt.handle().clone(), ++ scanner_poison, ++ ); ++ let mut ffi = FFI_ArrowArrayStream::new(Box::new(reader)); ++ let mut ffi_schema = FFI_ArrowSchema::empty(); ++ let rc = unsafe { c_get_schema(&mut ffi, &mut ffi_schema) }; ++ panic!("invalid schema reached Arrow callback and returned rc={rc}"); ++ })); ++ assert!( ++ outcome.is_err(), ++ "invalid schema must be rejected while the Rust FFI guard can still catch it" ++ ); ++} ++ + /// A `get_next` call made from a thread that is currently driving a Tokio + /// runtime (inside `Runtime::block_on` or a spawned task — a merely + /// `enter()`ed context does not trip tokio's check) makes `Handle::block_on`