diff --git a/be/src/exprs/function/ai/ai_adapter.h b/be/src/exprs/function/ai/ai_adapter.h index b83aa26c51a857..39f16e73420472 100644 --- a/be/src/exprs/function/ai/ai_adapter.h +++ b/be/src/exprs/function/ai/ai_adapter.h @@ -1568,6 +1568,14 @@ class AnthropicAdapter : public VoyageAIAdapter { // Mock adapter used only for UT to bypass real HTTP calls and return deterministic data. class MockAdapter : public AIAdapter { public: +#ifdef BE_TEST + static void clear_embedding_inputs_for_test() { _embedding_inputs_for_test().clear(); } + + static const std::vector& get_embedding_inputs_for_test() { + return _embedding_inputs_for_test(); + } +#endif + Status set_authentication(HttpClient* client) const override { return Status::OK(); } Status build_request_payload(const std::vector& inputs, @@ -1583,6 +1591,10 @@ class MockAdapter : public AIAdapter { Status build_embedding_request(const std::vector& inputs, std::string& request_body) const override { +#ifdef BE_TEST + auto& embedding_inputs = _embedding_inputs_for_test(); + embedding_inputs.insert(embedding_inputs.end(), inputs.begin(), inputs.end()); +#endif return Status::OK(); } @@ -1612,6 +1624,14 @@ class MockAdapter : public AIAdapter { [](const auto& val) { return val.GetFloat(); }); return Status::OK(); } + +private: +#ifdef BE_TEST + static std::vector& _embedding_inputs_for_test() { + static thread_local std::vector embedding_inputs; + return embedding_inputs; + } +#endif }; class AIAdapterFactory { diff --git a/be/src/exprs/function/ai/ai_classify.h b/be/src/exprs/function/ai/ai_classify.h index 58048a1ed805ae..db21a3fd352c12 100644 --- a/be/src/exprs/function/ai/ai_classify.h +++ b/be/src/exprs/function/ai/ai_classify.h @@ -37,13 +37,13 @@ class FunctionAIClassify : public AIFunction { static constexpr size_t number_of_arguments = 3; - DataTypePtr get_return_type_impl(const DataTypes& arguments) const override { + DataTypePtr get_nested_return_type_impl(const DataTypes& /*arguments*/) const { return std::make_shared(); } static FunctionPtr create() { return std::make_shared(); } - Status build_prompt(const Block& block, const ColumnNumbers& arguments, size_t row_num, + Status build_prompt(const Columns& prompt_columns, size_t row_num, std::string& prompt) const override; }; } // namespace doris \ No newline at end of file diff --git a/be/src/exprs/function/ai/ai_extract.h b/be/src/exprs/function/ai/ai_extract.h index d2564554d82623..bca4a7319ec3ad 100644 --- a/be/src/exprs/function/ai/ai_extract.h +++ b/be/src/exprs/function/ai/ai_extract.h @@ -38,13 +38,13 @@ class FunctionAIExtract : public AIFunction { static constexpr size_t number_of_arguments = 3; - DataTypePtr get_return_type_impl(const DataTypes& arguments) const override { + DataTypePtr get_nested_return_type_impl(const DataTypes& /*arguments*/) const { return std::make_shared(); } static FunctionPtr create() { return std::make_shared(); } - Status build_prompt(const Block& block, const ColumnNumbers& arguments, size_t row_num, + Status build_prompt(const Columns& prompt_columns, size_t row_num, std::string& prompt) const override; }; diff --git a/be/src/exprs/function/ai/ai_filter.h b/be/src/exprs/function/ai/ai_filter.h index 6d6962e81dd62d..e92c5991405f49 100644 --- a/be/src/exprs/function/ai/ai_filter.h +++ b/be/src/exprs/function/ai/ai_filter.h @@ -38,7 +38,7 @@ class FunctionAIFilter : public AIFunction { static constexpr size_t number_of_arguments = 2; - DataTypePtr get_return_type_impl(const DataTypes& arguments) const override { + DataTypePtr get_nested_return_type_impl(const DataTypes& /*arguments*/) const { return std::make_shared(); } diff --git a/be/src/exprs/function/ai/ai_fix_grammar.h b/be/src/exprs/function/ai/ai_fix_grammar.h index 43f9d7a639481c..50ad2deb918db1 100644 --- a/be/src/exprs/function/ai/ai_fix_grammar.h +++ b/be/src/exprs/function/ai/ai_fix_grammar.h @@ -38,7 +38,7 @@ class FunctionAIFixGrammar : public AIFunction { static constexpr size_t number_of_arguments = 2; - DataTypePtr get_return_type_impl(const DataTypes& arguments) const override { + DataTypePtr get_nested_return_type_impl(const DataTypes& /*arguments*/) const { return std::make_shared(); } diff --git a/be/src/exprs/function/ai/ai_functions.cpp b/be/src/exprs/function/ai/ai_functions.cpp index ce6111f1fa47f8..4a06c3e5040652 100644 --- a/be/src/exprs/function/ai/ai_functions.cpp +++ b/be/src/exprs/function/ai/ai_functions.cpp @@ -30,17 +30,15 @@ #include "exprs/function/simple_function_factory.h" namespace doris { -Status FunctionAIClassify::build_prompt(const Block& block, const ColumnNumbers& arguments, - size_t row_num, std::string& prompt) const { +Status FunctionAIClassify::build_prompt(const Columns& prompt_columns, size_t row_num, + std::string& prompt) const { // Get the text column - const ColumnWithTypeAndName& text_column = block.get_by_position(arguments[1]); - StringRef text = text_column.column->get_data_at(row_num); + StringRef text = prompt_columns[0]->get_data_at(row_num); std::string text_str = std::string(text.data, text.size); // Get the labels array column - const ColumnWithTypeAndName& labels_column = block.get_by_position(arguments[2]); const auto& [array_column, array_row_num] = - check_column_const_set_readability(*labels_column.column, row_num); + check_column_const_set_readability(*prompt_columns[1], row_num); const auto* col_array = check_and_get_column(*array_column); if (col_array == nullptr) { return Status::InternalError( @@ -72,17 +70,15 @@ Status FunctionAIClassify::build_prompt(const Block& block, const ColumnNumbers& return Status::OK(); } -Status FunctionAIExtract::build_prompt(const Block& block, const ColumnNumbers& arguments, - size_t row_num, std::string& prompt) const { +Status FunctionAIExtract::build_prompt(const Columns& prompt_columns, size_t row_num, + std::string& prompt) const { // Get the text column - const ColumnWithTypeAndName& text_column = block.get_by_position(arguments[1]); - StringRef text = text_column.column->get_data_at(row_num); + StringRef text = prompt_columns[0]->get_data_at(row_num); std::string text_str = std::string(text.data, text.size); // Get the labels array column - const ColumnWithTypeAndName& labels_column = block.get_by_position(arguments[2]); const auto& [array_column, array_row_num] = - check_column_const_set_readability(*labels_column.column, row_num); + check_column_const_set_readability(*prompt_columns[1], row_num); const auto* col_array = check_and_get_column(*array_column); if (col_array == nullptr) { return Status::InternalError( @@ -114,26 +110,23 @@ Status FunctionAIExtract::build_prompt(const Block& block, const ColumnNumbers& return Status::OK(); } -Status FunctionAIGenerate::build_prompt(const Block& block, const ColumnNumbers& arguments, - size_t row_num, std::string& prompt) const { - const ColumnWithTypeAndName& text_column = block.get_by_position(arguments[1]); - StringRef text_ref = text_column.column->get_data_at(row_num); +Status FunctionAIGenerate::build_prompt(const Columns& prompt_columns, size_t row_num, + std::string& prompt) const { + StringRef text_ref = prompt_columns[0]->get_data_at(row_num); prompt = std::string(text_ref.data, text_ref.size); return Status::OK(); } -Status FunctionAIMask::build_prompt(const Block& block, const ColumnNumbers& arguments, - size_t row_num, std::string& prompt) const { +Status FunctionAIMask::build_prompt(const Columns& prompt_columns, size_t row_num, + std::string& prompt) const { // Get the text column - const ColumnWithTypeAndName& text_column = block.get_by_position(arguments[1]); - StringRef text = text_column.column->get_data_at(row_num); + StringRef text = prompt_columns[0]->get_data_at(row_num); std::string text_str = std::string(text.data, text.size); // Get the labels array column - const ColumnWithTypeAndName& labels_column = block.get_by_position(arguments[2]); const auto& [array_column, array_row_num] = - check_column_const_set_readability(*labels_column.column, row_num); + check_column_const_set_readability(*prompt_columns[1], row_num); const auto* col_array = check_and_get_column(*array_column); if (col_array == nullptr) { return Status::InternalError( @@ -165,16 +158,14 @@ Status FunctionAIMask::build_prompt(const Block& block, const ColumnNumbers& arg return Status::OK(); } -Status FunctionAISimilarity::build_prompt(const Block& block, const ColumnNumbers& arguments, - size_t row_num, std::string& prompt) const { +Status FunctionAISimilarity::build_prompt(const Columns& prompt_columns, size_t row_num, + std::string& prompt) const { // text1 - const ColumnWithTypeAndName& text_column_1 = block.get_by_position(arguments[1]); - StringRef text_1 = text_column_1.column.get()->get_data_at(row_num); + StringRef text_1 = prompt_columns[0]->get_data_at(row_num); std::string text_str_1 = std::string(text_1.data, text_1.size); // text2 - const ColumnWithTypeAndName& text_column_2 = block.get_by_position(arguments[2]); - StringRef text_2 = text_column_2.column.get()->get_data_at(row_num); + StringRef text_2 = prompt_columns[1]->get_data_at(row_num); std::string text_str_2 = std::string(text_2.data, text_2.size); prompt = "Text 1: " + text_str_1 + "\nText 2: " + text_str_2; @@ -182,16 +173,14 @@ Status FunctionAISimilarity::build_prompt(const Block& block, const ColumnNumber return Status::OK(); } -Status FunctionAITranslate::build_prompt(const Block& block, const ColumnNumbers& arguments, - size_t row_num, std::string& prompt) const { +Status FunctionAITranslate::build_prompt(const Columns& prompt_columns, size_t row_num, + std::string& prompt) const { // text - const ColumnWithTypeAndName& text_column = block.get_by_position(arguments[1]); - StringRef text = text_column.column.get()->get_data_at(row_num); + StringRef text = prompt_columns[0]->get_data_at(row_num); std::string text_str = std::string(text.data, text.size); // target language - const ColumnWithTypeAndName& lang_column = block.get_by_position(arguments[2]); - StringRef lang = lang_column.column.get()->get_data_at(row_num); + StringRef lang = prompt_columns[1]->get_data_at(row_num); std::string target_lang = std::string(lang.data, lang.size); prompt = "Translate the following text to " + target_lang + ".\nText: " + text_str; diff --git a/be/src/exprs/function/ai/ai_functions.h b/be/src/exprs/function/ai/ai_functions.h index db8d0245e52993..b074b6cc58473d 100644 --- a/be/src/exprs/function/ai/ai_functions.h +++ b/be/src/exprs/function/ai/ai_functions.h @@ -36,9 +36,11 @@ #include "core/column/column_nullable.h" #include "core/cow.h" #include "core/data_type/data_type_array.h" +#include "core/data_type/data_type_nullable.h" #include "core/data_type/data_type_number.h" #include "core/data_type/define_primitive_type.h" #include "core/data_type/primitive_type.h" +#include "exec/common/util.hpp" #include "exprs/function/ai/ai_adapter.h" #include "exprs/function/function.h" #include "runtime/query_context.h" @@ -65,10 +67,21 @@ class AIFunction : public IFunction { bool is_blockable() const override { return true; } - virtual Status build_prompt(const Block& block, const ColumnNumbers& arguments, size_t row_num, + bool use_default_implementation_for_nulls() const final { return false; } + + DataTypePtr get_return_type_impl(const DataTypes& arguments) const final { + bool has_nullable_argument = std::ranges::any_of( + arguments, [](const auto& argument) { return argument->is_nullable(); }); + DataTypePtr return_type = + assert_cast(*this).get_nested_return_type_impl(arguments); + return has_nullable_argument ? make_nullable(return_type) : return_type; + } + + using PreparedFunctionImpl::execute; + + virtual Status build_prompt(const Columns& prompt_columns, size_t row_num, std::string& prompt) const { - const ColumnWithTypeAndName& text_column = block.get_by_position(arguments[1]); - StringRef text_ref = text_column.column->get_data_at(row_num); + StringRef text_ref = prompt_columns[0]->get_data_at(row_num); prompt = std::string(text_ref.data, text_ref.size); return Status::OK(); @@ -76,6 +89,13 @@ class AIFunction : public IFunction { Status execute_impl(FunctionContext* context, Block& block, const ColumnNumbers& arguments, uint32_t result, size_t input_rows_count) const override { + if (block.get_by_position(arguments[0]).column->only_null()) { + block.get_by_position(result).column = + block.get_by_position(result).type->create_column_const(input_rows_count, + Field()); + return Status::OK(); + } + TAIResource config; std::shared_ptr adapter; if (Status status = this->_init_from_resource(context, block, arguments, config, adapter); @@ -83,8 +103,8 @@ class AIFunction : public IFunction { return status; } - return assert_cast(*this).execute_with_adapter( - context, block, arguments, result, input_rows_count, config, adapter); + return assert_cast(*this).execute(context, block, arguments, result, + input_rows_count, config, adapter); } protected: @@ -98,20 +118,6 @@ class AIFunction : public IFunction { return query_ctx->query_options().ai_context_window_size; } - // Derived classes can override this method for non-text/default behavior. - // The base implementation handles all string-input/string-output batchable functions. - Status execute_with_adapter(FunctionContext* context, Block& block, - const ColumnNumbers& arguments, uint32_t result, - size_t input_rows_count, const TAIResource& config, - std::shared_ptr& adapter) const { - auto col_result = assert_cast(*this).create_result_column(); - RETURN_IF_ERROR(execute_batched_prompts(context, block, arguments, input_rows_count, config, - adapter, *col_result)); - - block.replace_by_position(result, std::move(col_result)); - return Status::OK(); - } - MutableColumnPtr create_result_column() const { return ColumnString::create(); } // Provider-reusable hook for AI functions(string) -> string. @@ -285,19 +291,50 @@ class AIFunction : public IFunction { // Provider-reusable helper for string-returning functions. // Runs the common batch execution flow; derived classes only need to define how one batch of // string results is inserted into the final output column. - Status execute_batched_prompts(FunctionContext* context, Block& block, - const ColumnNumbers& arguments, size_t input_rows_count, - const TAIResource& config, std::shared_ptr& adapter, - IColumn& col_result) const { + Status execute(FunctionContext* context, Block& block, const ColumnNumbers& arguments, + uint32_t result, size_t input_rows_count, const TAIResource& config, + std::shared_ptr& adapter) const { + Columns prompt_columns; + prompt_columns.reserve(arguments.size() - 1); + ColumnUInt8::MutablePtr result_null_map; + for (size_t i = 1; i < arguments.size(); ++i) { + const auto& argument = block.get_by_position(arguments[i]); + if (argument.type->is_nullable()) { + const auto& [column, is_const] = unpack_if_const(argument.column); + const auto& nullable = + assert_cast(*column); + if (!result_null_map) { + result_null_map = ColumnUInt8::create(input_rows_count, 0); + } + VectorizedUtils::update_null_map(result_null_map->get_data(), + nullable.get_null_map_data(), is_const); + } + prompt_columns.emplace_back(argument.unnest_nullable().column); + } + + if (result_null_map && + !simd::contain_zero(result_null_map->get_data().data(), input_rows_count)) { + block.get_by_position(result).column = + block.get_by_position(result).type->create_column_const(input_rows_count, + Field()); + return Status::OK(); + } + + auto col_result = assert_cast(*this).create_result_column(); std::vector batch_prompts; size_t current_batch_size = 2; // [] const size_t max_batch_prompt_size = static_cast(get_ai_context_window_size(context)); + const NullMap* null_map = result_null_map ? &result_null_map->get_data() : nullptr; for (size_t i = 0; i < input_rows_count; ++i) { + if (null_map && (*null_map)[i]) { + continue; + } + std::string prompt; RETURN_IF_ERROR( - assert_cast(*this).build_prompt(block, arguments, i, prompt)); + assert_cast(*this).build_prompt(prompt_columns, i, prompt)); size_t entry_size = estimate_batch_entry_size(batch_prompts.size(), prompt); if (entry_size > max_batch_prompt_size) { @@ -306,7 +343,7 @@ class AIFunction : public IFunction { RETURN_IF_ERROR(this->execute_batch_request(batch_prompts, batch_results, config, adapter, context)); RETURN_IF_ERROR(assert_cast(*this).append_batch_results( - batch_results, col_result)); + batch_results, *col_result)); batch_prompts.clear(); current_batch_size = 2; } @@ -317,7 +354,7 @@ class AIFunction : public IFunction { RETURN_IF_ERROR(this->execute_batch_request(single_prompts, single_results, config, adapter, context)); RETURN_IF_ERROR(assert_cast(*this).append_batch_results( - single_results, col_result)); + single_results, *col_result)); continue; } @@ -328,7 +365,7 @@ class AIFunction : public IFunction { RETURN_IF_ERROR(this->execute_batch_request(batch_prompts, batch_results, config, adapter, context)); RETURN_IF_ERROR(assert_cast(*this).append_batch_results( - batch_results, col_result)); + batch_results, *col_result)); batch_prompts.clear(); current_batch_size = 2; additional_size = entry_size; @@ -343,8 +380,32 @@ class AIFunction : public IFunction { RETURN_IF_ERROR(this->execute_batch_request(batch_prompts, batch_results, config, adapter, context)); RETURN_IF_ERROR(assert_cast(*this).append_batch_results(batch_results, - col_result)); + *col_result)); + } + + if (!result_null_map) { + block.replace_by_position(result, std::move(col_result)); + return Status::OK(); + } + + if (!simd::contain_one(result_null_map->get_data().data(), input_rows_count)) { + block.replace_by_position(result, ColumnNullable::create(std::move(col_result), + std::move(result_null_map))); + return Status::OK(); } + + auto nested_result = col_result->clone_empty(); + size_t result_row = 0; + for (UInt8 is_null : result_null_map->get_data()) { + if (is_null) { + nested_result->insert_default(); + } else { + nested_result->insert_from(*col_result, result_row++); + } + } + + block.replace_by_position(result, ColumnNullable::create(std::move(nested_result), + std::move(result_null_map))); return Status::OK(); } diff --git a/be/src/exprs/function/ai/ai_generate.h b/be/src/exprs/function/ai/ai_generate.h index e8960864e1f04e..10f12f2a2bea29 100644 --- a/be/src/exprs/function/ai/ai_generate.h +++ b/be/src/exprs/function/ai/ai_generate.h @@ -36,13 +36,13 @@ class FunctionAIGenerate : public AIFunction { static constexpr size_t number_of_arguments = 2; - DataTypePtr get_return_type_impl(const DataTypes& arguments) const override { + DataTypePtr get_nested_return_type_impl(const DataTypes& /*arguments*/) const { return std::make_shared(); } static FunctionPtr create() { return std::make_shared(); } - Status build_prompt(const Block& block, const ColumnNumbers& arguments, size_t row_num, + Status build_prompt(const Columns& prompt_columns, size_t row_num, std::string& prompt) const override; }; diff --git a/be/src/exprs/function/ai/ai_mask.h b/be/src/exprs/function/ai/ai_mask.h index 35077f78dfa0c8..1052c30f739e4e 100644 --- a/be/src/exprs/function/ai/ai_mask.h +++ b/be/src/exprs/function/ai/ai_mask.h @@ -37,13 +37,13 @@ class FunctionAIMask : public AIFunction { static constexpr size_t number_of_arguments = 3; - DataTypePtr get_return_type_impl(const DataTypes& arguments) const override { + DataTypePtr get_nested_return_type_impl(const DataTypes& /*arguments*/) const { return std::make_shared(); } static FunctionPtr create() { return std::make_shared(); } - Status build_prompt(const Block& block, const ColumnNumbers& arguments, size_t row_num, + Status build_prompt(const Columns& prompt_columns, size_t row_num, std::string& prompt) const override; }; diff --git a/be/src/exprs/function/ai/ai_sentiment.h b/be/src/exprs/function/ai/ai_sentiment.h index 8e50125b430d4f..5fce0e06b8dbe2 100644 --- a/be/src/exprs/function/ai/ai_sentiment.h +++ b/be/src/exprs/function/ai/ai_sentiment.h @@ -36,7 +36,7 @@ class FunctionAISentiment : public AIFunction { static constexpr size_t number_of_arguments = 2; - DataTypePtr get_return_type_impl(const DataTypes& arguments) const override { + DataTypePtr get_nested_return_type_impl(const DataTypes& /*arguments*/) const { return std::make_shared(); } diff --git a/be/src/exprs/function/ai/ai_similarity.h b/be/src/exprs/function/ai/ai_similarity.h index 55705b588b6691..93b588b04271af 100644 --- a/be/src/exprs/function/ai/ai_similarity.h +++ b/be/src/exprs/function/ai/ai_similarity.h @@ -41,13 +41,13 @@ class FunctionAISimilarity : public AIFunction { static constexpr size_t number_of_arguments = 3; - DataTypePtr get_return_type_impl(const DataTypes& arguments) const override { + DataTypePtr get_nested_return_type_impl(const DataTypes& /*arguments*/) const { return std::make_shared(); } static FunctionPtr create() { return std::make_shared(); } - Status build_prompt(const Block& block, const ColumnNumbers& arguments, size_t row_num, + Status build_prompt(const Columns& prompt_columns, size_t row_num, std::string& prompt) const override; private: diff --git a/be/src/exprs/function/ai/ai_summarize.h b/be/src/exprs/function/ai/ai_summarize.h index 23963968e9f424..ad1d5fede118e2 100644 --- a/be/src/exprs/function/ai/ai_summarize.h +++ b/be/src/exprs/function/ai/ai_summarize.h @@ -37,7 +37,7 @@ class FunctionAISummarize : public AIFunction { static constexpr size_t number_of_arguments = 2; - DataTypePtr get_return_type_impl(const DataTypes& arguments) const override { + DataTypePtr get_nested_return_type_impl(const DataTypes& /*arguments*/) const { return std::make_shared(); } diff --git a/be/src/exprs/function/ai/ai_translate.h b/be/src/exprs/function/ai/ai_translate.h index 2f6514c47a136a..8c5365b804e7c0 100644 --- a/be/src/exprs/function/ai/ai_translate.h +++ b/be/src/exprs/function/ai/ai_translate.h @@ -35,13 +35,13 @@ class FunctionAITranslate : public AIFunction { "corresponding item, with no explanation, markdown, or extra text."; static constexpr size_t number_of_arguments = 3; - DataTypePtr get_return_type_impl(const DataTypes& arguments) const override { + DataTypePtr get_nested_return_type_impl(const DataTypes& /*arguments*/) const { return std::make_shared(); } static FunctionPtr create() { return std::make_shared(); } - Status build_prompt(const Block& block, const ColumnNumbers& arguments, size_t row_num, + Status build_prompt(const Columns& prompt_columns, size_t row_num, std::string& prompt) const override; }; diff --git a/be/src/exprs/function/ai/embed.h b/be/src/exprs/function/ai/embed.h index 2367e4b9459540..f193a1c171c1b5 100644 --- a/be/src/exprs/function/ai/embed.h +++ b/be/src/exprs/function/ai/embed.h @@ -38,33 +38,53 @@ class FunctionEmbed : public AIFunction { static constexpr auto system_prompt = ""; - DataTypePtr get_return_type_impl(const DataTypes& arguments) const override { + DataTypePtr get_nested_return_type_impl(const DataTypes& /*arguments*/) const { return std::make_shared(make_nullable(std::make_shared())); } - Status execute_with_adapter(FunctionContext* context, Block& block, - const ColumnNumbers& arguments, uint32_t result, - size_t input_rows_count, const TAIResource& config, - std::shared_ptr& adapter) const { + using PreparedFunctionImpl::execute; + + Status execute(FunctionContext* context, Block& block, const ColumnNumbers& arguments, + uint32_t result, size_t input_rows_count, const TAIResource& config, + std::shared_ptr& adapter) const { if (arguments.size() != 2) { return Status::InvalidArgument("Function EMBED expects 2 arguments, but got {}", arguments.size()); } - PrimitiveType input_type = - remove_nullable(block.get_by_position(arguments[1]).type)->get_primitive_type(); + const auto& input = block.get_by_position(arguments[1]); + ColumnUInt8::MutablePtr result_null_map; + if (input.type->is_nullable()) { + const auto& [column, is_const] = unpack_if_const(input.column); + const auto& nullable = + assert_cast(*column); + result_null_map = ColumnUInt8::create(input_rows_count, 0); + VectorizedUtils::update_null_map(result_null_map->get_data(), + nullable.get_null_map_data(), is_const); + } + + if (result_null_map && + !simd::contain_zero(result_null_map->get_data().data(), input_rows_count)) { + block.get_by_position(result).column = + block.get_by_position(result).type->create_column_const(input_rows_count, + Field()); + return Status::OK(); + } + + ColumnPtr input_column = input.unnest_nullable().column; + PrimitiveType input_type = remove_nullable(input.type)->get_primitive_type(); if (input_type == PrimitiveType::TYPE_JSONB) { - return _execute_multimodal_embed(context, block, arguments, result, input_rows_count, - config, adapter); + return _execute_multimodal_embed(context, block, result, input_rows_count, config, + adapter, input_column, std::move(result_null_map)); } if (input_type == PrimitiveType::TYPE_STRING || input_type == PrimitiveType::TYPE_VARCHAR || input_type == PrimitiveType::TYPE_CHAR) { - return _execute_text_embed(context, block, arguments, result, input_rows_count, config, - adapter); + return _execute_text_embed(context, block, result, input_rows_count, config, adapter, + input_column, std::move(result_null_map)); } return Status::InvalidArgument( "Function EMBED expects the second argument to be STRING or JSON, but got type {}", - block.get_by_position(arguments[1]).type->get_name()); + input.type->get_name()); } static FunctionPtr create() { return std::make_shared(); } @@ -77,10 +97,10 @@ class FunctionEmbed : public AIFunction { return query_ctx->query_options().embed_max_batch_size; } - Status _execute_text_embed(FunctionContext* context, Block& block, - const ColumnNumbers& arguments, uint32_t result, + Status _execute_text_embed(FunctionContext* context, Block& block, uint32_t result, size_t input_rows_count, const TAIResource& config, - std::shared_ptr& adapter) const { + std::shared_ptr& adapter, const ColumnPtr& input_column, + ColumnUInt8::MutablePtr result_null_map) const { auto col_result = ColumnArray::create( ColumnNullable::create(ColumnFloat32::create(), ColumnUInt8::create())); std::vector batch_prompts; @@ -88,10 +108,16 @@ class FunctionEmbed : public AIFunction { const int32_t max_batch_size = _get_embed_max_batch_size(context); const size_t max_context_window_size = static_cast(get_ai_context_window_size(context)); + const NullMap* null_map = result_null_map ? &result_null_map->get_data() : nullptr; + const Columns prompt_columns {input_column}; for (size_t i = 0; i < input_rows_count; ++i) { + if (null_map && (*null_map)[i]) { + continue; + } + std::string prompt; - RETURN_IF_ERROR(build_prompt(block, arguments, i, prompt)); + RETURN_IF_ERROR(build_prompt(prompt_columns, i, prompt)); const size_t prompt_size = prompt.size(); @@ -122,19 +148,23 @@ class FunctionEmbed : public AIFunction { RETURN_IF_ERROR( _flush_text_embedding_batch(batch_prompts, *col_result, config, adapter, context)); - block.replace_by_position(result, std::move(col_result)); + block.replace_by_position(result, _expand_and_wrap_nullable_result( + std::move(col_result), std::move(result_null_map), + input_rows_count)); return Status::OK(); } - Status _execute_multimodal_embed(FunctionContext* context, Block& block, - const ColumnNumbers& arguments, uint32_t result, + Status _execute_multimodal_embed(FunctionContext* context, Block& block, uint32_t result, size_t input_rows_count, const TAIResource& config, - std::shared_ptr& adapter) const { + std::shared_ptr& adapter, + const ColumnPtr& input_column, + ColumnUInt8::MutablePtr result_null_map) const { auto col_result = ColumnArray::create( ColumnNullable::create(ColumnFloat32::create(), ColumnUInt8::create())); std::vector batch_media_types; std::vector batch_media_content_types; std::vector batch_media_urls; + const NullMap* null_map = result_null_map ? &result_null_map->get_data() : nullptr; int64_t ttl_seconds = 3600; QueryContext* query_ctx = context->state()->get_query_ctx(); @@ -147,10 +177,13 @@ class FunctionEmbed : public AIFunction { const int32_t max_batch_size = _get_embed_max_batch_size(context); - const ColumnWithTypeAndName& file_column = block.get_by_position(arguments[1]); for (size_t i = 0; i < input_rows_count; ++i) { + if (null_map && (*null_map)[i]) { + continue; + } + rapidjson::Document file_input; - RETURN_IF_ERROR(_parse_file_input(file_column, i, file_input)); + RETURN_IF_ERROR(_parse_file_input(*input_column, i, file_input)); std::string content_type; MultimodalType media_type; @@ -175,7 +208,9 @@ class FunctionEmbed : public AIFunction { batch_media_types, batch_media_content_types, batch_media_urls, *col_result, config, adapter, context)); - block.replace_by_position(result, std::move(col_result)); + block.replace_by_position(result, _expand_and_wrap_nullable_result( + std::move(col_result), std::move(result_null_map), + input_rows_count)); return Status::OK(); } @@ -279,6 +314,29 @@ class FunctionEmbed : public AIFunction { null_map.insert_many_vals(0, float_result.size()); } + static ColumnPtr _expand_and_wrap_nullable_result(ColumnArray::MutablePtr result, + ColumnUInt8::MutablePtr result_null_map, + size_t input_rows_count) { + if (!result_null_map) { + return result; + } + + auto& offsets = result->get_offsets(); + size_t compact_row = offsets.size(); + offsets.resize(input_rows_count); + // For example, embedding rows 1 and 3 produces compact offsets [5, 10]. Given + // result_null_map [1, 0, 1, 0, 1], expand them to [0, 5, 5, 10, 10], where NULL rows + // reuse the previous offset. Fill backwards to avoid overwriting unread compact offsets. + for (size_t row = input_rows_count; row-- > 0;) { + if (result_null_map->get_data()[row]) { + offsets[row] = compact_row == 0 ? 0 : offsets[compact_row - 1]; + } else { + offsets[row] = offsets[--compact_row]; + } + } + return ColumnNullable::create(std::move(result), std::move(result_null_map)); + } + static bool _starts_with_ignore_case(std::string_view s, std::string_view prefix) { if (s.size() < prefix.size()) { return false; @@ -308,11 +366,10 @@ class FunctionEmbed : public AIFunction { } // Parse the FILE-like JSONB argument into a JSON object for downstream field reads. - static Status _parse_file_input(const ColumnWithTypeAndName& file_column, size_t row_num, + static Status _parse_file_input(const IColumn& file_column, size_t row_num, rapidjson::Document& file_input) { - std::string file_json = - JsonbToJson::jsonb_to_json_string(file_column.column->get_data_at(row_num).data, - file_column.column->get_data_at(row_num).size); + StringRef file_ref = file_column.get_data_at(row_num); + std::string file_json = JsonbToJson::jsonb_to_json_string(file_ref.data, file_ref.size); file_input.Parse(file_json.c_str()); DORIS_CHECK(!file_input.HasParseError() && file_input.IsObject()); return Status::OK(); diff --git a/be/test/ai/ai_function_test.cpp b/be/test/ai/ai_function_test.cpp index 23855611861733..f48f1e8518362f 100644 --- a/be/test/ai/ai_function_test.cpp +++ b/be/test/ai/ai_function_test.cpp @@ -27,6 +27,7 @@ #include "core/block/block.h" #include "core/column/column_array.h" +#include "core/column/column_const.h" #include "core/column/column_nullable.h" #include "core/column/column_string.h" #include "core/column/column_vector.h" @@ -43,6 +44,7 @@ #include "exprs/function/ai/ai_summarize.h" #include "exprs/function/ai/ai_translate.h" #include "exprs/function/ai/embed.h" +#include "exprs/function/simple_function_factory.h" #include "testutil/column_helper.h" #include "testutil/mock/mock_runtime_state.h" @@ -63,7 +65,7 @@ class FunctionAIFilterBatchTestHelper : public AIFunction::execute_batch_request; - DataTypePtr get_return_type_impl(const DataTypes& arguments) const override { + DataTypePtr get_nested_return_type_impl(const DataTypes& /*arguments*/) const { return std::make_shared(); } @@ -202,6 +204,21 @@ MutableColumnPtr create_string_array_column(const std::vector(); + auto ctx = FunctionContext::create_context(runtime_state.get(), {}, {}); + setenv("AI_TEST_RESULT", R"(["answer-a","answer-c"])", 1); + + std::vector texts = {"unused-null", "text-a", "unused-null", "text-c", + "unused-null"}; + std::vector null_map = {1, 0, 1, 0, 1}; + Block block; + block.insert({ColumnHelper::create_column( + std::vector(texts.size(), "mock_resource")), + std::make_shared(), "resource"}); + block.insert({ColumnHelper::create_nullable_column(texts, null_map), + make_nullable(std::make_shared()), "text"}); + + auto return_type = make_nullable(std::make_shared()); + auto function = get_ai_function("ai_generate", block, return_type); + ASSERT_NE(function, nullptr); + EXPECT_TRUE(function->get_return_type()->equals(*return_type)); + + block.insert({nullptr, return_type, "result"}); + Status status = function->execute(ctx.get(), block, {0, 1}, 2, texts.size()); + unsetenv("AI_TEST_RESULT"); + + ASSERT_TRUE(status.ok()) << status.to_string(); + const auto& result = assert_cast(*block.get_by_position(2).column); + const auto& nested = assert_cast(result.get_nested_column()); + ASSERT_EQ(result.size(), texts.size()); + for (size_t row = 0; row < null_map.size(); ++row) { + EXPECT_EQ(result.is_null_at(row), null_map[row] != 0); + } + EXPECT_EQ(nested.get_data_at(1).to_string(), "answer-a"); + EXPECT_EQ(nested.get_data_at(3).to_string(), "answer-c"); +} + +TEST(AIFunctionTest, NullableInputWithoutNullsThroughPreparedFunction) { + auto runtime_state = std::make_unique(); + auto ctx = FunctionContext::create_context(runtime_state.get(), {}, {}); + setenv("AI_TEST_RESULT", R"(["answer-a","answer-b","answer-c"])", 1); + + std::vector texts = {"text-a", "text-b", "text-c"}; + Block block; + block.insert({ColumnHelper::create_column( + std::vector(texts.size(), "mock_resource")), + std::make_shared(), "resource"}); + block.insert({ColumnHelper::create_nullable_column( + texts, std::vector(texts.size(), 0)), + make_nullable(std::make_shared()), "text"}); + + auto return_type = make_nullable(std::make_shared()); + auto function = get_ai_function("ai_generate", block, return_type); + ASSERT_NE(function, nullptr); + + block.insert({nullptr, return_type, "result"}); + Status status = function->execute(ctx.get(), block, {0, 1}, 2, texts.size()); + unsetenv("AI_TEST_RESULT"); + + ASSERT_TRUE(status.ok()) << status.to_string(); + const auto& result = assert_cast(*block.get_by_position(2).column); + const auto& nested = assert_cast(result.get_nested_column()); + ASSERT_EQ(result.size(), texts.size()); + for (size_t row = 0; row < texts.size(); ++row) { + EXPECT_FALSE(result.is_null_at(row)); + } + EXPECT_EQ(nested.get_data_at(0).to_string(), "answer-a"); + EXPECT_EQ(nested.get_data_at(1).to_string(), "answer-b"); + EXPECT_EQ(nested.get_data_at(2).to_string(), "answer-c"); +} + +TEST(AIFunctionTest, NullableBoolResultThroughPreparedFunction) { + auto runtime_state = std::make_unique(); + auto ctx = FunctionContext::create_context(runtime_state.get(), {}, {}); + setenv("AI_TEST_RESULT", R"(["1","0"])", 1); + + std::vector texts = {"unused-null", "valid", "unused-null", "invalid"}; + std::vector null_map = {1, 0, 1, 0}; + Block block; + block.insert({ColumnHelper::create_column( + std::vector(texts.size(), "mock_resource")), + std::make_shared(), "resource"}); + block.insert({ColumnHelper::create_nullable_column(texts, null_map), + make_nullable(std::make_shared()), "text"}); + + auto return_type = make_nullable(std::make_shared()); + auto function = get_ai_function("ai_filter", block, return_type); + ASSERT_NE(function, nullptr); + EXPECT_TRUE(function->get_return_type()->equals(*return_type)); + + block.insert({nullptr, return_type, "result"}); + Status status = function->execute(ctx.get(), block, {0, 1}, 2, texts.size()); + unsetenv("AI_TEST_RESULT"); + + ASSERT_TRUE(status.ok()) << status.to_string(); + const auto& result = assert_cast(*block.get_by_position(2).column); + const auto& nested = assert_cast(result.get_nested_column()); + EXPECT_TRUE(result.is_null_at(0)); + EXPECT_EQ(nested.get_element(1), 1); + EXPECT_TRUE(result.is_null_at(2)); + EXPECT_EQ(nested.get_element(3), 0); +} + +TEST(AIFunctionTest, NullableFloatResultMergesArgumentNullMaps) { + auto runtime_state = std::make_unique(); + auto ctx = FunctionContext::create_context(runtime_state.get(), {}, {}); + setenv("AI_TEST_RESULT", R"(["0.5","1.5"])", 1); + + std::vector text1 = {"left-a", "unused-null", "left-c", "left-d"}; + std::vector text2 = {"right-a", "right-b", "unused-null", "right-d"}; + std::vector null_map1 = {0, 1, 0, 0}; + std::vector null_map2 = {0, 0, 1, 0}; + Block block; + block.insert({ColumnHelper::create_column( + std::vector(text1.size(), "mock_resource")), + std::make_shared(), "resource"}); + block.insert({ColumnHelper::create_nullable_column(text1, null_map1), + make_nullable(std::make_shared()), "text1"}); + block.insert({ColumnHelper::create_nullable_column(text2, null_map2), + make_nullable(std::make_shared()), "text2"}); + + auto return_type = make_nullable(std::make_shared()); + auto function = get_ai_function("ai_similarity", block, return_type); + ASSERT_NE(function, nullptr); + EXPECT_TRUE(function->get_return_type()->equals(*return_type)); + + block.insert({nullptr, return_type, "result"}); + Status status = function->execute(ctx.get(), block, {0, 1, 2}, 3, text1.size()); + unsetenv("AI_TEST_RESULT"); + + ASSERT_TRUE(status.ok()) << status.to_string(); + const auto& result = assert_cast(*block.get_by_position(3).column); + const auto& nested = assert_cast(result.get_nested_column()); + EXPECT_FALSE(result.is_null_at(0)); + EXPECT_FLOAT_EQ(nested.get_element(0), 0.5f); + EXPECT_TRUE(result.is_null_at(1)); + EXPECT_TRUE(result.is_null_at(2)); + EXPECT_FALSE(result.is_null_at(3)); + EXPECT_FLOAT_EQ(nested.get_element(3), 1.5f); +} + +TEST(AIFunctionTest, NullableArrayArgumentThroughPreparedFunction) { + auto runtime_state = std::make_unique(); + auto ctx = FunctionContext::create_context(runtime_state.get(), {}, {}); + setenv("AI_TEST_RESULT", R"(["positive"])", 1); + + std::vector texts = {"unused-null", "good product", "unused-null"}; + std::vector labels_null_map = {1, 0, 1}; + auto labels = create_string_array_column({{}, {"positive", "negative"}, {}}); + Block block; + block.insert({ColumnHelper::create_column( + std::vector(texts.size(), "mock_resource")), + std::make_shared(), "resource"}); + block.insert({ColumnHelper::create_column(texts), + std::make_shared(), "text"}); + block.insert( + {ColumnNullable::create(std::move(labels), + ColumnHelper::create_column(labels_null_map)), + make_nullable(std::make_shared( + make_nullable(std::make_shared()))), + "labels"}); + + auto return_type = make_nullable(std::make_shared()); + auto function = get_ai_function("ai_classify", block, return_type); + ASSERT_NE(function, nullptr); + + block.insert({nullptr, return_type, "result"}); + Status status = function->execute(ctx.get(), block, {0, 1, 2}, 3, texts.size()); + unsetenv("AI_TEST_RESULT"); + + ASSERT_TRUE(status.ok()) << status.to_string(); + const auto& result = assert_cast(*block.get_by_position(3).column); + const auto& nested = assert_cast(result.get_nested_column()); + EXPECT_TRUE(result.is_null_at(0)); + EXPECT_EQ(nested.get_data_at(1).to_string(), "positive"); + EXPECT_TRUE(result.is_null_at(2)); +} + +TEST(AIFunctionTest, AllNullConstArgumentReturnsConstNull) { + auto runtime_state = std::make_unique(); + auto ctx = FunctionContext::create_context(runtime_state.get(), {}, {}); + constexpr size_t row_count = 5; + + auto resource = ColumnConst::create( + ColumnHelper::create_column({"mock_resource"}), row_count); + auto nullable_text = + ColumnHelper::create_nullable_column({""}, std::vector {1}); + auto text = ColumnConst::create(std::move(nullable_text), row_count); + Block block; + block.insert({std::move(resource), std::make_shared(), "resource"}); + block.insert({std::move(text), make_nullable(std::make_shared()), "text"}); + + auto return_type = make_nullable(std::make_shared()); + auto function = get_ai_function("ai_generate", block, return_type); + ASSERT_NE(function, nullptr); + + block.insert({nullptr, return_type, "result"}); + Status status = function->execute(ctx.get(), block, {0, 1}, 2, row_count); + + ASSERT_TRUE(status.ok()) << status.to_string(); + ASSERT_TRUE(is_column_const(*block.get_by_position(2).column)); + ColumnPtr full_result = block.get_by_position(2).column->convert_to_full_column_if_const(); + const auto& nullable_result = assert_cast(*full_result); + ASSERT_EQ(nullable_result.size(), row_count); + for (size_t row = 0; row < row_count; ++row) { + EXPECT_TRUE(nullable_result.is_null_at(row)); + } +} + +TEST(AIFunctionTest, NullResourceReturnsConstNullBeforeLookup) { + auto runtime_state = std::make_unique(); + auto ctx = FunctionContext::create_context(runtime_state.get(), {}, {}); + constexpr size_t row_count = 3; + + auto nullable_resource = + ColumnHelper::create_nullable_column({""}, std::vector {1}); + auto resource = ColumnConst::create(std::move(nullable_resource), row_count); + auto text = ColumnHelper::create_column( + std::vector(row_count, "prompt")); + Block block; + block.insert( + {std::move(resource), make_nullable(std::make_shared()), "resource"}); + block.insert({std::move(text), std::make_shared(), "text"}); + + auto return_type = make_nullable(std::make_shared()); + auto function = get_ai_function("ai_generate", block, return_type); + ASSERT_NE(function, nullptr); + + block.insert({nullptr, return_type, "result"}); + Status status = function->execute(ctx.get(), block, {0, 1}, 2, row_count); + + ASSERT_TRUE(status.ok()) << status.to_string(); + ASSERT_TRUE(is_column_const(*block.get_by_position(2).column)); + EXPECT_TRUE(block.get_by_position(2).column->only_null()); +} + TEST(AIFunctionTest, MissingAIResourcesMetadataTest) { auto query_ctx = MockQueryContext::create(); TQueryOptions query_options; diff --git a/be/test/ai/embed_test.cpp b/be/test/ai/embed_test.cpp index 2074697614252a..c9bd32ed17cb7a 100644 --- a/be/test/ai/embed_test.cpp +++ b/be/test/ai/embed_test.cpp @@ -26,10 +26,12 @@ #include #include +#include "core/column/column_const.h" #include "core/data_type/data_type_jsonb.h" #include "core/data_type/data_type_number.h" #include "core/value/jsonb_value.h" #include "exprs/function/ai/ai_adapter.h" +#include "exprs/function/simple_function_factory.h" #include "io/fs/obj_storage_client.h" #include "testutil/column_helper.h" #include "testutil/mock/mock_runtime_state.h" @@ -151,6 +153,25 @@ static ColumnString::MutablePtr create_jsonb_column(const std::vector& json_rows, + const std::vector& null_map) { + EXPECT_EQ(json_rows.size(), null_map.size()); + auto column = ColumnString::create(); + auto null_column = ColumnUInt8::create(); + for (size_t i = 0; i < json_rows.size(); ++i) { + if (null_map[i]) { + column->insert_default(); + } else { + JsonBinaryValue jsonb_value; + Status st = jsonb_value.from_json_string(json_rows[i]); + EXPECT_TRUE(st.ok()) << st.to_string(); + column->insert_data(jsonb_value.value(), jsonb_value.size()); + } + null_column->insert_value(null_map[i]); + } + return ColumnNullable::create(std::move(column), std::move(null_column)); +} + static void assert_mock_embedding_column(const ColumnArray& col_array, size_t row_count) { const auto& offsets = col_array.get_offsets(); ASSERT_EQ(offsets.size(), row_count); @@ -168,6 +189,39 @@ static void assert_mock_embedding_column(const ColumnArray& col_array, size_t ro } } +static void assert_mock_nullable_embedding_column(const IColumn& column, + const std::vector& expected_null_map) { + const auto& nullable_column = assert_cast(column); + ASSERT_EQ(nullable_column.size(), expected_null_map.size()); + + const auto& col_array = assert_cast(nullable_column.get_nested_column()); + const auto& offsets = col_array.get_offsets(); + const auto& nested_nullable_col = assert_cast(col_array.get_data()); + const auto& nested_col = + assert_cast(*nested_nullable_col.get_nested_column_ptr()); + + size_t expected_offset = 0; + for (size_t row = 0; row < expected_null_map.size(); ++row) { + ASSERT_EQ(nullable_column.is_null_at(row), expected_null_map[row] != 0); + if (expected_null_map[row]) { + ASSERT_EQ(offsets[row], expected_offset); + continue; + } + + expected_offset += 5; + ASSERT_EQ(offsets[row], expected_offset); + for (size_t i = 0; i < 5; ++i) { + ASSERT_FLOAT_EQ(nested_col.get_element(expected_offset - 5 + i), static_cast(i)); + } + } + ASSERT_EQ(nested_col.size(), expected_offset); +} + +static FunctionBasePtr get_embed_function(const Block& block, const DataTypePtr& return_type) { + return SimpleFunctionFactory::instance().get_function( + "embed", block.get_columns_with_type_and_name(), return_type); +} + TEST(EMBED_TEST, embed_function_build_test) { FunctionEmbed function; @@ -183,7 +237,7 @@ TEST(EMBED_TEST, embed_function_build_test) { ColumnNumbers arguments = {0, 1}; std::string prompt; - Status status = function.build_prompt(block, arguments, 0, prompt); + Status status = function.build_prompt({block.get_by_position(arguments[1]).column}, 0, prompt); ASSERT_TRUE(status.ok()); ASSERT_EQ(prompt, "this is a test prompt"); @@ -297,6 +351,133 @@ TEST(EMBED_TEST, embed_function_multimodal_direct_url) { assert_mock_embedding_column(col_array, file_json_rows.size()); } +TEST(EMBED_TEST, embed_function_partial_null_through_framework) { + auto runtime_state = std::make_unique(); + auto ctx = FunctionContext::create_context(runtime_state.get(), {}, {}); + + std::vector resources(5, "mock_resource"); + std::vector texts = {"", "text-a", "", "text-c", ""}; + std::vector null_map = {1, 0, 1, 0, 1}; + auto col_resource = ColumnHelper::create_column(resources); + auto col_text = ColumnHelper::create_nullable_column(texts, null_map); + auto return_type = make_nullable( + std::make_shared(make_nullable(std::make_shared()))); + + Block block; + block.insert({std::move(col_resource), std::make_shared(), "resource"}); + block.insert({std::move(col_text), make_nullable(std::make_shared()), "text"}); + + auto function = get_embed_function(block, return_type); + ASSERT_NE(function, nullptr); + EXPECT_TRUE(function->get_return_type()->equals(*return_type)); + + block.insert({nullptr, return_type, "result"}); + const size_t result_idx = 2; + MockAdapter::clear_embedding_inputs_for_test(); + Status exec_status = function->execute(ctx.get(), block, {0, 1}, result_idx, texts.size()); + + ASSERT_TRUE(exec_status.ok()) << exec_status.to_string(); + EXPECT_THAT(MockAdapter::get_embedding_inputs_for_test(), + ::testing::ElementsAre("text-a", "text-c")); + assert_mock_nullable_embedding_column(*block.get_by_position(result_idx).column, null_map); +} + +TEST(EMBED_TEST, embed_function_multimodal_partial_null_through_framework) { + auto runtime_state = std::make_unique(); + auto ctx = FunctionContext::create_context(runtime_state.get(), {}, {}); + + std::vector resources(5, "mock_resource"); + std::vector file_json_rows = { + "", R"({"content_type":"image/png","uri":"https://example.com/a.png"})", "", + R"({"content_type":"video/mp4","uri":"https://example.com/b.mp4"})", ""}; + std::vector null_map = {1, 0, 1, 0, 1}; + auto col_resource = ColumnHelper::create_column(resources); + auto col_file = create_nullable_jsonb_column(file_json_rows, null_map); + auto return_type = make_nullable( + std::make_shared(make_nullable(std::make_shared()))); + + Block block; + block.insert({std::move(col_resource), std::make_shared(), "resource"}); + block.insert({std::move(col_file), make_nullable(std::make_shared()), "file"}); + + auto function = get_embed_function(block, return_type); + ASSERT_NE(function, nullptr); + + block.insert({nullptr, return_type, "result"}); + const size_t result_idx = 2; + Status exec_status = + function->execute(ctx.get(), block, {0, 1}, result_idx, file_json_rows.size()); + + ASSERT_TRUE(exec_status.ok()) << exec_status.to_string(); + assert_mock_nullable_embedding_column(*block.get_by_position(result_idx).column, null_map); +} + +TEST(EMBED_TEST, embed_function_all_null_const_nullable_through_framework) { + auto runtime_state = std::make_unique(); + auto ctx = FunctionContext::create_context(runtime_state.get(), {}, {}); + + constexpr size_t row_count = 5; + std::vector resources(row_count, "mock_resource"); + auto col_resource = ColumnHelper::create_column(resources); + auto nullable_text = + ColumnHelper::create_nullable_column({""}, std::vector {1}); + auto col_text = ColumnConst::create(std::move(nullable_text), row_count); + auto return_type = make_nullable( + std::make_shared(make_nullable(std::make_shared()))); + + Block block; + block.insert({std::move(col_resource), std::make_shared(), "resource"}); + block.insert({std::move(col_text), make_nullable(std::make_shared()), "text"}); + + auto function = get_embed_function(block, return_type); + ASSERT_NE(function, nullptr); + + block.insert({nullptr, return_type, "result"}); + const size_t result_idx = 2; + Status exec_status = function->execute(ctx.get(), block, {0, 1}, result_idx, row_count); + + ASSERT_TRUE(exec_status.ok()) << exec_status.to_string(); + const auto& result_column = block.get_by_position(result_idx).column; + ASSERT_TRUE(is_column_const(*result_column)); + EXPECT_TRUE(result_column->only_null()); + ColumnPtr full_result = result_column->convert_to_full_column_if_const(); + assert_mock_nullable_embedding_column(*full_result, std::vector(row_count, 1)); +} + +TEST(EMBED_TEST, embed_function_null_rows_across_batches_through_framework) { + TQueryOptions query_options = create_fake_query_options(); + query_options.__set_embed_max_batch_size(2); + auto query_ctx = MockQueryContext::create(TUniqueId(), ExecEnv::GetInstance(), query_options); + query_ctx->set_mock_ai_resource(); + TQueryGlobals query_globals; + RuntimeState runtime_state(TUniqueId(), 0, query_options, query_globals, nullptr, + query_ctx.get()); + auto ctx = FunctionContext::create_context(&runtime_state, {}, {}); + + std::vector texts = {"", "text-a", "", "text-b", "", "text-c", + "", "text-d", "", "text-e", ""}; + std::vector null_map = {1, 0, 1, 0, 1, 0, 1, 0, 1, 0, 1}; + std::vector resources(texts.size(), "mock_resource"); + auto col_resource = ColumnHelper::create_column(resources); + auto col_text = ColumnHelper::create_nullable_column(texts, null_map); + auto return_type = make_nullable( + std::make_shared(make_nullable(std::make_shared()))); + + Block block; + block.insert({std::move(col_resource), std::make_shared(), "resource"}); + block.insert({std::move(col_text), make_nullable(std::make_shared()), "text"}); + + auto function = get_embed_function(block, return_type); + ASSERT_NE(function, nullptr); + + block.insert({nullptr, return_type, "result"}); + const size_t result_idx = 2; + Status exec_status = function->execute(ctx.get(), block, {0, 1}, result_idx, texts.size()); + + ASSERT_TRUE(exec_status.ok()) << exec_status.to_string(); + assert_mock_nullable_embedding_column(*block.get_by_position(result_idx).column, null_map); +} + TEST(EMBED_TEST, embed_function_multimodal_batch_request) { auto runtime_state = std::make_unique(); auto ctx = FunctionContext::create_context(runtime_state.get(), {}, {}); @@ -327,8 +508,8 @@ TEST(EMBED_TEST, embed_function_multimodal_batch_request) { ColumnNumbers arguments = {0, 1}; size_t result_idx = 2; FunctionEmbed embed_func; - Status exec_status = embed_func.execute_with_adapter(ctx.get(), block, arguments, result_idx, - file_json_rows.size(), config, adapter); + Status exec_status = embed_func.execute(ctx.get(), block, arguments, result_idx, + file_json_rows.size(), config, adapter); ASSERT_TRUE(exec_status.ok()) << exec_status.to_string(); EXPECT_THAT(counting_adapter->batch_sizes, ::testing::ElementsAre(3)); @@ -373,8 +554,8 @@ TEST(EMBED_TEST, embed_function_multimodal_batch_split_by_session_variable) { ColumnNumbers arguments = {0, 1}; size_t result_idx = 2; FunctionEmbed embed_func; - Status exec_status = embed_func.execute_with_adapter(ctx.get(), block, arguments, result_idx, - file_json_rows.size(), config, adapter); + Status exec_status = embed_func.execute(ctx.get(), block, arguments, result_idx, + file_json_rows.size(), config, adapter); ASSERT_TRUE(exec_status.ok()) << exec_status.to_string(); EXPECT_THAT(counting_adapter->batch_sizes, ::testing::ElementsAre(2, 1)); @@ -416,8 +597,8 @@ TEST(EMBED_TEST, embed_function_text_batch_split_by_session_variable) { ColumnNumbers arguments = {0, 1}; size_t result_idx = 2; FunctionEmbed embed_func; - Status exec_status = embed_func.execute_with_adapter(ctx.get(), block, arguments, result_idx, - texts.size(), config, adapter); + Status exec_status = embed_func.execute(ctx.get(), block, arguments, result_idx, texts.size(), + config, adapter); ASSERT_TRUE(exec_status.ok()) << exec_status.to_string(); EXPECT_THAT(counting_adapter->batch_sizes, ::testing::ElementsAre(2, 1));