diff --git a/be/src/exprs/function/ai/ai_adapter.h b/be/src/exprs/function/ai/ai_adapter.h index a50fa7123eb240..ba3a0039eb8beb 100644 --- a/be/src/exprs/function/ai/ai_adapter.h +++ b/be/src/exprs/function/ai/ai_adapter.h @@ -24,7 +24,6 @@ #include #include #include -#include #include #include @@ -35,6 +34,7 @@ #include "rapidjson/writer.h" #include "service/http/http_client.h" #include "service/http/http_headers.h" +#include "util/security.h" namespace doris { #include "common/compile_check_begin.h" @@ -91,6 +91,20 @@ struct AIResource { } }; +enum class MultimodalType { IMAGE, VIDEO, AUDIO }; + +inline const char* multimodal_type_to_string(MultimodalType type) { + switch (type) { + case MultimodalType::IMAGE: + return "image"; + case MultimodalType::VIDEO: + return "video"; + case MultimodalType::AUDIO: + return "audio"; + } + return "unknown"; +} + class AIAdapter { public: virtual ~AIAdapter() = default; @@ -126,19 +140,33 @@ class AIAdapter { virtual Status build_embedding_request(const std::vector& inputs, std::string& request_body) const { - return Status::NotSupported("{} does not support the Embed feature.", + return embed_not_supported_status(); + } + + virtual Status build_multimodal_embedding_request( + const std::vector& /*media_types*/, + const std::vector& /*media_urls*/, + const std::vector& /*media_content_types*/, + std::string& /*request_body*/) const { + return Status::NotSupported("{} does not support multimodal Embed feature.", _config.provider_type); } virtual Status parse_embedding_response(const std::string& response_body, std::vector>& results) const { - return Status::NotSupported("{} does not support the Embed feature.", - _config.provider_type); + return embed_not_supported_status(); } protected: TAIResource _config; + Status embed_not_supported_status() const { + return Status::NotSupported( + "{} does not support the Embed feature. Currently supported providers are " + "OpenAI, Gemini, Voyage, Jina, Qwen, and Minimax.", + _config.provider_type); + } + // Appends one provider-parsed text result to `results`. // The adapter has already parsed the provider's outer response envelope before calling here. // Example: @@ -188,6 +216,50 @@ class AIAdapter { doc.AddMember(name, _config.dimensions, allocator); } } + + // Validates common multimodal embedding request invariants shared by providers. + Status validate_multimodal_embedding_inputs( + std::string_view provider_name, const std::vector& media_types, + const std::vector& media_urls, + std::initializer_list supported_types) const { + if (media_urls.empty()) { + return Status::InvalidArgument("{} multimodal embed inputs can not be empty", + provider_name); + } + if (media_types.size() != media_urls.size()) { + return Status::InvalidArgument( + "{} multimodal embed input size mismatch, media_types={}, media_urls={}", + provider_name, media_types.size(), media_urls.size()); + } + for (MultimodalType media_type : media_types) { + bool supported = false; + for (MultimodalType supported_type : supported_types) { + if (media_type == supported_type) { + supported = true; + break; + } + } + if (!supported) [[unlikely]] { + return Status::InvalidArgument( + "{} only supports {} multimodal embed, got {}", provider_name, + supported_multimodal_types_to_string(supported_types), + multimodal_type_to_string(media_type)); + } + } + return Status::OK(); + } + + static std::string supported_multimodal_types_to_string( + std::initializer_list supported_types) { + std::string result; + for (MultimodalType type : supported_types) { + if (!result.empty()) { + result += "/"; + } + result += multimodal_type_to_string(type); + } + return result; + } }; // Most LLM-providers' Embedding formats are based on VoyageAI. @@ -233,6 +305,68 @@ class VoyageAIAdapter : public AIAdapter { return Status::OK(); } + Status build_multimodal_embedding_request( + const std::vector& media_types, + const std::vector& media_urls, + const std::vector& /*media_content_types*/, + std::string& request_body) const override { + RETURN_IF_ERROR(validate_multimodal_embedding_inputs( + "VoyageAI", media_types, media_urls, + {MultimodalType::IMAGE, MultimodalType::VIDEO})); + if (_config.dimensions != -1) { + LOG(WARNING) << "VoyageAI multimodal embedding currently ignores dimensions parameter, " + << "model=" << _config.model_name << ", dimensions=" << _config.dimensions; + } + + rapidjson::Document doc; + doc.SetObject(); + auto& allocator = doc.GetAllocator(); + + /*{ + "inputs": [ + { + "content": [ + {"type": "image_url", "image_url": ""} + ] + }, + { + "content": [ + {"type": "video_url", "video_url": ""} + ] + } + ], + "model": "voyage-multimodal-3.5" + }*/ + doc.AddMember("model", rapidjson::Value(_config.model_name.c_str(), allocator), allocator); + + rapidjson::Value request_inputs(rapidjson::kArrayType); + for (size_t i = 0; i < media_urls.size(); ++i) { + rapidjson::Value input(rapidjson::kObjectType); + rapidjson::Value content(rapidjson::kArrayType); + rapidjson::Value media_item(rapidjson::kObjectType); + if (media_types[i] == MultimodalType::IMAGE) { + media_item.AddMember("type", "image_url", allocator); + media_item.AddMember("image_url", + rapidjson::Value(media_urls[i].c_str(), allocator), allocator); + } else { + media_item.AddMember("type", "video_url", allocator); + media_item.AddMember("video_url", + rapidjson::Value(media_urls[i].c_str(), allocator), allocator); + } + content.PushBack(media_item, allocator); + input.AddMember("content", content, allocator); + request_inputs.PushBack(input, allocator); + } + + doc.AddMember("inputs", request_inputs, allocator); + + rapidjson::StringBuffer buffer; + rapidjson::Writer writer(buffer); + doc.Accept(writer); + request_body = buffer.GetString(); + return Status::OK(); + } + Status parse_embedding_response(const std::string& response_body, std::vector>& results) const override { rapidjson::Document doc; @@ -365,7 +499,6 @@ class LocalAdapter : public AIAdapter { } else { return Status::NotSupported("Unsupported response format from local AI."); } - return Status::OK(); } @@ -396,6 +529,15 @@ class LocalAdapter : public AIAdapter { return Status::OK(); } + Status build_multimodal_embedding_request( + const std::vector& /*media_types*/, + const std::vector& /*media_urls*/, + const std::vector& /*media_content_types*/, + std::string& /*request_body*/) const override { + return Status::NotSupported("{} does not support multimodal Embed feature.", + _config.provider_type); + } + Status parse_embedding_response(const std::string& response_body, std::vector>& results) const override { rapidjson::Document doc; @@ -746,6 +888,15 @@ class OpenAIAdapter : public VoyageAIAdapter { return Status::OK(); } + Status build_multimodal_embedding_request( + const std::vector& /*media_types*/, + const std::vector& /*media_urls*/, + const std::vector& /*media_content_types*/, + std::string& /*request_body*/) const override { + return Status::NotSupported("{} does not support multimodal Embed feature.", + _config.provider_type); + } + protected: bool supports_dimension_param(const std::string& model_name) const override { return !(model_name == "text-embedding-ada-002"); @@ -758,14 +909,12 @@ class DeepSeekAdapter : public OpenAIAdapter { public: Status build_embedding_request(const std::vector& inputs, std::string& request_body) const override { - return Status::NotSupported("{} does not support the Embed feature.", - _config.provider_type); + return embed_not_supported_status(); } Status parse_embedding_response(const std::string& response_body, std::vector>& results) const override { - return Status::NotSupported("{} does not support the Embed feature.", - _config.provider_type); + return embed_not_supported_status(); } }; @@ -773,14 +922,12 @@ class MoonShotAdapter : public OpenAIAdapter { public: Status build_embedding_request(const std::vector& inputs, std::string& request_body) const override { - return Status::NotSupported("{} does not support the Embed feature.", - _config.provider_type); + return embed_not_supported_status(); } Status parse_embedding_response(const std::string& response_body, std::vector>& results) const override { - return Status::NotSupported("{} does not support the Embed feature.", - _config.provider_type); + return embed_not_supported_status(); } }; @@ -822,6 +969,105 @@ class ZhipuAdapter : public OpenAIAdapter { }; class QwenAdapter : public OpenAIAdapter { +public: + Status build_multimodal_embedding_request( + const std::vector& media_types, + const std::vector& media_urls, + const std::vector& /*media_content_types*/, + std::string& request_body) const override { + RETURN_IF_ERROR(validate_multimodal_embedding_inputs( + "QWEN", media_types, media_urls, {MultimodalType::IMAGE, MultimodalType::VIDEO})); + + rapidjson::Document doc; + doc.SetObject(); + auto& allocator = doc.GetAllocator(); + + /*{ + "model": "tongyi-embedding-vision-plus", + "input": { + "contents": [ + {"image": ""}, + {"video": ""} + ] + } + "parameters": { + "dimension": 512 + } + }*/ + doc.AddMember("model", rapidjson::Value(_config.model_name.c_str(), allocator), allocator); + rapidjson::Value input(rapidjson::kObjectType); + rapidjson::Value contents(rapidjson::kArrayType); + + for (size_t i = 0; i < media_urls.size(); ++i) { + rapidjson::Value media_item(rapidjson::kObjectType); + if (media_types[i] == MultimodalType::IMAGE) { + media_item.AddMember("image", rapidjson::Value(media_urls[i].c_str(), allocator), + allocator); + } else { + media_item.AddMember("video", rapidjson::Value(media_urls[i].c_str(), allocator), + allocator); + } + contents.PushBack(media_item, allocator); + } + + input.AddMember("contents", contents, allocator); + doc.AddMember("input", input, allocator); + if (_config.dimensions != -1 && supports_dimension_param(_config.model_name)) { + rapidjson::Value parameters(rapidjson::kObjectType); + std::string param_name = get_dimension_param_name(); + rapidjson::Value dimension_name(param_name.c_str(), allocator); + parameters.AddMember(dimension_name, _config.dimensions, allocator); + doc.AddMember("parameters", parameters, allocator); + } + + rapidjson::StringBuffer buffer; + rapidjson::Writer writer(buffer); + doc.Accept(writer); + request_body = buffer.GetString(); + return Status::OK(); + } + + Status parse_embedding_response(const std::string& response_body, + std::vector>& results) const override { + rapidjson::Document doc; + doc.Parse(response_body.c_str()); + + if (doc.HasParseError() || !doc.IsObject()) [[unlikely]] { + return Status::InternalError("Failed to parse {} response: {}", _config.provider_type, + response_body); + } + // Qwen multimodal embedding usually returns: + // { + // "output": { + // "embeddings": [ + // {"index":0, "embedding":[...], "type":"image|video|text"}, + // ... + // ] + // } + // } + // + // In text-only or compatibility endpoints, Qwen may also return OpenAI-style + // "data":[{"embedding":[...]}]. For compatibility we first parse native + // output.embeddings and then fallback to OpenAIAdapter parser. + if (doc.HasMember("output") && doc["output"].IsObject() && + doc["output"].HasMember("embeddings") && doc["output"]["embeddings"].IsArray()) { + const auto& embeddings = doc["output"]["embeddings"]; + results.reserve(embeddings.Size()); + for (rapidjson::SizeType i = 0; i < embeddings.Size(); i++) { + if (!embeddings[i].HasMember("embedding") || + !embeddings[i]["embedding"].IsArray()) { + return Status::InternalError("Invalid {} response format: {}", + _config.provider_type, response_body); + } + std::transform(embeddings[i]["embedding"].Begin(), embeddings[i]["embedding"].End(), + std::back_inserter(results.emplace_back()), + [](const auto& val) { return val.GetFloat(); }); + } + return Status::OK(); + } + return OpenAIAdapter::parse_embedding_response(response_body, results); + } + protected: bool supports_dimension_param(const std::string& model_name) const override { static const std::unordered_set no_dimension_models = { @@ -832,6 +1078,56 @@ class QwenAdapter : public OpenAIAdapter { std::string get_dimension_param_name() const override { return "dimension"; } }; +class JinaAdapter : public VoyageAIAdapter { +public: + Status build_multimodal_embedding_request( + const std::vector& media_types, + const std::vector& media_urls, + const std::vector& /*media_content_types*/, + std::string& request_body) const override { + RETURN_IF_ERROR(validate_multimodal_embedding_inputs( + "JINA", media_types, media_urls, {MultimodalType::IMAGE, MultimodalType::VIDEO})); + + rapidjson::Document doc; + doc.SetObject(); + auto& allocator = doc.GetAllocator(); + + /*{ + "model": "jina-embeddings-v4", + "task": "text-matching", + "input": [ + {"image": ""}, + {"video": ""} + ] + }*/ + doc.AddMember("model", rapidjson::Value(_config.model_name.c_str(), allocator), allocator); + doc.AddMember("task", "text-matching", allocator); + + rapidjson::Value input(rapidjson::kArrayType); + for (size_t i = 0; i < media_urls.size(); ++i) { + rapidjson::Value media_item(rapidjson::kObjectType); + if (media_types[i] == MultimodalType::IMAGE) { + media_item.AddMember("image", rapidjson::Value(media_urls[i].c_str(), allocator), + allocator); + } else { + media_item.AddMember("video", rapidjson::Value(media_urls[i].c_str(), allocator), + allocator); + } + input.PushBack(media_item, allocator); + } + if (_config.dimensions != -1 && supports_dimension_param(_config.model_name)) { + doc.AddMember("dimensions", _config.dimensions, allocator); + } + doc.AddMember("input", input, allocator); + + rapidjson::StringBuffer buffer; + rapidjson::Writer writer(buffer); + doc.Accept(writer); + request_body = buffer.GetString(); + return Status::OK(); + } +}; + class BaichuanAdapter : public OpenAIAdapter { protected: bool supports_dimension_param(const std::string& model_name) const override { return false; } @@ -965,7 +1261,6 @@ class GeminiAdapter : public AIAdapter { RETURN_IF_ERROR(append_parsed_text_result( candidates[i]["content"]["parts"][0]["text"].GetString(), results)); } - return Status::OK(); } @@ -1033,6 +1328,81 @@ class GeminiAdapter : public AIAdapter { return Status::OK(); } + Status build_multimodal_embedding_request(const std::vector& media_types, + const std::vector& media_urls, + const std::vector& media_content_types, + std::string& request_body) const override { + RETURN_IF_ERROR(validate_multimodal_embedding_inputs( + "Gemini", media_types, media_urls, + {MultimodalType::IMAGE, MultimodalType::AUDIO, MultimodalType::VIDEO})); + if (media_content_types.size() != media_urls.size()) { + return Status::InvalidArgument( + "Gemini multimodal embed input size mismatch, media_content_types={}, " + "media_urls={}", + media_content_types.size(), media_urls.size()); + } + + rapidjson::Document doc; + doc.SetObject(); + auto& allocator = doc.GetAllocator(); + + /*{ + "requests": [ + { + "model": "models/gemini-embedding-2-preview", + "content": { + "parts": [ + {"file_data": {"mime_type": "", "file_uri": ""}} + ] + }, + "outputDimensionality": 768 + }, + { + "model": "models/gemini-embedding-2-preview", + "content": { + "parts": [ + {"file_data": {"mime_type": "", "file_uri": ""}} + ] + }, + "outputDimensionality": 768 + } + ] + }*/ + std::string model_name = _config.model_name; + if (!model_name.starts_with("models/")) { + model_name = "models/" + model_name; + } + + rapidjson::Value requests(rapidjson::kArrayType); + for (size_t i = 0; i < media_urls.size(); ++i) { + rapidjson::Value request(rapidjson::kObjectType); + request.AddMember("model", rapidjson::Value(model_name.c_str(), allocator), allocator); + add_dimension_params(request, allocator); + + rapidjson::Value content(rapidjson::kObjectType); + rapidjson::Value parts(rapidjson::kArrayType); + rapidjson::Value part(rapidjson::kObjectType); + rapidjson::Value file_data(rapidjson::kObjectType); + file_data.AddMember("mime_type", + rapidjson::Value(media_content_types[i].c_str(), allocator), + allocator); + file_data.AddMember("file_uri", rapidjson::Value(media_urls[i].c_str(), allocator), + allocator); + part.AddMember("file_data", file_data, allocator); + parts.PushBack(part, allocator); + content.AddMember("parts", parts, allocator); + request.AddMember("content", content, allocator); + requests.PushBack(request, allocator); + } + doc.AddMember("requests", requests, allocator); + + rapidjson::StringBuffer buffer; + rapidjson::Writer writer(buffer); + doc.Accept(writer); + request_body = buffer.GetString(); + return Status::OK(); + } + Status parse_embedding_response(const std::string& response_body, std::vector>& results) const override { rapidjson::Document doc; @@ -1043,6 +1413,12 @@ class GeminiAdapter : public AIAdapter { response_body); } if (doc.HasMember("embeddings") && doc["embeddings"].IsArray()) { + /*{ + "embeddings": [ + {"values": [0.1, 0.2, 0.3]}, + {"values": [0.4, 0.5, 0.6]} + ] + }*/ const auto& embeddings = doc["embeddings"]; results.reserve(embeddings.Size()); for (rapidjson::SizeType i = 0; i < embeddings.Size(); i++) { @@ -1193,6 +1569,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, @@ -1208,6 +1592,18 @@ 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(); + } + + Status build_multimodal_embedding_request( + const std::vector& /*media_types*/, + const std::vector& /*media_urls*/, + const std::vector& /*media_content_types*/, + std::string& /*request_body*/) const override { return Status::OK(); } @@ -1229,6 +1625,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 { @@ -1242,6 +1646,7 @@ class AIAdapterFactory { {"MINIMAX", []() { return std::make_shared(); }}, {"ZHIPU", []() { return std::make_shared(); }}, {"QWEN", []() { return std::make_shared(); }}, + {"JINA", []() { return std::make_shared(); }}, {"BAICHUAN", []() { return std::make_shared(); }}, {"ANTHROPIC", []() { return std::make_shared(); }}, {"GEMINI", []() { return std::make_shared(); }}, diff --git a/be/src/exprs/function/ai/ai_classify.h b/be/src/exprs/function/ai/ai_classify.h index 3b5666647759fd..0b4fbb660481e7 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 diff --git a/be/src/exprs/function/ai/ai_extract.h b/be/src/exprs/function/ai/ai_extract.h index a0a310e41d629d..a550879b9d9b56 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 cf66ed0b835f97..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(); } @@ -47,6 +47,8 @@ class FunctionAIFilter : public AIFunction { private: MutableColumnPtr create_result_column() const { return ColumnUInt8::create(); } + // AI_FILTER-private helper. + // Converts one parsed batch of string flags into BOOL results. Status append_batch_results(const std::vector& batch_results, IColumn& col_result) const { auto& bool_col = assert_cast(col_result); diff --git a/be/src/exprs/function/ai/ai_fix_grammar.h b/be/src/exprs/function/ai/ai_fix_grammar.h index 4b9687f7b5b536..724f360a76121a 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..b7286d49240cc1 100644 --- a/be/src/exprs/function/ai/ai_functions.cpp +++ b/be/src/exprs/function/ai/ai_functions.cpp @@ -15,7 +15,7 @@ // specific language governing permissions and limitations // under the License. -#include "core/column/column_array.h" +#include "core/column/column_array_view.h" #include "exprs/function/ai/ai_classify.h" #include "exprs/function/ai/ai_extract.h" #include "exprs/function/ai/ai_filter.h" @@ -30,151 +30,93 @@ #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 { - // Get the text column - const ColumnWithTypeAndName& text_column = block.get_by_position(arguments[1]); - StringRef text = text_column.column->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); - const auto* col_array = check_and_get_column(*array_column); - if (col_array == nullptr) { +static Status format_labels(const ColumnPtr& labels_column, size_t row_num, + std::string_view function_name, std::string& labels_str) { + auto readable_column = check_column_const_set_readability(*labels_column, row_num); + if (!is_column(*readable_column.first)) { return Status::InternalError( - "labels argument for {} must be Array(String) or Array(Varchar)", name); - } - - std::vector label_values; - const auto& data = col_array->get_data(); - const auto& offsets = col_array->get_offsets(); - size_t start = array_row_num > 0 ? offsets[array_row_num - 1] : 0; - size_t end = offsets[array_row_num]; - for (size_t i = start; i < end; ++i) { - Field field; - data.get(i, field); - label_values.emplace_back(field.template get()); + "labels argument for {} must be Array(String) or Array(Varchar)", function_name); } - std::string labels_str = "["; - for (size_t i = 0; i < label_values.size(); ++i) { - if (i > 0) { + auto labels_view = ColumnArrayView::create(labels_column); + auto labels = labels_view[row_num]; + labels_str = "["; + bool is_first_label = true; + for (size_t i = 0; i < labels.size(); ++i) { + if (labels.is_null_at(i)) { + continue; + } + if (!is_first_label) { labels_str += ", "; } - labels_str += "\"" + label_values[i] + "\""; + StringRef label = labels.value_at(i); + labels_str += "\""; + labels_str.append(label.data, label.size); + labels_str += "\""; + is_first_label = false; } labels_str += "]"; + return Status::OK(); +} + +Status FunctionAIClassify::build_prompt(const Columns& prompt_columns, size_t row_num, + std::string& prompt) const { + // Get the text column + StringRef text = prompt_columns[0]->get_data_at(row_num); + std::string text_str = std::string(text.data, text.size); + + std::string labels_str; + RETURN_IF_ERROR(format_labels(prompt_columns[1], row_num, name, labels_str)); prompt = "Labels: " + labels_str + "\nText: " + text_str; 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); - const auto* col_array = check_and_get_column(*array_column); - if (col_array == nullptr) { - return Status::InternalError( - "labels argument for {} must be Array(String) or Array(Varchar)", name); - } - - std::vector label_values; - const auto& offsets = col_array->get_offsets(); - const auto& data = col_array->get_data(); - size_t start = array_row_num > 0 ? offsets[array_row_num - 1] : 0; - size_t end = offsets[array_row_num]; - for (size_t i = start; i < end; ++i) { - Field field; - data.get(i, field); - label_values.emplace_back(field.template get()); - } - - std::string labels_str = "["; - for (size_t i = 0; i < label_values.size(); ++i) { - if (i > 0) { - labels_str += ", "; - } - labels_str += "\"" + label_values[i] + "\""; - } - labels_str += "]"; + std::string labels_str; + RETURN_IF_ERROR(format_labels(prompt_columns[1], row_num, name, labels_str)); prompt = "Labels: " + labels_str + "\nText: " + text_str; 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); - const auto* col_array = check_and_get_column(*array_column); - if (col_array == nullptr) { - return Status::InternalError( - "labels argument for {} must be Array(String) or Array(Varchar)", name); - } - - std::vector label_values; - const auto& offsets = col_array->get_offsets(); - const auto& data = col_array->get_data(); - size_t start = array_row_num > 0 ? offsets[array_row_num - 1] : 0; - size_t end = offsets[array_row_num]; - for (size_t i = start; i < end; ++i) { - Field field; - data.get(i, field); - label_values.emplace_back(field.template get()); - } - - std::string labels_str = "["; - for (size_t i = 0; i < label_values.size(); ++i) { - if (i > 0) { - labels_str += ", "; - } - labels_str += "\"" + label_values[i] + "\""; - } - labels_str += "]"; + std::string labels_str; + RETURN_IF_ERROR(format_labels(prompt_columns[1], row_num, name, labels_str)); prompt = "Labels: " + labels_str + "\nText: " + text_str; 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 +124,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 528499992f5cf2..6fdd0414e47907 100644 --- a/be/src/exprs/function/ai/ai_functions.h +++ b/be/src/exprs/function/ai/ai_functions.h @@ -22,7 +22,6 @@ #include #include -#include #include #include #include @@ -37,14 +36,17 @@ #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" #include "runtime/runtime_state.h" #include "service/http/http_client.h" +#include "util/security.h" #include "util/string_util.h" #include "util/threadpool.h" @@ -66,10 +68,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(); @@ -77,38 +90,33 @@ 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); - !status.ok()) [[unlikely]] { + !status.ok()) { 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: // Reads the shared AI context window size from query options. String AI batch functions and // ai_agg both use the same byte-based session variable so batching behavior stays consistent. static int64_t get_ai_context_window_size(FunctionContext* context) { + DORIS_CHECK(context != nullptr); QueryContext* query_ctx = context->state()->get_query_ctx(); DORIS_CHECK(query_ctx != nullptr); - 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(); + return query_ctx->query_options().ai_context_window_size; } MutableColumnPtr create_result_column() const { return ColumnString::create(); } @@ -123,11 +131,11 @@ class AIFunction : public IFunction { return Status::OK(); } - // 1. If users configure only the version root like `.../v1` or `.../v1beta`, append - // `models/:batchEmbedContents` for `embed`, and `models/:generateContent` - // for other AI scalar functions. - // 2. `:embedContent` -> `:batchEmbedContents` static void normalize_endpoint(TAIResource& config) { + // 1. If users configure only the version root like `.../v1` or `.../v1beta`, append + // `models/:batchEmbedContents` for `embed`, and `models/:generateContent` + // for other AI scalar functions. + // 2. `:embedContent` -> `:batchEmbedContents` if (iequal(config.provider_type, "GEMINI")) { if (iequal(Derived::name, "embed") && config.endpoint.ends_with(":embedContent")) { static constexpr std::string_view legacy_suffix = ":embedContent"; @@ -144,6 +152,7 @@ class AIFunction : public IFunction { if (!model_name.starts_with("models/")) { model_name = "models/" + model_name; } + config.endpoint += "/"; config.endpoint += model_name; config.endpoint += @@ -165,7 +174,7 @@ class AIFunction : public IFunction { Status do_send_request(HttpClient* client, const std::string& request_body, std::string& response, const TAIResource& config, std::shared_ptr& adapter, FunctionContext* context) const { - RETURN_IF_ERROR(client->init(config.endpoint)); + RETURN_IF_ERROR(client->init(config.endpoint, false)); QueryContext* query_ctx = context->state()->get_query_ctx(); int64_t remaining_query_time = query_ctx->get_remaining_query_time_seconds(); @@ -179,7 +188,22 @@ class AIFunction : public IFunction { RETURN_IF_ERROR(adapter->set_authentication(client)); } - return client->execute_post_request(request_body, &response); + Status st = client->execute_post_request(request_body, &response); + long http_status = client->get_http_status(); + + if (!st.ok()) { + LOG(INFO) << "AI HTTP request failed before status validation, provider=" + << config.provider_type << ", model=" << config.model_name + << ", endpoint=" << mask_token(config.endpoint) + << ", exec_status=" << st.to_string() << ", response_body=" << response; + return st; + } + if (http_status != 200) { + return Status::HttpError( + "http status code is not 200, code={}, url={}, response_body={}", http_status, + mask_token(config.endpoint), response); + } + return Status::OK(); } // Sends the request with retry mechanism for handling transient failures @@ -229,13 +253,7 @@ class AIFunction : public IFunction { results.clear(); results.reserve(batch_prompts.size()); for (const auto& prompt : batch_prompts) { - if (get_name() == "ai_filter") { - results.emplace_back("0"); - } else if (get_name() == "ai_similarity") { - results.emplace_back("0.0"); - } else { - results.emplace_back("this is a mock response. " + prompt); - } + results.emplace_back("this is a mock response. " + prompt); } return Status::OK(); } @@ -243,7 +261,9 @@ class AIFunction : public IFunction { std::string batch_prompt; RETURN_IF_ERROR(build_batch_prompt(batch_prompts, batch_prompt)); + std::vector inputs = {batch_prompt}; + std::vector parsed_response; std::string request_body; RETURN_IF_ERROR(adapter->build_request_payload( @@ -251,7 +271,6 @@ class AIFunction : public IFunction { std::string response; RETURN_IF_ERROR(send_request_to_llm(request_body, response, config, adapter, context)); - std::vector parsed_response; RETURN_IF_ERROR(adapter->parse_response(response, parsed_response)); if (parsed_response.empty()) { return Status::InternalError("AI returned empty result"); @@ -273,40 +292,82 @@ 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; + 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) { if (!batch_prompts.empty()) { - RETURN_IF_ERROR(flush_batch_prompts(batch_prompts, col_result, config, adapter, - context)); + std::vector batch_results; + 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_prompts.clear(); current_batch_size = 2; } std::vector single_prompts; single_prompts.emplace_back(std::move(prompt)); - RETURN_IF_ERROR( - flush_batch_prompts(single_prompts, col_result, config, adapter, context)); + std::vector single_results; + 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)); continue; } size_t additional_size = entry_size + (batch_prompts.empty() ? 0 : 1); if (!batch_prompts.empty() && current_batch_size + additional_size > max_batch_prompt_size) { - RETURN_IF_ERROR( - flush_batch_prompts(batch_prompts, col_result, config, adapter, context)); + std::vector batch_results; + 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_prompts.clear(); current_batch_size = 2; additional_size = entry_size; } @@ -316,9 +377,36 @@ class AIFunction : public IFunction { } if (!batch_prompts.empty()) { - RETURN_IF_ERROR( - flush_batch_prompts(batch_prompts, col_result, config, adapter, context)); + std::vector batch_results; + 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)); + } + + 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(); } @@ -333,38 +421,17 @@ class AIFunction : public IFunction { const std::shared_ptr>& ai_resources = context->state()->get_query_ctx()->get_ai_resources(); - if (!ai_resources) { - return Status::InternalError("AI resources metadata missing in QueryContext"); - } + DORIS_CHECK(ai_resources); auto it = ai_resources->find(resource_name); - if (it == ai_resources->end()) { - return Status::InvalidArgument("AI resource not found: " + resource_name); - } + DORIS_CHECK(it != ai_resources->end()); config = it->second; normalize_endpoint(config); adapter = AIAdapterFactory::create_adapter(config.provider_type); - if (!adapter) { - return Status::InvalidArgument("Unsupported AI provider type: " + config.provider_type); - } - adapter->init(config); - - return Status::OK(); - } + DORIS_CHECK(adapter); - Status flush_batch_prompts(std::vector& batch_prompts, IColumn& col_result, - const TAIResource& config, std::shared_ptr& adapter, - FunctionContext* context) const { - if (batch_prompts.empty()) { - return Status::OK(); - } - std::vector batch_results; - RETURN_IF_ERROR( - execute_batch_request(batch_prompts, batch_results, config, adapter, context)); - RETURN_IF_ERROR( - assert_cast(*this).append_batch_results(batch_results, col_result)); - batch_prompts.clear(); + adapter->init(config); return Status::OK(); } diff --git a/be/src/exprs/function/ai/ai_generate.h b/be/src/exprs/function/ai/ai_generate.h index b15701024aab66..7ce5766b8ac01c 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 1202e28e53def3..5de6125d3f3d8a 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 50fa988b9232d7..fab30fe26dfc1e 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 168b9e665d5d34..b3a63c8e783476 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 4df5b732063369..f193a1c171c1b5 100644 --- a/be/src/exprs/function/ai/embed.h +++ b/be/src/exprs/function/ai/embed.h @@ -17,8 +17,17 @@ #pragma once +#include +#include + +#include + +#include "core/data_type/data_type_nullable.h" #include "core/data_type/primitive_type.h" #include "exprs/function/ai/ai_functions.h" +#include "util/jsonb_utils.h" +#include "util/s3_uri.h" +#include "util/s3_util.h" namespace doris { class FunctionEmbed : public AIFunction { @@ -29,50 +38,106 @@ 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())); } - static FunctionPtr create() { return std::make_shared(); } + using PreparedFunctionImpl::execute; - 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 { + 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()); } + 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, 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, 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 {}", + input.type->get_name()); + } + + static FunctionPtr create() { return std::make_shared(); } + +private: + static int32_t _get_embed_max_batch_size(FunctionContext* context) { + QueryContext* query_ctx = context->state()->get_query_ctx(); + DORIS_CHECK(query_ctx != nullptr); + + return query_ctx->query_options().embed_max_batch_size; + } + + Status _execute_text_embed(FunctionContext* context, Block& block, uint32_t result, + size_t input_rows_count, const TAIResource& config, + 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; size_t current_batch_size = 0; - const int32_t max_batch_size = get_embed_max_batch_size(context); + 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(); + if (prompt_size > max_context_window_size) { - RETURN_IF_ERROR(flush_text_embedding_batch(batch_prompts, *col_result, config, - adapter, context)); + // flush history batch + RETURN_IF_ERROR(_flush_text_embedding_batch(batch_prompts, *col_result, config, + adapter, context)); current_batch_size = 0; batch_prompts.emplace_back(std::move(prompt)); - RETURN_IF_ERROR(flush_text_embedding_batch(batch_prompts, *col_result, config, - adapter, context)); + RETURN_IF_ERROR(_flush_text_embedding_batch(batch_prompts, *col_result, config, + adapter, context)); continue; } if (!batch_prompts.empty() && (current_batch_size + prompt_size > max_context_window_size || batch_prompts.size() >= static_cast(max_batch_size))) { - RETURN_IF_ERROR(flush_text_embedding_batch(batch_prompts, *col_result, config, - adapter, context)); + RETURN_IF_ERROR(_flush_text_embedding_batch(batch_prompts, *col_result, config, + adapter, context)); current_batch_size = 0; } @@ -81,44 +146,82 @@ class FunctionEmbed : public AIFunction { } RETURN_IF_ERROR( - flush_text_embedding_batch(batch_prompts, *col_result, config, adapter, context)); + _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(); } -private: - static int32_t get_embed_max_batch_size(FunctionContext* context) { - QueryContext* query_ctx = context->state()->get_query_ctx(); - DORIS_CHECK(query_ctx != nullptr); - return query_ctx->query_options().embed_max_batch_size; - } + Status _execute_multimodal_embed(FunctionContext* context, Block& block, uint32_t result, + size_t input_rows_count, const TAIResource& config, + 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; - Status flush_text_embedding_batch(std::vector& batch_prompts, - ColumnArray& col_result, const TAIResource& config, - std::shared_ptr& adapter, - FunctionContext* context) const { - if (batch_prompts.empty()) { - return Status::OK(); + int64_t ttl_seconds = 3600; + QueryContext* query_ctx = context->state()->get_query_ctx(); + if (query_ctx && query_ctx->query_options().__isset.file_presigned_url_ttl_seconds) { + ttl_seconds = query_ctx->query_options().file_presigned_url_ttl_seconds; + if (ttl_seconds <= 0) { + ttl_seconds = 3600; + } } - std::string request_body; - RETURN_IF_ERROR(adapter->build_embedding_request(batch_prompts, request_body)); + const int32_t max_batch_size = _get_embed_max_batch_size(context); - std::vector> batch_results; - RETURN_IF_ERROR(execute_embedding_request(request_body, batch_results, batch_prompts.size(), - config, adapter, context)); - for (const auto& batch_result : batch_results) { - insert_embedding_result(col_result, batch_result); + 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(*input_column, i, file_input)); + + std::string content_type; + MultimodalType media_type; + RETURN_IF_ERROR(_infer_media_type(file_input, content_type, media_type)); + + std::string media_url; + RETURN_IF_ERROR(_resolve_media_url(file_input, ttl_seconds, media_url)); + + if (!batch_media_urls.empty() && + batch_media_urls.size() >= static_cast(max_batch_size)) { + RETURN_IF_ERROR(_flush_multimodal_embedding_batch( + batch_media_types, batch_media_content_types, batch_media_urls, *col_result, + config, adapter, context)); + } + + batch_media_types.emplace_back(media_type); + batch_media_content_types.emplace_back(std::move(content_type)); + batch_media_urls.emplace_back(std::move(media_url)); } - batch_prompts.clear(); + + RETURN_IF_ERROR(_flush_multimodal_embedding_batch( + batch_media_types, batch_media_content_types, batch_media_urls, *col_result, config, + adapter, context)); + + 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_embedding_request(const std::string& request_body, - std::vector>& results, size_t expected_size, - const TAIResource& config, std::shared_ptr& adapter, - FunctionContext* context) const { + // EMBED-private helper. + // Sends one embedding request with a prebuilt request body and validates returned row count. + Status _execute_prebuilt_embedding_request(const std::string& request_body, + std::vector>& results, + size_t expected_size, const TAIResource& config, + std::shared_ptr& adapter, + FunctionContext* context) const { + std::string response; #ifdef BE_TEST if (config.provider_type == "MOCK") { results.clear(); @@ -130,9 +233,9 @@ class FunctionEmbed : public AIFunction { } #endif - std::string response; RETURN_IF_ERROR( this->send_request_to_llm(request_body, response, config, adapter, context)); + RETURN_IF_ERROR(adapter->parse_embedding_response(response, results)); if (results.empty()) { return Status::InternalError("AI returned empty result"); @@ -145,8 +248,58 @@ class FunctionEmbed : public AIFunction { return Status::OK(); } - static void insert_embedding_result(ColumnArray& col_array, - const std::vector& float_result) { + // EMBED-private helper. + // Flushes one accumulated text embedding batch into the output array column. + Status _flush_text_embedding_batch(std::vector& batch_prompts, + ColumnArray& col_result, const TAIResource& config, + std::shared_ptr& adapter, + FunctionContext* context) const { + if (batch_prompts.empty()) { + return Status::OK(); + } + + std::string request_body; + RETURN_IF_ERROR(adapter->build_embedding_request(batch_prompts, request_body)); + std::vector> batch_results; + RETURN_IF_ERROR(_execute_prebuilt_embedding_request( + request_body, batch_results, batch_prompts.size(), config, adapter, context)); + for (const auto& batch_result : batch_results) { + _insert_embedding_result(col_result, batch_result); + } + batch_prompts.clear(); + return Status::OK(); + } + + // EMBED-private helper. + // Flushes one accumulated multimodal embedding batch into the output array column. + Status _flush_multimodal_embedding_batch(std::vector& batch_media_types, + std::vector& batch_media_content_types, + std::vector& batch_media_urls, + ColumnArray& col_result, const TAIResource& config, + std::shared_ptr& adapter, + FunctionContext* context) const { + if (batch_media_urls.empty()) { + return Status::OK(); + } + + std::string request_body; + RETURN_IF_ERROR(adapter->build_multimodal_embedding_request( + batch_media_types, batch_media_urls, batch_media_content_types, request_body)); + + std::vector> batch_results; + RETURN_IF_ERROR(_execute_prebuilt_embedding_request( + request_body, batch_results, batch_media_urls.size(), config, adapter, context)); + for (const auto& batch_result : batch_results) { + _insert_embedding_result(col_result, batch_result); + } + batch_media_types.clear(); + batch_media_content_types.clear(); + batch_media_urls.clear(); + return Status::OK(); + } + + static void _insert_embedding_result(ColumnArray& col_array, + const std::vector& float_result) { auto& offsets = col_array.get_offsets(); auto& nested_nullable_col = assert_cast(col_array.get_data()); auto& nested_col = @@ -160,6 +313,138 @@ class FunctionEmbed : public AIFunction { auto& null_map = nested_nullable_col.get_null_map_column(); 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; + } + return std::equal(prefix.begin(), prefix.end(), s.begin(), [](char a, char b) { + return std::tolower(static_cast(a)) == + std::tolower(static_cast(b)); + }); + } + + static Status _infer_media_type(const rapidjson::Value& file_input, std::string& content_type, + MultimodalType& media_type) { + RETURN_IF_ERROR(_get_required_string_field(file_input, "content_type", content_type)); + + if (_starts_with_ignore_case(content_type, "image/")) { + media_type = MultimodalType::IMAGE; + return Status::OK(); + } else if (_starts_with_ignore_case(content_type, "video/")) { + media_type = MultimodalType::VIDEO; + return Status::OK(); + } else if (_starts_with_ignore_case(content_type, "audio/")) { + media_type = MultimodalType::AUDIO; + return Status::OK(); + } + + return Status::InvalidArgument("Unsupported content_type for EMBED: {}", content_type); + } + + // Parse the FILE-like JSONB argument into a JSON object for downstream field reads. + static Status _parse_file_input(const IColumn& file_column, size_t row_num, + rapidjson::Document& file_input) { + 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(); + } + + // TODO(lzq): After support FILE type, We should use the interface provided by FILE to get the fields + // replacing this function + static Status _get_required_string_field(const rapidjson::Value& obj, const char* field_name, + std::string& value) { + auto iter = obj.FindMember(field_name); + if (iter == obj.MemberEnd() || !iter->value.IsString()) { + return Status::InvalidArgument( + "EMBED file json field '{}' is required and must be a string", field_name); + } + value = iter->value.GetString(); + if (value.empty()) { + return Status::InvalidArgument("EMBED file json field '{}' can not be empty", + field_name); + } + return Status::OK(); + } + + static Status init_s3_client_conf_from_json(const rapidjson::Value& file_input, + S3ClientConf& s3_client_conf) { + std::string endpoint; + RETURN_IF_ERROR(_get_required_string_field(file_input, "endpoint", endpoint)); + std::string region; + RETURN_IF_ERROR(_get_required_string_field(file_input, "region", region)); + + auto get_optional_string_field = [&](const char* field_name, std::string& value) { + auto iter = file_input.FindMember(field_name); + if (iter == file_input.MemberEnd() || iter->value.IsNull()) { + return; + } + DORIS_CHECK(iter->value.IsString()); + value = iter->value.GetString(); + }; + + get_optional_string_field("ak", s3_client_conf.ak); + get_optional_string_field("sk", s3_client_conf.sk); + get_optional_string_field("role_arn", s3_client_conf.role_arn); + get_optional_string_field("external_id", s3_client_conf.external_id); + s3_client_conf.endpoint = endpoint; + s3_client_conf.region = region; + + return Status::OK(); + } + + Status _resolve_media_url(const rapidjson::Value& file_input, int64_t ttl_seconds, + std::string& media_url) const { + std::string uri; + RETURN_IF_ERROR(_get_required_string_field(file_input, "uri", uri)); + + // If it's a direct http/https URL, use it as-is + if (_starts_with_ignore_case(uri, "http://") || _starts_with_ignore_case(uri, "https://")) { + media_url = uri; + return Status::OK(); + } + + S3ClientConf s3_client_conf; + RETURN_IF_ERROR(init_s3_client_conf_from_json(file_input, s3_client_conf)); + auto s3_client = S3ClientFactory::instance().create(s3_client_conf); + if (s3_client == nullptr) { + return Status::InternalError("Failed to create S3 client for EMBED file input"); + } + + S3URI s3_uri(uri); + RETURN_IF_ERROR(s3_uri.parse()); + std::string bucket = s3_uri.get_bucket(); + std::string key = s3_uri.get_key(); + DORIS_CHECK(!bucket.empty() && !key.empty()); + media_url = s3_client->generate_presigned_url({.bucket = bucket, .key = key}, ttl_seconds, + s3_client_conf); + return Status::OK(); + } }; }; // namespace doris diff --git a/be/test/ai/ai_adapter_test.cpp b/be/test/ai/ai_adapter_test.cpp index 6956af5e3fdcfb..2cfda9d715436c 100644 --- a/be/test/ai/ai_adapter_test.cpp +++ b/be/test/ai/ai_adapter_test.cpp @@ -391,6 +391,23 @@ TEST(AI_ADAPTER_TEST, openai_adapter_responses_parse_response) { ASSERT_EQ(results[0], "openai response result"); } +TEST(AI_ADAPTER_TEST, openai_adapter_parse_response_keeps_mask_literals) { + OpenAIAdapter adapter; + std::string resp = R"({"choices":[{"message":{"content":"[MSKED]"}}]})"; + std::vector results; + Status st = adapter.parse_response(resp, results); + ASSERT_TRUE(st.ok()) << st.to_string(); + ASSERT_EQ(results.size(), 1); + ASSERT_EQ(results[0], "[MSKED]"); + + resp = R"({"choices":[{"message":{"content":"[MASK]"}}]})"; + results.clear(); + st = adapter.parse_response(resp, results); + ASSERT_TRUE(st.ok()) << st.to_string(); + ASSERT_EQ(results.size(), 1); + ASSERT_EQ(results[0], "[MASK]"); +} + TEST(AI_ADAPTER_TEST, gemini_adapter_request) { GeminiAdapter adapter; TAIResource config; @@ -561,9 +578,9 @@ TEST(AI_ADAPTER_TEST, unsupported_provider_type) { } TEST(AI_ADAPTER_TEST, adapter_factory_all_types) { - std::vector types = {"LOCAL", "OPENAI", "MOONSHOT", "DEEPSEEK", - "MINIMAX", "ZHIPU", "QWEN", "BAICHUAN", - "ANTHROPIC", "GEMINI", "VOYAGEAI", "MOCK"}; + std::vector types = {"LOCAL", "OPENAI", "MOONSHOT", "DEEPSEEK", "MINIMAX", + "ZHIPU", "QWEN", "JINA", "BAICHUAN", "ANTHROPIC", + "GEMINI", "VOYAGEAI", "MOCK"}; for (const auto& type : types) { auto adapter = doris::AIAdapterFactory::create_adapter(type); ASSERT_TRUE(adapter != nullptr) << "Adapter not found for type: " << type; @@ -688,4 +705,404 @@ TEST(AI_ADAPTER_TEST, voyage_adapter_chat_test) { "[NOT_IMPLEMENTED_ERROR]VoyageAI don't support text generation"); } +TEST(AI_ADAPTER_TEST, qwen_multimodal_embedding_request_image) { + QwenAdapter adapter; + TAIResource config; + config.model_name = "tongyi-embedding-vision-plus"; + config.dimensions = 1024; + adapter.init(config); + + std::string request_body; + Status st = adapter.build_multimodal_embedding_request({MultimodalType::IMAGE}, + {"https://a/b/c.png"}, {}, request_body); + ASSERT_TRUE(st.ok()) << st.to_string(); + + rapidjson::Document doc; + doc.Parse(request_body.c_str()); + ASSERT_FALSE(doc.HasParseError()) << request_body; + ASSERT_TRUE(doc.IsObject()); + ASSERT_TRUE(doc.HasMember("model")); + ASSERT_STREQ(doc["model"].GetString(), "tongyi-embedding-vision-plus"); + ASSERT_TRUE(doc.HasMember("input")); + ASSERT_TRUE(doc["input"].HasMember("contents")); + const auto& contents = doc["input"]["contents"]; + ASSERT_TRUE(contents.IsArray()); + ASSERT_EQ(contents.Size(), 1); + ASSERT_TRUE(contents[0].HasMember("image")); + ASSERT_STREQ(contents[0]["image"].GetString(), "https://a/b/c.png"); + ASSERT_TRUE(doc.HasMember("parameters")); + ASSERT_TRUE(doc["parameters"].IsObject()); + ASSERT_TRUE(doc["parameters"].HasMember("dimension")); + ASSERT_EQ(doc["parameters"]["dimension"].GetInt(), 1024); +} + +TEST(AI_ADAPTER_TEST, qwen_multimodal_embedding_request_video) { + QwenAdapter adapter; + TAIResource config; + config.model_name = "tongyi-embedding-vision-plus"; + config.dimensions = 1024; + adapter.init(config); + + std::string request_body; + Status st = adapter.build_multimodal_embedding_request({MultimodalType::VIDEO}, + {"https://a/b/c.mp4"}, {}, request_body); + ASSERT_TRUE(st.ok()) << st.to_string(); + + rapidjson::Document doc; + doc.Parse(request_body.c_str()); + ASSERT_FALSE(doc.HasParseError()) << request_body; + ASSERT_TRUE(doc.IsObject()); + ASSERT_TRUE(doc.HasMember("model")); + ASSERT_STREQ(doc["model"].GetString(), "tongyi-embedding-vision-plus"); + ASSERT_TRUE(doc.HasMember("input")); + ASSERT_TRUE(doc["input"].HasMember("contents")); + const auto& contents = doc["input"]["contents"]; + ASSERT_TRUE(contents.IsArray()); + ASSERT_EQ(contents.Size(), 1); + ASSERT_TRUE(contents[0].HasMember("video")); + ASSERT_STREQ(contents[0]["video"].GetString(), "https://a/b/c.mp4"); + ASSERT_TRUE(doc.HasMember("parameters")); + ASSERT_TRUE(doc["parameters"].IsObject()); + ASSERT_TRUE(doc["parameters"].HasMember("dimension")); + ASSERT_EQ(doc["parameters"]["dimension"].GetInt(), 1024); +} + +TEST(AI_ADAPTER_TEST, qwen_multimodal_embedding_batch_request) { + QwenAdapter adapter; + TAIResource config; + config.model_name = "tongyi-embedding-vision-plus"; + config.dimensions = 1024; + adapter.init(config); + + std::vector media_types = {MultimodalType::IMAGE, MultimodalType::VIDEO}; + std::vector media_urls = {"https://a/b/c.png", "https://a/b/c.mp4"}; + std::string request_body; + Status st = + adapter.build_multimodal_embedding_request(media_types, media_urls, {}, request_body); + ASSERT_TRUE(st.ok()) << st.to_string(); + + rapidjson::Document doc; + doc.Parse(request_body.c_str()); + ASSERT_FALSE(doc.HasParseError()) << request_body; + ASSERT_TRUE(doc.IsObject()); + ASSERT_TRUE(doc.HasMember("input")); + ASSERT_TRUE(doc["input"].HasMember("contents")); + const auto& contents = doc["input"]["contents"]; + ASSERT_TRUE(contents.IsArray()); + ASSERT_EQ(contents.Size(), 2); + ASSERT_TRUE(contents[0].HasMember("image")); + ASSERT_STREQ(contents[0]["image"].GetString(), "https://a/b/c.png"); + ASSERT_TRUE(contents[1].HasMember("video")); + ASSERT_STREQ(contents[1]["video"].GetString(), "https://a/b/c.mp4"); +} + +TEST(AI_ADAPTER_TEST, qwen_multimodal_embedding_request_audio_not_supported) { + QwenAdapter adapter; + TAIResource config; + config.model_name = "tongyi-embedding-vision-plus"; + adapter.init(config); + + std::string request_body; + Status st = adapter.build_multimodal_embedding_request({MultimodalType::AUDIO}, + {"https://a/b/c.mp3"}, {}, request_body); + ASSERT_FALSE(st.ok()); + ASSERT_THAT(st.to_string(), + ::testing::HasSubstr("QWEN only supports image/video multimodal embed")); + ASSERT_THAT(st.to_string(), ::testing::HasSubstr("audio")); +} + +TEST(AI_ADAPTER_TEST, voyage_multimodal_embedding_request) { + VoyageAIAdapter adapter; + TAIResource config; + config.model_name = "voyage-multimodal-3.5"; + config.dimensions = 2048; + adapter.init(config); + + std::string request_body; + Status st = adapter.build_multimodal_embedding_request({MultimodalType::VIDEO}, + {"https://a/b/c.mp4"}, {}, request_body); + ASSERT_TRUE(st.ok()) << st.to_string(); + + rapidjson::Document doc; + doc.Parse(request_body.c_str()); + ASSERT_FALSE(doc.HasParseError()) << request_body; + ASSERT_TRUE(doc.IsObject()); + ASSERT_TRUE(doc.HasMember("inputs")); + const auto& inputs = doc["inputs"]; + ASSERT_TRUE(inputs.IsArray()); + ASSERT_EQ(inputs.Size(), 1); + ASSERT_TRUE(inputs[0].HasMember("content")); + const auto& content = inputs[0]["content"]; + ASSERT_TRUE(content.IsArray()); + ASSERT_EQ(content.Size(), 1); + ASSERT_TRUE(content[0].HasMember("type")); + ASSERT_STREQ(content[0]["type"].GetString(), "video_url"); + ASSERT_TRUE(content[0].HasMember("video_url")); + ASSERT_STREQ(content[0]["video_url"].GetString(), "https://a/b/c.mp4"); + ASSERT_FALSE(doc.HasMember("dimensions")); + ASSERT_FALSE(doc.HasMember("output_dimension")); +} + +TEST(AI_ADAPTER_TEST, voyage_multimodal_embedding_batch_request) { + VoyageAIAdapter adapter; + TAIResource config; + config.model_name = "voyage-multimodal-3.5"; + adapter.init(config); + + std::vector media_types = {MultimodalType::IMAGE, MultimodalType::VIDEO}; + std::vector media_urls = {"https://a/b/c.png", "https://a/b/c.mp4"}; + std::string request_body; + Status st = + adapter.build_multimodal_embedding_request(media_types, media_urls, {}, request_body); + ASSERT_TRUE(st.ok()) << st.to_string(); + + rapidjson::Document doc; + doc.Parse(request_body.c_str()); + ASSERT_FALSE(doc.HasParseError()) << request_body; + ASSERT_TRUE(doc.IsObject()); + ASSERT_TRUE(doc.HasMember("inputs")); + const auto& request_inputs = doc["inputs"]; + ASSERT_TRUE(request_inputs.IsArray()); + ASSERT_EQ(request_inputs.Size(), 2); + ASSERT_TRUE(request_inputs[0]["content"].IsArray()); + ASSERT_STREQ(request_inputs[0]["content"][0]["type"].GetString(), "image_url"); + ASSERT_STREQ(request_inputs[0]["content"][0]["image_url"].GetString(), "https://a/b/c.png"); + ASSERT_TRUE(request_inputs[1]["content"].IsArray()); + ASSERT_STREQ(request_inputs[1]["content"][0]["type"].GetString(), "video_url"); + ASSERT_STREQ(request_inputs[1]["content"][0]["video_url"].GetString(), "https://a/b/c.mp4"); +} + +TEST(AI_ADAPTER_TEST, jina_multimodal_embedding_request) { + JinaAdapter adapter; + TAIResource config; + config.model_name = "jina-embeddings-v4"; + config.dimensions = 512; + adapter.init(config); + + std::string request_body; + Status st = adapter.build_multimodal_embedding_request({MultimodalType::IMAGE}, + {"https://a/b/c.jpg"}, {}, request_body); + ASSERT_TRUE(st.ok()) << st.to_string(); + + rapidjson::Document doc; + doc.Parse(request_body.c_str()); + ASSERT_FALSE(doc.HasParseError()) << request_body; + ASSERT_TRUE(doc.IsObject()); + ASSERT_TRUE(doc.HasMember("task")); + ASSERT_STREQ(doc["task"].GetString(), "text-matching"); + ASSERT_TRUE(doc.HasMember("input")); + const auto& input = doc["input"]; + ASSERT_TRUE(input.IsArray()); + ASSERT_EQ(input.Size(), 1); + ASSERT_TRUE(input[0].HasMember("image")); + ASSERT_STREQ(input[0]["image"].GetString(), "https://a/b/c.jpg"); + ASSERT_TRUE(doc.HasMember("dimensions")); + ASSERT_EQ(doc["dimensions"].GetInt(), 512); +} + +TEST(AI_ADAPTER_TEST, jina_multimodal_embedding_batch_request) { + JinaAdapter adapter; + TAIResource config; + config.model_name = "jina-embeddings-v4"; + config.dimensions = 512; + adapter.init(config); + + std::vector media_types = {MultimodalType::IMAGE, MultimodalType::VIDEO}; + std::vector media_urls = {"https://a/b/c.jpg", "https://a/b/c.mp4"}; + std::string request_body; + Status st = + adapter.build_multimodal_embedding_request(media_types, media_urls, {}, request_body); + ASSERT_TRUE(st.ok()) << st.to_string(); + + rapidjson::Document doc; + doc.Parse(request_body.c_str()); + ASSERT_FALSE(doc.HasParseError()) << request_body; + ASSERT_TRUE(doc.IsObject()); + ASSERT_TRUE(doc.HasMember("input")); + const auto& input = doc["input"]; + ASSERT_TRUE(input.IsArray()); + ASSERT_EQ(input.Size(), 2); + ASSERT_TRUE(input[0].HasMember("image")); + ASSERT_STREQ(input[0]["image"].GetString(), "https://a/b/c.jpg"); + ASSERT_TRUE(input[1].HasMember("video")); + ASSERT_STREQ(input[1]["video"].GetString(), "https://a/b/c.mp4"); +} + +TEST(AI_ADAPTER_TEST, multimodal_provider_support) { + OpenAIAdapter openai_adapter; + TAIResource openai_config; + openai_config.provider_type = "OPENAI"; + openai_adapter.init(openai_config); + + std::string request_body; + Status st = openai_adapter.build_multimodal_embedding_request( + {MultimodalType::IMAGE}, {"https://a/b/c.png"}, {}, request_body); + ASSERT_FALSE(st.ok()); + ASSERT_THAT(st.to_string(), ::testing::HasSubstr("does not support multimodal Embed")); +} + +TEST(AI_ADAPTER_TEST, gemini_multimodal_embedding_request) { + GeminiAdapter gemini_adapter; + TAIResource gemini_config; + gemini_config.provider_type = "GEMINI"; + gemini_config.model_name = "gemini-embedding-2-preview"; + gemini_config.dimensions = 768; + gemini_adapter.init(gemini_config); + + struct GeminiMultimodalCase { + MultimodalType media_type; + const char* media_url; + const char* mime_type; + }; + const std::vector test_cases = { + {MultimodalType::IMAGE, "https://a/b/c.jpg", "image/jpeg"}, + {MultimodalType::IMAGE, "https://a/b/c.webp", "image/webp"}, + {MultimodalType::AUDIO, "https://a/b/c.wav", "audio/wav"}, + {MultimodalType::VIDEO, "https://a/b/c.webm", "video/webm"}, + }; + + for (const auto& test_case : test_cases) { + std::string request_body; + Status st = gemini_adapter.build_multimodal_embedding_request( + {test_case.media_type}, {test_case.media_url}, {test_case.mime_type}, request_body); + ASSERT_TRUE(st.ok()) << st.to_string(); + + rapidjson::Document doc; + doc.Parse(request_body.c_str()); + ASSERT_FALSE(doc.HasParseError()) << request_body; + ASSERT_TRUE(doc.IsObject()); + ASSERT_TRUE(doc.HasMember("requests")); + ASSERT_TRUE(doc["requests"].IsArray()); + ASSERT_EQ(doc["requests"].Size(), 1); + const auto& request = doc["requests"][0]; + ASSERT_TRUE(request.HasMember("model")); + ASSERT_STREQ(request["model"].GetString(), "models/gemini-embedding-2-preview"); + ASSERT_TRUE(request.HasMember("outputDimensionality")); + ASSERT_EQ(request["outputDimensionality"].GetInt(), 768); + ASSERT_TRUE(request.HasMember("content")); + ASSERT_TRUE(request["content"].HasMember("parts")); + ASSERT_TRUE(request["content"]["parts"].IsArray()); + ASSERT_EQ(request["content"]["parts"].Size(), 1); + ASSERT_TRUE(request["content"]["parts"][0].HasMember("file_data")); + ASSERT_TRUE(request["content"]["parts"][0]["file_data"].IsObject()); + ASSERT_STREQ(request["content"]["parts"][0]["file_data"]["mime_type"].GetString(), + test_case.mime_type); + ASSERT_STREQ(request["content"]["parts"][0]["file_data"]["file_uri"].GetString(), + test_case.media_url); + } +} + +TEST(AI_ADAPTER_TEST, gemini_multimodal_embedding_batch_request) { + GeminiAdapter adapter; + TAIResource config; + config.provider_type = "GEMINI"; + config.model_name = "gemini-embedding-2-preview"; + config.dimensions = 768; + adapter.init(config); + + std::vector media_types = {MultimodalType::IMAGE, MultimodalType::AUDIO, + MultimodalType::VIDEO}; + std::vector media_urls = {"https://a/b/c.jpg", "https://a/b/c.wav", + "https://a/b/c.webm"}; + std::vector media_content_types = {"image/jpeg", "audio/wav", "video/webm"}; + std::string request_body; + Status st = adapter.build_multimodal_embedding_request(media_types, media_urls, + media_content_types, request_body); + ASSERT_TRUE(st.ok()) << st.to_string(); + + rapidjson::Document doc; + doc.Parse(request_body.c_str()); + ASSERT_FALSE(doc.HasParseError()) << request_body; + ASSERT_TRUE(doc.IsObject()); + ASSERT_TRUE(doc.HasMember("requests")); + const auto& requests = doc["requests"]; + ASSERT_TRUE(requests.IsArray()); + ASSERT_EQ(requests.Size(), 3); + + ASSERT_STREQ(requests[0]["model"].GetString(), "models/gemini-embedding-2-preview"); + ASSERT_EQ(requests[0]["outputDimensionality"].GetInt(), 768); + ASSERT_STREQ(requests[0]["content"]["parts"][0]["file_data"]["mime_type"].GetString(), + "image/jpeg"); + ASSERT_STREQ(requests[0]["content"]["parts"][0]["file_data"]["file_uri"].GetString(), + "https://a/b/c.jpg"); + + ASSERT_STREQ(requests[1]["content"]["parts"][0]["file_data"]["mime_type"].GetString(), + "audio/wav"); + ASSERT_STREQ(requests[1]["content"]["parts"][0]["file_data"]["file_uri"].GetString(), + "https://a/b/c.wav"); + + ASSERT_STREQ(requests[2]["content"]["parts"][0]["file_data"]["mime_type"].GetString(), + "video/webm"); + ASSERT_STREQ(requests[2]["content"]["parts"][0]["file_data"]["file_uri"].GetString(), + "https://a/b/c.webm"); +} + +TEST(AI_ADAPTER_TEST, gemini_multimodal_embedding_request_empty_inputs) { + GeminiAdapter adapter; + TAIResource config; + config.provider_type = "GEMINI"; + config.model_name = "gemini-embedding-2-preview"; + adapter.init(config); + + std::string request_body; + Status st = adapter.build_multimodal_embedding_request({}, {}, {}, request_body); + ASSERT_FALSE(st.ok()); + ASSERT_THAT(st.to_string(), + ::testing::HasSubstr("Gemini multimodal embed inputs can not be empty")); +} + +TEST(AI_ADAPTER_TEST, gemini_multimodal_embedding_request_size_mismatch) { + GeminiAdapter adapter; + TAIResource config; + config.provider_type = "GEMINI"; + config.model_name = "gemini-embedding-2-preview"; + adapter.init(config); + + std::string request_body; + Status st = adapter.build_multimodal_embedding_request( + {MultimodalType::IMAGE, MultimodalType::VIDEO}, {"https://a/b/c.png"}, {}, + request_body); + ASSERT_FALSE(st.ok()); + ASSERT_THAT( + st.to_string(), + ::testing::HasSubstr( + "Gemini multimodal embed input size mismatch, media_types=2, media_urls=1")); +} + +TEST(AI_ADAPTER_TEST, gemini_multimodal_embedding_content_type_size_mismatch) { + GeminiAdapter adapter; + TAIResource config; + config.provider_type = "GEMINI"; + config.model_name = "gemini-embedding-2-preview"; + adapter.init(config); + + std::string request_body; + Status st = adapter.build_multimodal_embedding_request({MultimodalType::IMAGE}, + {"https://a/b/c.jpg"}, {}, request_body); + ASSERT_FALSE(st.ok()); + ASSERT_THAT(st.to_string(), ::testing::HasSubstr("Gemini multimodal embed input size mismatch, " + "media_content_types=0, media_urls=1")); +} + +TEST(AI_ADAPTER_TEST, gemini_parse_batch_embedding_response) { + GeminiAdapter adapter; + std::string resp = R"({ + "embeddings": [ + {"values": [0.1, 0.2, 0.3]}, + {"values": [0.4, 0.5]} + ] + })"; + + std::vector> results; + Status st = adapter.parse_embedding_response(resp, results); + ASSERT_TRUE(st.ok()) << st.to_string(); + ASSERT_EQ(results.size(), 2); + ASSERT_EQ(results[0].size(), 3); + ASSERT_EQ(results[1].size(), 2); + ASSERT_FLOAT_EQ(results[0][0], 0.1F); + ASSERT_FLOAT_EQ(results[0][2], 0.3F); + ASSERT_FLOAT_EQ(results[1][0], 0.4F); + ASSERT_FLOAT_EQ(results[1][1], 0.5F); +} + } // namespace doris diff --git a/be/test/ai/ai_function_test.cpp b/be/test/ai/ai_function_test.cpp index 5a8777f7377bbd..976eaddd5bcb40 100644 --- a/be/test/ai/ai_function_test.cpp +++ b/be/test/ai/ai_function_test.cpp @@ -15,17 +15,24 @@ // specific language governing permissions and limitations // under the License. +#include #include #include +#include +#include +#include #include +#include #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" #include "core/data_type/data_type_array.h" +#include "exprs/function/ai/ai_adapter.h" #include "exprs/function/ai/ai_classify.h" #include "exprs/function/ai/ai_extract.h" #include "exprs/function/ai/ai_filter.h" @@ -37,22 +44,160 @@ #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" namespace doris { +class FunctionAITransportTestHelper : public FunctionAISentiment { +public: + using FunctionAISentiment::do_send_request; +}; + +class FunctionAIFilterBatchTestHelper : public AIFunction { +public: + friend class AIFunction; + + static constexpr auto name = "ai_filter"; + static constexpr auto system_prompt = FunctionAIFilter::system_prompt; + static constexpr size_t number_of_arguments = FunctionAIFilter::number_of_arguments; + + using AIFunction::execute_batch_request; + + DataTypePtr get_nested_return_type_impl(const DataTypes& /*arguments*/) const { + return std::make_shared(); + } + + MutableColumnPtr create_result_column() const { return ColumnUInt8::create(); } + + Status append_batch_results(const std::vector& batch_results, + IColumn& col_result) const { + auto& bool_col = assert_cast(col_result); + for (const auto& batch_result : batch_results) { + std::string_view trimmed = doris::trim(batch_result); + if (trimmed != "1" && trimmed != "0") { + return Status::RuntimeError("Failed to parse boolean value: " + + std::string(trimmed)); + } + bool_col.insert_value(static_cast(trimmed == "1")); + } + return Status::OK(); + } +}; + +class OneShotHttpServer { +public: + OneShotHttpServer(int status_code, std::string response_body) + : _status_code(status_code), _response_body(std::move(response_body)) { + _listen_fd = socket(AF_INET, SOCK_STREAM, 0); + DCHECK_GE(_listen_fd, 0); + + int reuse_addr = 1; + int ret = setsockopt(_listen_fd, SOL_SOCKET, SO_REUSEADDR, &reuse_addr, sizeof(reuse_addr)); + DCHECK_EQ(ret, 0); + + sockaddr_in addr {}; + addr.sin_family = AF_INET; + addr.sin_addr.s_addr = htonl(INADDR_LOOPBACK); + addr.sin_port = 0; + ret = bind(_listen_fd, reinterpret_cast(&addr), sizeof(addr)); + DCHECK_EQ(ret, 0); + ret = listen(_listen_fd, 1); + DCHECK_EQ(ret, 0); + + socklen_t addr_len = sizeof(addr); + ret = getsockname(_listen_fd, reinterpret_cast(&addr), &addr_len); + DCHECK_EQ(ret, 0); + _port = ntohs(addr.sin_port); + + _thread = std::thread([this] { _serve_once(); }); + } + + ~OneShotHttpServer() { + if (_thread.joinable()) { + _thread.join(); + } + if (_listen_fd >= 0) { + close(_listen_fd); + } + } + + std::string endpoint() const { return "http://127.0.0.1:" + std::to_string(_port); } + + std::string join_and_get_request() { + if (_thread.joinable()) { + _thread.join(); + } + return _request; + } + +private: + void _serve_once() { + int client_fd = accept(_listen_fd, nullptr, nullptr); + DCHECK_GE(client_fd, 0); + + char buffer[4096]; + while (true) { + ssize_t read_bytes = recv(client_fd, buffer, sizeof(buffer), 0); + if (read_bytes <= 0) { + break; + } + _request.append(buffer, read_bytes); + if (_request.find("\r\n\r\n") != std::string::npos) { + auto header_end = _request.find("\r\n\r\n"); + size_t body_length = 0; + size_t length_pos = _request.find("Content-Length:"); + if (length_pos != std::string::npos) { + size_t value_begin = length_pos + sizeof("Content-Length:") - 1; + size_t value_end = _request.find("\r\n", value_begin); + std::string length_str = _request.substr(value_begin, value_end - value_begin); + body_length = std::stoul(length_str); + } + if (_request.size() >= header_end + 4 + body_length) { + break; + } + } + } + + std::string status_text = _status_code == 200 ? "OK" : "Internal Server Error"; + std::string response = "HTTP/1.1 " + std::to_string(_status_code) + " " + status_text + + "\r\nContent-Type: application/json\r\nContent-Length: " + + std::to_string(_response_body.size()) + + "\r\nConnection: close\r\n\r\n" + _response_body; + size_t sent_bytes = 0; + while (sent_bytes < response.size()) { + ssize_t sent = + send(client_fd, response.data() + sent_bytes, response.size() - sent_bytes, 0); + DCHECK_GT(sent, 0); + sent_bytes += sent; + } + shutdown(client_fd, SHUT_RDWR); + close(client_fd); + } + + int _listen_fd = -1; + uint16_t _port = 0; + int _status_code; + std::string _response_body; + std::string _request; + std::thread _thread; +}; + namespace { -MutableColumnPtr create_string_array_column(const std::vector>& rows) { +MutableColumnPtr create_string_array_column(const std::vector>& rows, + const std::vector& null_map = {}) { auto nested_column = ColumnString::create(); auto null_map_column = ColumnUInt8::create(); auto offsets_column = ColumnOffset64::create(); IColumn::Offset offset = 0; + size_t element = 0; for (const auto& row : rows) { for (const auto& value : row) { nested_column->insert_data(value.data(), value.size()); - null_map_column->insert_value(0); + null_map_column->insert_value(null_map.empty() ? 0 : null_map[element]); + ++element; } offset += row.size(); offsets_column->insert_value(offset); @@ -62,7 +207,23 @@ MutableColumnPtr create_string_array_column(const std::vector texts = {"good product"}; + auto labels = create_string_array_column({{"positive", "unused-null", "negative"}}, + std::vector {0, 1, 0}); + + Block block; + block.insert({ColumnHelper::create_column({"resource_name"}), + std::make_shared(), "resource"}); + block.insert({ColumnHelper::create_column(texts), + std::make_shared(), "text"}); + block.insert({std::move(labels), + std::make_shared(std::make_shared()), "labels"}); + + Columns prompt_columns = get_prompt_columns(block, {0, 1, 2}); + std::string prompt; + const std::string expected = + "Labels: [\"positive\", \"negative\"]\n" + "Text: good product"; + + FunctionAIClassify classify; + ASSERT_TRUE(classify.build_prompt(prompt_columns, 0, prompt).ok()); + EXPECT_EQ(prompt, expected); + + FunctionAIExtract extract; + ASSERT_TRUE(extract.build_prompt(prompt_columns, 0, prompt).ok()); + EXPECT_EQ(prompt, expected); + + FunctionAIMask mask; + ASSERT_TRUE(mask.build_prompt(prompt_columns, 0, prompt).ok()); + EXPECT_EQ(prompt, expected); +} + TEST(AIFunctionTest, AITranslateTest) { FunctionAITranslate function; @@ -247,7 +440,7 @@ TEST(AIFunctionTest, AITranslateTest) { ColumnNumbers arguments = {0, 1, 2}; std::string prompt; - Status status = function.build_prompt(block, arguments, 0, prompt); + Status status = function.build_prompt(get_prompt_columns(block, arguments), 0, prompt); ASSERT_TRUE(status.ok()); ASSERT_EQ(prompt, @@ -273,7 +466,7 @@ TEST(AIFunctionTest, AISimilarityTest) { ColumnNumbers arguments = {0, 1, 2}; std::string prompt; - Status status = function.build_prompt(block, arguments, 0, prompt); + Status status = function.build_prompt(get_prompt_columns(block, arguments), 0, prompt); ASSERT_TRUE(status.ok()); ASSERT_EQ(prompt, "Text 1: I like this dish\nText 2: This dish is very good"); @@ -282,6 +475,7 @@ TEST(AIFunctionTest, AISimilarityTest) { TEST(AIFunctionTest, AISimilarityExecuteTest) { auto runtime_state = std::make_unique(); auto ctx = FunctionContext::create_context(runtime_state.get(), {}, {}); + setenv("AI_TEST_RESULT", R"(["0.5"])", 1); std::vector resources = {"mock_resource"}; std::vector text1 = {"I like this dish"}; @@ -304,6 +498,7 @@ TEST(AIFunctionTest, AISimilarityExecuteTest) { similarity_func->execute_impl(ctx.get(), block, arguments, result_idx, text1.size()); ASSERT_TRUE(exec_status.ok()); + unsetenv("AI_TEST_RESULT"); } TEST(AIFunctionTest, AISimilarityTrimWhitespace) { @@ -311,10 +506,19 @@ TEST(AIFunctionTest, AISimilarityTrimWhitespace) { auto ctx = FunctionContext::create_context(runtime_state.get(), {}, {}); std::vector> test_cases = { - {"0.5", 0.5f}, {"1.0", 1.0f}, {"0.0", 0.0f}, {" 0.5", 0.5f}, - {"0.5 ", 0.5f}, {" 0.5 ", 0.5f}, {"\n0.8", 0.8f}, {"0.3\n", 0.3f}, - {"\n0.7\n", 0.7f}, {"\t0.2\t", 0.2f}, {" \n\t0.9 \n\t", 0.9f}, {" 0.1 ", 0.1f}, - {"\r\n0.6\r\n", 0.6f}}; + {R"(["0.5"])", 0.5f}, + {R"(["1.0"])", 1.0f}, + {R"(["0.0"])", 0.0f}, + {" " + std::string(R"(["0.5"])"), 0.5f}, + {std::string(R"(["0.5"])") + " ", 0.5f}, + {" " + std::string(R"(["0.5"])") + " ", 0.5f}, + {"\n" + std::string(R"(["0.8"])"), 0.8f}, + {std::string(R"(["0.3"])") + "\n", 0.3f}, + {"\n" + std::string(R"(["0.7"])") + "\n", 0.7f}, + {"\t" + std::string(R"(["0.2"])") + "\t", 0.2f}, + {" \n\t" + std::string(R"(["0.9"])") + " \n\t", 0.9f}, + {" " + std::string(R"(["0.1"])") + " ", 0.1f}, + {"\r\n" + std::string(R"(["0.6"])") + "\r\n", 0.6f}}; for (const auto& test_case : test_cases) { setenv("AI_TEST_RESULT", test_case.first.c_str(), 1); @@ -352,6 +556,80 @@ TEST(AIFunctionTest, AISimilarityTrimWhitespace) { unsetenv("AI_TEST_RESULT"); } +TEST(AIFunctionTest, AISimilarityInvalidValue) { + auto runtime_state = std::make_unique(); + auto ctx = FunctionContext::create_context(runtime_state.get(), {}, {}); + + std::vector invalid_cases = {"abc", "1.2x", "", " ", "\n\t", "1,2"}; + + for (const auto& invalid_value : invalid_cases) { + setenv("AI_TEST_RESULT", invalid_value.c_str(), 1); + + std::vector resources = {"mock_resource"}; + std::vector text1 = {"Test text 1"}; + std::vector text2 = {"Test text 2"}; + auto col_resource = ColumnHelper::create_column(resources); + auto col_text1 = ColumnHelper::create_column(text1); + auto col_text2 = ColumnHelper::create_column(text2); + + Block block; + block.insert({std::move(col_resource), std::make_shared(), "resource"}); + block.insert({std::move(col_text1), std::make_shared(), "text1"}); + block.insert({std::move(col_text2), std::make_shared(), "text2"}); + block.insert({nullptr, std::make_shared(), "result"}); + + ColumnNumbers arguments = {0, 1, 2}; + size_t result_idx = 3; + + auto similarity_func = FunctionAISimilarity::create(); + Status exec_status = similarity_func->execute_impl(ctx.get(), block, arguments, result_idx, + text1.size()); + + ASSERT_FALSE(exec_status.ok()) + << "Should have failed for invalid value: '" << invalid_value << "'"; + ASSERT_NE(exec_status.to_string().find("Failed to parse float value"), std::string::npos); + } + unsetenv("AI_TEST_RESULT"); +} + +TEST(AIFunctionTest, AISimilarityBatchExecuteTest) { + auto runtime_state = std::make_unique(); + auto ctx = FunctionContext::create_context(runtime_state.get(), {}, {}); + + setenv("AI_TEST_RESULT", R"(["0.5","1.0","0.0"])", 1); + + std::vector resources = {"mock_resource"}; + std::vector text1 = {"a", "b", "c"}; + std::vector text2 = {"d", "e", "f"}; + auto col_resource = ColumnHelper::create_column(resources); + auto col_text1 = ColumnHelper::create_column(text1); + auto col_text2 = ColumnHelper::create_column(text2); + + Block block; + block.insert({std::move(col_resource), std::make_shared(), "resource"}); + block.insert({std::move(col_text1), std::make_shared(), "text1"}); + block.insert({std::move(col_text2), std::make_shared(), "text2"}); + block.insert({nullptr, std::make_shared(), "result"}); + + ColumnNumbers arguments = {0, 1, 2}; + size_t result_idx = 3; + + auto similarity_func = FunctionAISimilarity::create(); + Status exec_status = + similarity_func->execute_impl(ctx.get(), block, arguments, result_idx, text1.size()); + + ASSERT_TRUE(exec_status.ok()); + + const auto& res_col = + assert_cast(*block.get_by_position(result_idx).column); + ASSERT_EQ(res_col.size(), 3); + EXPECT_FLOAT_EQ(res_col.get_data()[0], 0.5f); + EXPECT_FLOAT_EQ(res_col.get_data()[1], 1.0f); + EXPECT_FLOAT_EQ(res_col.get_data()[2], 0.0f); + + unsetenv("AI_TEST_RESULT"); +} + TEST(AIFunctionTest, AIFilterTest) { FunctionAIFilter function; @@ -367,7 +645,7 @@ TEST(AIFunctionTest, AIFilterTest) { ColumnNumbers arguments = {0, 1}; std::string prompt; - Status status = function.build_prompt(block, arguments, 0, prompt); + Status status = function.build_prompt(get_prompt_columns(block, arguments), 0, prompt); ASSERT_TRUE(status.ok()); ASSERT_EQ(prompt, "This is a valid sentence."); @@ -376,6 +654,7 @@ TEST(AIFunctionTest, AIFilterTest) { TEST(AIFunctionTest, AIFilterExecuteTest) { auto runtime_state = std::make_unique(); auto ctx = FunctionContext::create_context(runtime_state.get(), {}, {}); + setenv("AI_TEST_RESULT", R"(["0"])", 1); std::vector resources = {"mock_resource"}; std::vector texts = {"This is a valid sentence."}; @@ -394,10 +673,46 @@ TEST(AIFunctionTest, AIFilterExecuteTest) { Status exec_status = filter_func->execute_impl(ctx.get(), block, arguments, result_idx, texts.size()); + ASSERT_TRUE(exec_status.ok()); + const auto& res_col = assert_cast(*block.get_by_position(result_idx).column); UInt8 val = res_col.get_data()[0]; ASSERT_TRUE(val == 0); + unsetenv("AI_TEST_RESULT"); +} + +TEST(AIFunctionTest, AIFilterExecuteMultipleRows) { + auto runtime_state = std::make_unique(); + auto ctx = FunctionContext::create_context(runtime_state.get(), {}, {}); + + setenv("AI_TEST_RESULT", R"(["1","1"])", 1); + + std::vector resources = {"mock_resource", "mock_resource"}; + std::vector texts = {"This is valid.", "This is also valid."}; + auto col_resource = ColumnHelper::create_column(resources); + auto col_text = ColumnHelper::create_column(texts); + + Block block; + block.insert({std::move(col_resource), std::make_shared(), "resource"}); + block.insert({std::move(col_text), std::make_shared(), "text"}); + block.insert({nullptr, std::make_shared(), "result"}); + + ColumnNumbers arguments = {0, 1}; + size_t result_idx = 2; + + auto filter_func = FunctionAIFilter::create(); + Status exec_status = + filter_func->execute_impl(ctx.get(), block, arguments, result_idx, texts.size()); + + unsetenv("AI_TEST_RESULT"); + + ASSERT_TRUE(exec_status.ok()) << exec_status.to_string(); + const auto& res_col = + assert_cast(*block.get_by_position(result_idx).column); + ASSERT_EQ(res_col.size(), 2); + ASSERT_EQ(res_col.get_data()[0], 1); + ASSERT_EQ(res_col.get_data()[1], 1); } TEST(AIFunctionTest, AIFilterTrimWhitespace) { @@ -405,9 +720,18 @@ TEST(AIFunctionTest, AIFilterTrimWhitespace) { auto ctx = FunctionContext::create_context(runtime_state.get(), {}, {}); std::vector> test_cases = { - {"0", 0}, {"1", 1}, {" 0", 0}, {"0 ", 0}, - {" 0 ", 0}, {"\n0", 0}, {"0\n", 0}, {"\n0\n", 0}, - {"\t1\t", 1}, {" \n\t1 \n\t", 1}, {" 1 ", 1}, {"\r\n0\r\n", 0}}; + {R"(["0"])", 0}, + {R"(["1"])", 1}, + {" " + std::string(R"(["0"])"), 0}, + {std::string(R"(["0"])") + " ", 0}, + {" " + std::string(R"(["0"])") + " ", 0}, + {"\n" + std::string(R"(["0"])"), 0}, + {std::string(R"(["0"])") + "\n", 0}, + {"\n" + std::string(R"(["0"])") + "\n", 0}, + {"\t" + std::string(R"(["1"])") + "\t", 1}, + {" \n\t" + std::string(R"(["1"])") + " \n\t", 1}, + {" " + std::string(R"(["1"])") + " ", 1}, + {"\r\n" + std::string(R"(["0"])") + "\r\n", 0}}; for (const auto& test_case : test_cases) { setenv("AI_TEST_RESULT", test_case.first.c_str(), 1); @@ -447,8 +771,10 @@ TEST(AIFunctionTest, AIFilterInvalidValue) { auto ctx = FunctionContext::create_context(runtime_state.get(), {}, {}); std::vector invalid_cases = { - "2", "maybe", "ok", "", " ", "01", "0.5", "sure", "truee", "falsee", - "yess", "noo", "true", "false", "yes", "no", "TRUE", "FALSE", "YES", "NO"}; + R"(["2"])", R"(["maybe"])", R"(["ok"])", R"([""])", R"(["01"])", + R"(["0.5"])", R"(["sure"])", R"(["truee"])", R"(["falsee"])", R"(["yess"])", + R"(["noo"])", R"(["true"])", R"(["false"])", R"(["yes"])", R"(["no"])", + R"(["TRUE"])", R"(["FALSE"])", R"(["YES"])", "[\"NO\"]"}; for (const auto& invalid_value : invalid_cases) { setenv("AI_TEST_RESULT", invalid_value.c_str(), 1); @@ -481,101 +807,158 @@ TEST(AIFunctionTest, AIFilterInvalidValue) { unsetenv("AI_TEST_RESULT"); } -TEST(AIFunctionTest, ResourceNotFound) { +TEST(AIFunctionTest, AIFilterBatchExecuteTest) { auto runtime_state = std::make_unique(); auto ctx = FunctionContext::create_context(runtime_state.get(), {}, {}); - std::vector resources = {"not_exist_resource"}; - std::vector texts = {"test"}; + setenv("AI_TEST_RESULT", R"(["1","0","1"])", 1); + + std::vector resources = {"mock_resource"}; + std::vector texts = {"valid text", "invalid text", "valid again"}; auto col_resource = ColumnHelper::create_column(resources); auto col_text = ColumnHelper::create_column(texts); Block block; block.insert({std::move(col_resource), std::make_shared(), "resource"}); block.insert({std::move(col_text), std::make_shared(), "text"}); - block.insert({nullptr, std::make_shared(), "result"}); + block.insert({nullptr, std::make_shared(), "result"}); ColumnNumbers arguments = {0, 1}; size_t result_idx = 2; - auto sentiment_func = FunctionAISentiment::create(); + auto filter_func = FunctionAIFilter::create(); Status exec_status = - sentiment_func->execute_impl(ctx.get(), block, arguments, result_idx, texts.size()); + filter_func->execute_impl(ctx.get(), block, arguments, result_idx, texts.size()); - ASSERT_FALSE(exec_status.ok()); - ASSERT_TRUE(exec_status.to_string().find("AI resource not found") != std::string::npos); + ASSERT_TRUE(exec_status.ok()); + + const auto& res_col = + assert_cast(*block.get_by_position(result_idx).column); + ASSERT_EQ(res_col.size(), 3); + EXPECT_EQ(res_col.get_data()[0], 1); + EXPECT_EQ(res_col.get_data()[1], 0); + EXPECT_EQ(res_col.get_data()[2], 1); + + unsetenv("AI_TEST_RESULT"); } -TEST(AIFunctionTest, MockResourceSendRequest) { +TEST(AIFunctionTest, AIFilterBatchLengthMismatch) { auto runtime_state = std::make_unique(); auto ctx = FunctionContext::create_context(runtime_state.get(), {}, {}); + setenv("AI_TEST_RESULT", R"(["1","0"])", 1); + std::vector resources = {"mock_resource"}; - std::vector texts = {"test input"}; + std::vector texts = {"row1", "row2", "row3"}; auto col_resource = ColumnHelper::create_column(resources); auto col_text = ColumnHelper::create_column(texts); Block block; block.insert({std::move(col_resource), std::make_shared(), "resource"}); block.insert({std::move(col_text), std::make_shared(), "text"}); - block.insert({nullptr, std::make_shared(), "result"}); + block.insert({nullptr, std::make_shared(), "result"}); ColumnNumbers arguments = {0, 1}; size_t result_idx = 2; - auto sentiment_func = FunctionAISentiment::create(); + auto filter_func = FunctionAIFilter::create(); Status exec_status = - sentiment_func->execute_impl(ctx.get(), block, arguments, result_idx, texts.size()); + filter_func->execute_impl(ctx.get(), block, arguments, result_idx, texts.size()); - ASSERT_TRUE(exec_status.ok()) << exec_status.to_string(); - const auto& res_col = - assert_cast(*block.get_by_position(result_idx).column); - StringRef ref = res_col.get_data_at(0); - std::string val(ref.data, ref.size); - ASSERT_EQ(val, "this is a mock response. test input"); -} + ASSERT_FALSE(exec_status.ok()); + ASSERT_TRUE(exec_status.to_string().find("expected 3 items but got 2") != std::string::npos); -TEST(AIFunctionTest, MockResourceBatchStringResult) { - setenv("AI_TEST_RESULT", R"(["first result","second result"])", 1); + unsetenv("AI_TEST_RESULT"); +} +TEST(AIFunctionTest, AIFilterBatchInvalidJson) { auto runtime_state = std::make_unique(); auto ctx = FunctionContext::create_context(runtime_state.get(), {}, {}); - std::vector resources = {"mock_resource", "mock_resource"}; - std::vector texts = {"first input", "second input"}; - auto col_resource = ColumnHelper::create_column(resources); - auto col_text = ColumnHelper::create_column(texts); + std::vector invalid_cases = {"1,0", "{}", "", " "}; - Block block; - block.insert({std::move(col_resource), std::make_shared(), "resource"}); - block.insert({std::move(col_text), std::make_shared(), "text"}); - block.insert({nullptr, std::make_shared(), "result"}); + for (const auto& invalid_value : invalid_cases) { + setenv("AI_TEST_RESULT", invalid_value.c_str(), 1); - ColumnNumbers arguments = {0, 1}; - size_t result_idx = 2; + std::vector resources = {"mock_resource"}; + std::vector texts = {"row1", "row2"}; + auto col_resource = ColumnHelper::create_column(resources); + auto col_text = ColumnHelper::create_column(texts); - auto sentiment_func = FunctionAISentiment::create(); - Status exec_status = - sentiment_func->execute_impl(ctx.get(), block, arguments, result_idx, texts.size()); + Block block; + block.insert({std::move(col_resource), std::make_shared(), "resource"}); + block.insert({std::move(col_text), std::make_shared(), "text"}); + block.insert({nullptr, std::make_shared(), "result"}); - unsetenv("AI_TEST_RESULT"); + ColumnNumbers arguments = {0, 1}; + size_t result_idx = 2; - ASSERT_TRUE(exec_status.ok()) << exec_status.to_string(); - const auto& res_col = - assert_cast(*block.get_by_position(result_idx).column); - ASSERT_EQ(res_col.size(), 2); - ASSERT_EQ(res_col.get_data_at(0).to_string(), "first result"); - ASSERT_EQ(res_col.get_data_at(1).to_string(), "second result"); -} + auto filter_func = FunctionAIFilter::create(); + Status exec_status = + filter_func->execute_impl(ctx.get(), block, arguments, result_idx, texts.size()); -TEST(AIFunctionTest, MockResourceBatchBoolResult) { - setenv("AI_TEST_RESULT", R"(["1","0"])", 1); + ASSERT_FALSE(exec_status.ok()) + << "Should have failed for invalid batch json: '" << invalid_value << "'"; + ASSERT_TRUE( + exec_status.to_string().find("Invalid batch result format") != std::string::npos || + exec_status.to_string().find("expected 2 items but got 1") != std::string::npos); + } + unsetenv("AI_TEST_RESULT"); +} + +TEST(AIFunctionTest, AIFilterBatchInvalidElement) { auto runtime_state = std::make_unique(); auto ctx = FunctionContext::create_context(runtime_state.get(), {}, {}); - std::vector resources = {"mock_resource", "mock_resource"}; - std::vector texts = {"valid input", "invalid input"}; + std::vector invalid_cases = {R"(["1","2"])", R"(["1",0])", R"(["yes","no"])"}; + + for (const auto& invalid_value : invalid_cases) { + setenv("AI_TEST_RESULT", invalid_value.c_str(), 1); + + std::vector resources = {"mock_resource"}; + std::vector texts = {"row1", "row2"}; + auto col_resource = ColumnHelper::create_column(resources); + auto col_text = ColumnHelper::create_column(texts); + + Block block; + block.insert({std::move(col_resource), std::make_shared(), "resource"}); + block.insert({std::move(col_text), std::make_shared(), "text"}); + block.insert({nullptr, std::make_shared(), "result"}); + + ColumnNumbers arguments = {0, 1}; + size_t result_idx = 2; + + auto filter_func = FunctionAIFilter::create(); + Status exec_status = + filter_func->execute_impl(ctx.get(), block, arguments, result_idx, texts.size()); + + ASSERT_FALSE(exec_status.ok()) + << "Should have failed for invalid batch element: '" << invalid_value << "'"; + ASSERT_TRUE(exec_status.to_string().find("Failed to parse boolean value") != + std::string::npos || + exec_status.to_string().find("Invalid batch result format") != + std::string::npos); + } + + unsetenv("AI_TEST_RESULT"); +} + +TEST(AIFunctionTest, AIFilterBatchSplitByWindow) { + TQueryOptions query_options = create_fake_query_options(); + query_options.__set_ai_context_window_size(128 * 1024); + auto query_ctx = MockQueryContext::create(TUniqueId(), ExecEnv::GetInstance(), query_options); + query_ctx->set_mock_ai_resource(); + TQueryGlobals query_globals; + auto runtime_state = std::make_unique( + TUniqueId(), 0, query_options, query_globals, nullptr, query_ctx.get()); + auto ctx = FunctionContext::create_context(runtime_state.get(), {}, {}); + + setenv("AI_TEST_RESULT", R"(["1"])", 1); + + std::vector resources = {"mock_resource"}; + std::vector texts = {std::string(70 * 1024, 'a'), std::string(70 * 1024, 'b'), + std::string(70 * 1024, 'c')}; auto col_resource = ColumnHelper::create_column(resources); auto col_text = ColumnHelper::create_column(texts); @@ -591,50 +974,488 @@ TEST(AIFunctionTest, MockResourceBatchBoolResult) { Status exec_status = filter_func->execute_impl(ctx.get(), block, arguments, result_idx, texts.size()); - unsetenv("AI_TEST_RESULT"); + ASSERT_TRUE(exec_status.ok()); - ASSERT_TRUE(exec_status.ok()) << exec_status.to_string(); const auto& res_col = assert_cast(*block.get_by_position(result_idx).column); - ASSERT_EQ(res_col.size(), 2); - ASSERT_EQ(res_col.get_data()[0], 1); - ASSERT_EQ(res_col.get_data()[1], 0); + ASSERT_EQ(res_col.size(), 3); + EXPECT_EQ(res_col.get_data()[0], 1); + EXPECT_EQ(res_col.get_data()[1], 1); + EXPECT_EQ(res_col.get_data()[2], 1); + + unsetenv("AI_TEST_RESULT"); } -TEST(AIFunctionTest, MockResourceBatchFloatResult) { - setenv("AI_TEST_RESULT", R"(["0.5","1.25"])", 1); +TEST(AIFunctionTest, AIFilterSingleRowExceedsBatchWindow) { + TQueryOptions query_options = create_fake_query_options(); + query_options.__set_ai_context_window_size(128 * 1024); + auto query_ctx = MockQueryContext::create(TUniqueId(), ExecEnv::GetInstance(), query_options); + query_ctx->set_mock_ai_resource(); + TQueryGlobals query_globals; + auto runtime_state = std::make_unique( + TUniqueId(), 0, query_options, query_globals, nullptr, query_ctx.get()); + auto ctx = FunctionContext::create_context(runtime_state.get(), {}, {}); + + std::vector resources = {"mock_resource"}; + std::vector texts = {std::string(130 * 1024, 'x')}; + auto col_resource = ColumnHelper::create_column(resources); + auto col_text = ColumnHelper::create_column(texts); - auto runtime_state = std::make_unique(); + Block block; + block.insert({std::move(col_resource), std::make_shared(), "resource"}); + block.insert({std::move(col_text), std::make_shared(), "text"}); + block.insert({nullptr, std::make_shared(), "result"}); + + ColumnNumbers arguments = {0, 1}; + size_t result_idx = 2; + + auto filter_func = FunctionAIFilter::create(); + // Even if a single row exceeds the batch window, it should be sent as a standalone request. + setenv("AI_TEST_RESULT", "[\"1\"]", 1); + Status exec_status = + filter_func->execute_impl(ctx.get(), block, arguments, result_idx, texts.size()); + + ASSERT_TRUE(exec_status.ok()); + + const auto& res_col = + assert_cast(*block.get_by_position(result_idx).column); + ASSERT_EQ(res_col.size(), 1); + EXPECT_EQ(res_col.get_data()[0], 1); + + unsetenv("AI_TEST_RESULT"); +} + +TEST(AIFunctionTest, AIFilterOversizedRowFlushesHistoryBatchFirst) { + TQueryOptions query_options = create_fake_query_options(); + query_options.__set_ai_context_window_size(128 * 1024); + auto query_ctx = MockQueryContext::create(TUniqueId(), ExecEnv::GetInstance(), query_options); + query_ctx->set_mock_ai_resource(); + TQueryGlobals query_globals; + auto runtime_state = std::make_unique( + TUniqueId(), 0, query_options, query_globals, nullptr, query_ctx.get()); auto ctx = FunctionContext::create_context(runtime_state.get(), {}, {}); std::vector resources = {"mock_resource", "mock_resource"}; - std::vector text1 = {"first text", "second text"}; - std::vector text2 = {"first compare", "second compare"}; + std::vector texts = {"small row", std::string(130 * 1024, 'x')}; auto col_resource = ColumnHelper::create_column(resources); - auto col_text1 = ColumnHelper::create_column(text1); - auto col_text2 = ColumnHelper::create_column(text2); + auto col_text = ColumnHelper::create_column(texts); Block block; block.insert({std::move(col_resource), std::make_shared(), "resource"}); - block.insert({std::move(col_text1), std::make_shared(), "text1"}); - block.insert({std::move(col_text2), std::make_shared(), "text2"}); - block.insert({nullptr, std::make_shared(), "result"}); + block.insert({std::move(col_text), std::make_shared(), "text"}); + block.insert({nullptr, std::make_shared(), "result"}); - ColumnNumbers arguments = {0, 1, 2}; - size_t result_idx = 3; + ColumnNumbers arguments = {0, 1}; + size_t result_idx = 2; - auto similarity_func = FunctionAISimilarity::create(); + auto filter_func = FunctionAIFilter::create(); + setenv("AI_TEST_RESULT", R"(["1"])", 1); Status exec_status = - similarity_func->execute_impl(ctx.get(), block, arguments, result_idx, text1.size()); + filter_func->execute_impl(ctx.get(), block, arguments, result_idx, texts.size()); + + ASSERT_TRUE(exec_status.ok()) << exec_status.to_string(); + + const auto& res_col = + assert_cast(*block.get_by_position(result_idx).column); + ASSERT_EQ(res_col.size(), 2); + EXPECT_EQ(res_col.get_data()[0], 1); + EXPECT_EQ(res_col.get_data()[1], 1); unsetenv("AI_TEST_RESULT"); +} + +TEST(AIFunctionTest, AIFilterUsesAiContextWindowSizeSessionVariable) { + TQueryOptions query_options = create_fake_query_options(); + query_options.__set_ai_context_window_size(16); + auto query_ctx = MockQueryContext::create(TUniqueId(), ExecEnv::GetInstance(), query_options); + query_ctx->set_mock_ai_resource(); + TQueryGlobals query_globals; + auto runtime_state = std::make_unique( + TUniqueId(), 0, query_options, query_globals, nullptr, query_ctx.get()); + auto ctx = FunctionContext::create_context(runtime_state.get(), {}, {}); + + setenv("AI_TEST_RESULT", R"(["1"])", 1); + + std::vector resources = {"mock_resource"}; + std::vector texts = {"12345678901234567890", "abcdefghijabcdefghij"}; + auto col_resource = ColumnHelper::create_column(resources); + auto col_text = ColumnHelper::create_column(texts); + + Block block; + block.insert({std::move(col_resource), std::make_shared(), "resource"}); + block.insert({std::move(col_text), std::make_shared(), "text"}); + block.insert({nullptr, std::make_shared(), "result"}); + + ColumnNumbers arguments = {0, 1}; + size_t result_idx = 2; + + auto filter_func = FunctionAIFilter::create(); + Status exec_status = + filter_func->execute_impl(ctx.get(), block, arguments, result_idx, texts.size()); ASSERT_TRUE(exec_status.ok()) << exec_status.to_string(); + const auto& res_col = - assert_cast(*block.get_by_position(result_idx).column); + assert_cast(*block.get_by_position(result_idx).column); ASSERT_EQ(res_col.size(), 2); - ASSERT_FLOAT_EQ(res_col.get_data()[0], 0.5F); - ASSERT_FLOAT_EQ(res_col.get_data()[1], 1.25F); + EXPECT_EQ(res_col.get_data()[0], 1); + EXPECT_EQ(res_col.get_data()[1], 1); + + unsetenv("AI_TEST_RESULT"); +} + +TEST(AIFunctionTest, ResourceNotFound) { + auto runtime_state = std::make_unique(); + auto ctx = FunctionContext::create_context(runtime_state.get(), {}, {}); + + std::vector resources = {"not_exist_resource"}; + std::vector texts = {"test"}; + auto col_resource = ColumnHelper::create_column(resources); + auto col_text = ColumnHelper::create_column(texts); + + Block block; + block.insert({std::move(col_resource), std::make_shared(), "resource"}); + block.insert({std::move(col_text), std::make_shared(), "text"}); + block.insert({nullptr, std::make_shared(), "result"}); + + ColumnNumbers arguments = {0, 1}; + size_t result_idx = 2; + + ASSERT_DEATH( + { + auto sentiment_func = FunctionAISentiment::create(); + Status exec_status = sentiment_func->execute_impl(ctx.get(), block, arguments, + result_idx, texts.size()); + static_cast(exec_status); + }, + "it != ai_resources->end"); +} + +TEST(AIFunctionTest, MockResourceSendRequest) { + auto runtime_state = std::make_unique(); + auto ctx = FunctionContext::create_context(runtime_state.get(), {}, {}); + + std::vector resources = {"mock_resource"}; + std::vector texts = {"test input"}; + auto col_resource = ColumnHelper::create_column(resources); + auto col_text = ColumnHelper::create_column(texts); + + Block block; + block.insert({std::move(col_resource), std::make_shared(), "resource"}); + block.insert({std::move(col_text), std::make_shared(), "text"}); + block.insert({nullptr, std::make_shared(), "result"}); + + ColumnNumbers arguments = {0, 1}; + size_t result_idx = 2; + + auto sentiment_func = FunctionAISentiment::create(); + Status exec_status = + sentiment_func->execute_impl(ctx.get(), block, arguments, result_idx, texts.size()); + + ASSERT_TRUE(exec_status.ok()) << exec_status.to_string(); + const auto& res_col = + assert_cast(*block.get_by_position(result_idx).column); + StringRef ref = res_col.get_data_at(0); + std::string val(ref.data, ref.size); + ASSERT_EQ(val, "this is a mock response. test input"); +} + +TEST(AIFunctionTest, AIStringFunctionBatchExecuteTest) { + auto runtime_state = std::make_unique(); + auto ctx = FunctionContext::create_context(runtime_state.get(), {}, {}); + + setenv("AI_TEST_RESULT", R"(["positive","negative","neutral"])", 1); + + std::vector resources = {"mock_resource"}; + std::vector texts = {"great", "bad", "okay"}; + auto col_resource = ColumnHelper::create_column(resources); + auto col_text = ColumnHelper::create_column(texts); + + Block block; + block.insert({std::move(col_resource), std::make_shared(), "resource"}); + block.insert({std::move(col_text), std::make_shared(), "text"}); + block.insert({nullptr, std::make_shared(), "result"}); + + ColumnNumbers arguments = {0, 1}; + size_t result_idx = 2; + + auto sentiment_func = FunctionAISentiment::create(); + Status exec_status = + sentiment_func->execute_impl(ctx.get(), block, arguments, result_idx, texts.size()); + + ASSERT_TRUE(exec_status.ok()); + + const auto& res_col = + assert_cast(*block.get_by_position(result_idx).column); + ASSERT_EQ(res_col.size(), 3); + EXPECT_EQ(res_col.get_data_at(0).to_string(), "positive"); + EXPECT_EQ(res_col.get_data_at(1).to_string(), "negative"); + EXPECT_EQ(res_col.get_data_at(2).to_string(), "neutral"); + + unsetenv("AI_TEST_RESULT"); +} + +TEST(AIFunctionTest, NullableStringResultThroughPreparedFunction) { + auto runtime_state = std::make_unique(); + 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, NullableLabelElementsThroughPreparedFunction) { + auto runtime_state = std::make_unique(); + auto ctx = FunctionContext::create_context(runtime_state.get(), {}, {}); + setenv("AI_TEST_RESULT", R"(["positive"])", 1); + + auto labels = create_string_array_column({{"positive", "unused-null", "negative"}}, + std::vector {0, 1, 0}); + Block block; + block.insert({ColumnHelper::create_column({"mock_resource"}), + std::make_shared(), "resource"}); + block.insert({ColumnHelper::create_column({"good product"}), + std::make_shared(), "text"}); + block.insert({std::move(labels), + std::make_shared(std::make_shared()), "labels"}); + + auto return_type = 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, 1); + unsetenv("AI_TEST_RESULT"); + + ASSERT_TRUE(status.ok()) << status.to_string(); + const auto& result = assert_cast(*block.get_by_position(3).column); + ASSERT_EQ(result.size(), 1); + EXPECT_EQ(result.get_data_at(0).to_string(), "positive"); +} + +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) { @@ -658,12 +1479,14 @@ TEST(AIFunctionTest, MissingAIResourcesMetadataTest) { ColumnNumbers arguments = {0, 1}; size_t result_idx = 2; - auto sentiment_func = FunctionAISentiment::create(); - Status exec_status = - sentiment_func->execute_impl(ctx.get(), block, arguments, result_idx, texts.size()); - - ASSERT_FALSE(exec_status.ok()); - ASSERT_NE(exec_status.to_string().find("AI resources metadata missing"), std::string::npos); + ASSERT_DEATH( + { + auto sentiment_func = FunctionAISentiment::create(); + Status exec_status = sentiment_func->execute_impl(ctx.get(), block, arguments, + result_idx, texts.size()); + static_cast(exec_status); + }, + "ai_resources"); } TEST(AIFunctionTest, ReturnTypeTest) { @@ -729,6 +1552,11 @@ class FunctionAISentimentTestHelper : public FunctionAISentiment { using FunctionAISentiment::normalize_endpoint; }; +class FunctionEmbedTestHelper : public FunctionEmbed { +public: + using FunctionEmbed::normalize_endpoint; +}; + TEST(AIFunctionTest, NormalizeLegacyCompletionsEndpoint) { TAIResource resource; resource.endpoint = "https://api.openai.com/v1/completions"; @@ -747,4 +1575,220 @@ TEST(AIFunctionTest, NormalizeEndpointNoopForOtherPaths) { ASSERT_EQ(resource.endpoint, "https://localhost/v1/responses"); } +TEST(AIFunctionTest, NormalizeGeminiGenerateEndpointFromBaseVersion) { + TAIResource resource; + resource.provider_type = "gemini"; + resource.model_name = "gemini-pro"; + resource.endpoint = "https://generativelanguage.googleapis.com/v1beta"; + + FunctionAISentimentTestHelper::normalize_endpoint(resource); + ASSERT_EQ(resource.endpoint, + "https://generativelanguage.googleapis.com/v1beta/models/gemini-pro:generateContent"); +} + +TEST(AIFunctionTest, NormalizeGeminiEmbedEndpointFromBaseVersion) { + TAIResource resource; + resource.provider_type = "GEMINI"; + resource.model_name = "gemini-embedding-2-preview"; + resource.endpoint = "https://generativelanguage.googleapis.com/v1beta"; + + FunctionEmbedTestHelper::normalize_endpoint(resource); + ASSERT_EQ(resource.endpoint, + "https://generativelanguage.googleapis.com/v1beta/models/" + "gemini-embedding-2-preview:batchEmbedContents"); +} + +TEST(AIFunctionTest, NormalizeGeminiEndpointNoopForNonBasePath) { + TAIResource resource; + resource.provider_type = "gemini"; + resource.model_name = "gemini-pro"; + resource.endpoint = + "https://generativelanguage.googleapis.com/v1beta/models/gemini-pro:generateContent"; + + FunctionAISentimentTestHelper::normalize_endpoint(resource); + ASSERT_EQ(resource.endpoint, + "https://generativelanguage.googleapis.com/v1beta/models/gemini-pro:generateContent"); +} + +TEST(AIFunctionTest, NormalizeGeminiEmbedLegacySingleEndpointToBatchEndpoint) { + TAIResource resource; + resource.provider_type = "gemini"; + resource.model_name = "gemini-embedding-2-preview"; + resource.endpoint = + "https://generativelanguage.googleapis.com/v1beta/models/" + "gemini-embedding-2-preview:embedContent"; + + FunctionEmbedTestHelper::normalize_endpoint(resource); + ASSERT_EQ(resource.endpoint, + "https://generativelanguage.googleapis.com/v1beta/models/" + "gemini-embedding-2-preview:batchEmbedContents"); +} + +TEST(AIFunctionTest, ExecuteBatchRequestSuccess) { + TQueryOptions query_options = create_fake_query_options(); + query_options.__set_query_timeout(5); + auto query_ctx = MockQueryContext::create(TUniqueId(), ExecEnv::GetInstance(), query_options); + TQueryGlobals query_globals; + RuntimeState runtime_state(TUniqueId(), 0, query_options, query_globals, nullptr, + query_ctx.get()); + auto ctx = FunctionContext::create_context(&runtime_state, {}, {}); + + OneShotHttpServer server(200, R"({"choices":[{"message":{"content":"[\"1\",\"0\"]"}}]})"); + + TAIResource config; + config.endpoint = server.endpoint(); + config.provider_type = "OPENAI"; + config.model_name = "test-model"; + config.api_key = "secret"; + config.max_retries = 1; + + std::shared_ptr adapter = std::make_shared(); + adapter->init(config); + + FunctionAIFilterBatchTestHelper helper; + std::vector results; + Status st = helper.execute_batch_request({"first row", "second row"}, results, config, adapter, + ctx.get()); + + ASSERT_TRUE(st.ok()) << st.to_string(); + ASSERT_EQ(results.size(), 2); + EXPECT_EQ(results[0], "1"); + EXPECT_EQ(results[1], "0"); + + std::string request = server.join_and_get_request(); + ASSERT_NE(request.find("Authorization: Bearer secret"), std::string::npos); + ASSERT_NE( + request.find( + R"([{\"idx\":0,\"input\":\"first row\"},{\"idx\":1,\"input\":\"second row\"}])"), + std::string::npos); +} + +TEST(AIFunctionTest, ExecuteBatchRequestResultSizeMismatch) { + TQueryOptions query_options = create_fake_query_options(); + query_options.__set_query_timeout(5); + auto query_ctx = MockQueryContext::create(TUniqueId(), ExecEnv::GetInstance(), query_options); + TQueryGlobals query_globals; + RuntimeState runtime_state(TUniqueId(), 0, query_options, query_globals, nullptr, + query_ctx.get()); + auto ctx = FunctionContext::create_context(&runtime_state, {}, {}); + + OneShotHttpServer server(200, R"({"choices":[{"message":{"content":"[\"1\"]"}}]})"); + + TAIResource config; + config.endpoint = server.endpoint(); + config.provider_type = "OPENAI"; + config.model_name = "test-model"; + config.api_key = "secret"; + config.max_retries = 1; + + std::shared_ptr adapter = std::make_shared(); + adapter->init(config); + + FunctionAIFilterBatchTestHelper helper; + std::vector results; + Status st = helper.execute_batch_request({"first row", "second row"}, results, config, adapter, + ctx.get()); + + ASSERT_FALSE(st.ok()); + ASSERT_NE(st.to_string().find( + "Failed to parse ai_filter batch result, expected 2 items but got 1"), + std::string::npos); +} + +TEST(AIFunctionTest, DoSendRequestTransportError) { + TQueryOptions query_options = create_fake_query_options(); + query_options.__set_query_timeout(5); + auto query_ctx = MockQueryContext::create(TUniqueId(), ExecEnv::GetInstance(), query_options); + TQueryGlobals query_globals; + RuntimeState runtime_state(TUniqueId(), 0, query_options, query_globals, nullptr, + query_ctx.get()); + auto ctx = FunctionContext::create_context(&runtime_state, {}, {}); + + TAIResource config; + config.endpoint = "http://127.0.0.1:1"; + config.provider_type = "OPENAI"; + config.model_name = "test-model"; + config.api_key = "secret"; + + std::shared_ptr adapter = std::make_shared(); + adapter->init(config); + + HttpClient client; + std::string response; + FunctionAITransportTestHelper helper; + Status st = helper.do_send_request(&client, "{}", response, config, adapter, ctx.get()); + + ASSERT_FALSE(st.ok()); + ASSERT_EQ(response, ""); +} + +TEST(AIFunctionTest, DoSendRequestNon200) { + TQueryOptions query_options = create_fake_query_options(); + query_options.__set_query_timeout(5); + auto query_ctx = MockQueryContext::create(TUniqueId(), ExecEnv::GetInstance(), query_options); + TQueryGlobals query_globals; + RuntimeState runtime_state(TUniqueId(), 0, query_options, query_globals, nullptr, + query_ctx.get()); + auto ctx = FunctionContext::create_context(&runtime_state, {}, {}); + + OneShotHttpServer server(500, R"({"error":"bad request"})"); + + TAIResource config; + config.endpoint = server.endpoint(); + config.provider_type = "OPENAI"; + config.model_name = "test-model"; + config.api_key = "secret"; + + std::shared_ptr adapter = std::make_shared(); + adapter->init(config); + + HttpClient client; + std::string response; + FunctionAITransportTestHelper helper; + Status st = helper.do_send_request(&client, R"({"message":"hello"})", response, config, adapter, + ctx.get()); + + ASSERT_FALSE(st.ok()); + ASSERT_NE(st.to_string().find("http status code is not 200"), std::string::npos); + ASSERT_EQ(response, R"({"error":"bad request"})"); + + std::string request = server.join_and_get_request(); + ASSERT_NE(request.find("Authorization: Bearer secret"), std::string::npos); + ASSERT_NE(request.find("Content-Type: application/json"), std::string::npos); +} + +TEST(AIFunctionTest, DoSendRequestSuccess) { + TQueryOptions query_options = create_fake_query_options(); + query_options.__set_query_timeout(5); + auto query_ctx = MockQueryContext::create(TUniqueId(), ExecEnv::GetInstance(), query_options); + TQueryGlobals query_globals; + RuntimeState runtime_state(TUniqueId(), 0, query_options, query_globals, nullptr, + query_ctx.get()); + auto ctx = FunctionContext::create_context(&runtime_state, {}, {}); + + OneShotHttpServer server(200, R"({"ok":true})"); + + TAIResource config; + config.endpoint = server.endpoint(); + config.provider_type = "OPENAI"; + config.model_name = "test-model"; + config.api_key = "secret"; + + std::shared_ptr adapter = std::make_shared(); + adapter->init(config); + + HttpClient client; + std::string response; + FunctionAITransportTestHelper helper; + Status st = helper.do_send_request(&client, R"({"message":"hello"})", response, config, adapter, + ctx.get()); + + ASSERT_TRUE(st.ok()) << st.to_string(); + ASSERT_EQ(response, R"({"ok":true})"); + + std::string request = server.join_and_get_request(); + ASSERT_NE(request.find("POST "), std::string::npos); + ASSERT_NE(request.find(R"({"message":"hello"})"), std::string::npos); +} + } // namespace doris diff --git a/be/test/ai/embed_test.cpp b/be/test/ai/embed_test.cpp index 2c26fccb6addfe..c9bd32ed17cb7a 100644 --- a/be/test/ai/embed_test.cpp +++ b/be/test/ai/embed_test.cpp @@ -18,6 +18,7 @@ #include "exprs/function/ai/embed.h" #include +#include #include #include #include @@ -25,7 +26,13 @@ #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" @@ -40,6 +47,181 @@ class MockHttpClient : public HttpClient { std::string _content_type; }; +class MockEmbedObjStorageClient : public io::ObjStorageClient { +public: + io::ObjectStorageUploadResponse create_multipart_upload( + const io::ObjectStoragePathOptions& /*opts*/) override { + return {}; + } + + io::ObjectStorageResponse put_object(const io::ObjectStoragePathOptions& /*opts*/, + std::string_view /*stream*/) override { + return io::ObjectStorageResponse::OK(); + } + + io::ObjectStorageUploadResponse upload_part(const io::ObjectStoragePathOptions& /*opts*/, + std::string_view /*stream*/, + int /*part_num*/) override { + return {}; + } + + io::ObjectStorageResponse complete_multipart_upload( + const io::ObjectStoragePathOptions& /*opts*/, + const std::vector& /*completed_parts*/) override { + return io::ObjectStorageResponse::OK(); + } + + io::ObjectStorageHeadResponse head_object( + const io::ObjectStoragePathOptions& /*opts*/) override { + return {}; + } + + io::ObjectStorageResponse get_object(const io::ObjectStoragePathOptions& /*opts*/, + void* /*buffer*/, size_t /*offset*/, size_t /*bytes_read*/, + size_t* /*size_return*/) override { + return io::ObjectStorageResponse::OK(); + } + + io::ObjectStorageResponse list_objects(const io::ObjectStoragePathOptions& /*opts*/, + std::vector* /*files*/) override { + return io::ObjectStorageResponse::OK(); + } + + io::ObjectStorageResponse delete_objects(const io::ObjectStoragePathOptions& /*opts*/, + std::vector /*objs*/) override { + return io::ObjectStorageResponse::OK(); + } + + io::ObjectStorageResponse delete_object(const io::ObjectStoragePathOptions& /*opts*/) override { + return io::ObjectStorageResponse::OK(); + } + + io::ObjectStorageResponse delete_objects_recursively( + const io::ObjectStoragePathOptions& /*opts*/) override { + return io::ObjectStorageResponse::OK(); + } + + std::string generate_presigned_url(const io::ObjectStoragePathOptions& opts, + int64_t expiration_secs, const S3ClientConf& conf) override { + last_opts = opts; + last_expiration_secs = expiration_secs; + last_conf = conf; + return fmt::format("mock-s3://{}/{}?ttl={}", opts.bucket, opts.key, expiration_secs); + } + + io::ObjectStoragePathOptions last_opts; + int64_t last_expiration_secs = 0; + S3ClientConf last_conf; +}; + +class CountingMultimodalMockAdapter : public MockAdapter { +public: + Status build_multimodal_embedding_request(const std::vector& media_types, + const std::vector& media_urls, + const std::vector& media_content_types, + std::string& request_body) const override { + EXPECT_EQ(media_types.size(), media_urls.size()); + EXPECT_EQ(media_content_types.size(), media_urls.size()); + batch_sizes.push_back(media_urls.size()); + request_body = "{}"; + return Status::OK(); + } + + mutable std::vector batch_sizes; +}; + +class CountingTextMockAdapter : public MockAdapter { +public: + Status build_embedding_request(const std::vector& inputs, + std::string& request_body) const override { + batch_sizes.push_back(inputs.size()); + request_body = "{}"; + return Status::OK(); + } + + mutable std::vector batch_sizes; +}; + +static ColumnString::MutablePtr create_jsonb_column(const std::vector& json_rows) { + auto column = ColumnString::create(); + for (const auto& json_row : json_rows) { + JsonBinaryValue jsonb_value; + Status st = jsonb_value.from_json_string(json_row); + EXPECT_TRUE(st.ok()) << st.to_string(); + column->insert_data(jsonb_value.value(), jsonb_value.size()); + } + return column; +} + +static ColumnPtr create_nullable_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); + + const auto& nested_nullable_col = assert_cast(col_array.get_data()); + const auto& nested_col = + assert_cast(*nested_nullable_col.get_nested_column_ptr()); + ASSERT_EQ(nested_col.size(), row_count * 5); + + for (size_t row = 0; row < row_count; ++row) { + ASSERT_EQ(offsets[row], (row + 1) * 5); + for (size_t i = 0; i < 5; ++i) { + ASSERT_FLOAT_EQ(nested_col.get_element(row * 5 + i), static_cast(i)); + } + } +} + +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; @@ -55,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"); @@ -73,7 +255,10 @@ TEST(EMBED_TEST, embed_function_test) { Block block; block.insert({std::move(col_resource), std::make_shared(), "resource"}); block.insert({std::move(col_text), std::make_shared(), "text"}); - block.insert({nullptr, std::make_shared(), "result"}); + block.insert( + {nullptr, + std::make_shared(make_nullable(std::make_shared())), + "result"}); ColumnNumbers arguments = {0, 1}; size_t result_idx = 2; @@ -96,19 +281,22 @@ TEST(EMBED_TEST, embed_function_test) { } } -TEST(EMBED_TEST, embed_function_batch_test) { +TEST(EMBED_TEST, embed_function_text_multi_rows) { auto runtime_state = std::make_unique(); auto ctx = FunctionContext::create_context(runtime_state.get(), {}, {}); std::vector resources = {"mock_resource", "mock_resource"}; - std::vector texts = {"first input", "second input"}; + std::vector texts = {"test input 1", "test input 2"}; auto col_resource = ColumnHelper::create_column(resources); auto col_text = ColumnHelper::create_column(texts); Block block; block.insert({std::move(col_resource), std::make_shared(), "resource"}); block.insert({std::move(col_text), std::make_shared(), "text"}); - block.insert({nullptr, std::make_shared(), "result"}); + block.insert( + {nullptr, + std::make_shared(make_nullable(std::make_shared())), + "result"}); ColumnNumbers arguments = {0, 1}; size_t result_idx = 2; @@ -120,19 +308,544 @@ TEST(EMBED_TEST, embed_function_batch_test) { ASSERT_TRUE(exec_status.ok()) << exec_status.to_string(); const auto& col_array = assert_cast(*block.get_by_position(result_idx).column); - const auto& offsets = col_array.get_offsets(); - ASSERT_EQ(offsets.size(), 2U); - ASSERT_EQ(offsets[0], 5); - ASSERT_EQ(offsets[1], 10); - const auto& nested_nullable_col = assert_cast(col_array.get_data()); - const auto& nested_col = - assert_cast(*nested_nullable_col.get_nested_column_ptr()); - ASSERT_EQ(nested_col.size(), 10U); - for (int row = 0; row < 2; ++row) { - for (int i = 0; i < 5; ++i) { - ASSERT_FLOAT_EQ(nested_col.get_element(row * 5 + i), static_cast(i)); - } - } + assert_mock_embedding_column(col_array, texts.size()); +} + +TEST(EMBED_TEST, embed_function_multimodal_direct_url) { + TQueryOptions query_options = create_fake_query_options(); + query_options.__set_file_presigned_url_ttl_seconds(0); + 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 resources = {"mock_resource", "mock_resource", "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"})", + R"({"content_type":"audio/mpeg","uri":"https://example.com/c.mp3"})"}; + + auto col_resource = ColumnHelper::create_column(resources); + auto col_file = create_jsonb_column(file_json_rows); + + Block block; + block.insert({std::move(col_resource), std::make_shared(), "resource"}); + block.insert({std::move(col_file), std::make_shared(), "file"}); + block.insert( + {nullptr, + std::make_shared(make_nullable(std::make_shared())), + "result"}); + + ColumnNumbers arguments = {0, 1}; + size_t result_idx = 2; + + auto embed_func = FunctionEmbed::create(); + Status exec_status = embed_func->execute_impl(ctx.get(), block, arguments, result_idx, + file_json_rows.size()); + + ASSERT_TRUE(exec_status.ok()) << exec_status.to_string(); + const auto& col_array = + assert_cast(*block.get_by_position(result_idx).column); + 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(), {}, {}); + + std::vector resources = {"mock_resource", "mock_resource", "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"})", + R"({"content_type":"audio/mpeg","uri":"https://example.com/c.mp3"})"}; + + auto col_resource = ColumnHelper::create_column(resources); + auto col_file = create_jsonb_column(file_json_rows); + + Block block; + block.insert({std::move(col_resource), std::make_shared(), "resource"}); + block.insert({std::move(col_file), std::make_shared(), "file"}); + block.insert( + {nullptr, + std::make_shared(make_nullable(std::make_shared())), + "result"}); + + TAIResource config; + config.provider_type = "MOCK"; + auto counting_adapter = std::make_shared(); + std::shared_ptr adapter = counting_adapter; + adapter->init(config); + + ColumnNumbers arguments = {0, 1}; + size_t result_idx = 2; + FunctionEmbed embed_func; + 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)); + + const auto& col_array = + assert_cast(*block.get_by_position(result_idx).column); + assert_mock_embedding_column(col_array, file_json_rows.size()); +} + +TEST(EMBED_TEST, embed_function_multimodal_batch_split_by_session_variable) { + 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); + 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 resources = {"mock_resource", "mock_resource", "mock_resource"}; + std::vector file_json_rows = { + R"({"content_type":"image/png","uri":"https://example.com/a.png"})", + R"({"content_type":"image/png","uri":"https://example.com/b.png"})", + R"({"content_type":"image/png","uri":"https://example.com/c.png"})"}; + + auto col_resource = ColumnHelper::create_column(resources); + auto col_file = create_jsonb_column(file_json_rows); + + Block block; + block.insert({std::move(col_resource), std::make_shared(), "resource"}); + block.insert({std::move(col_file), std::make_shared(), "file"}); + block.insert( + {nullptr, + std::make_shared(make_nullable(std::make_shared())), + "result"}); + + TAIResource config; + config.provider_type = "MOCK"; + auto counting_adapter = std::make_shared(); + std::shared_ptr adapter = counting_adapter; + adapter->init(config); + + ColumnNumbers arguments = {0, 1}; + size_t result_idx = 2; + FunctionEmbed embed_func; + 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)); + + const auto& col_array = + assert_cast(*block.get_by_position(result_idx).column); + assert_mock_embedding_column(col_array, file_json_rows.size()); +} + +TEST(EMBED_TEST, embed_function_text_batch_split_by_session_variable) { + 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); + 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 resources = {"mock_resource", "mock_resource", "mock_resource"}; + std::vector texts = {"text-a", "text-b", "text-c"}; + + auto col_resource = ColumnHelper::create_column(resources); + auto col_text = ColumnHelper::create_column(texts); + + Block block; + block.insert({std::move(col_resource), std::make_shared(), "resource"}); + block.insert({std::move(col_text), std::make_shared(), "text"}); + block.insert( + {nullptr, + std::make_shared(make_nullable(std::make_shared())), + "result"}); + + TAIResource config; + config.provider_type = "MOCK"; + auto counting_adapter = std::make_shared(); + std::shared_ptr adapter = counting_adapter; + adapter->init(config); + + ColumnNumbers arguments = {0, 1}; + size_t result_idx = 2; + FunctionEmbed embed_func; + 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)); + + const auto& col_array = + assert_cast(*block.get_by_position(result_idx).column); + assert_mock_embedding_column(col_array, texts.size()); +} + +TEST(EMBED_TEST, embed_function_multimodal_s3_presigned_url) { + TQueryOptions query_options = create_fake_query_options(); + query_options.__set_file_presigned_url_ttl_seconds(123); + 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, {}, {}); + + auto mock_client = std::make_shared(); + S3ClientFactory::instance().set_client_creator_for_test( + [mock_client](const S3ClientConf&) { return mock_client; }); + + std::vector resources = {"mock_resource"}; + std::vector file_json_rows = {R"({ + "content_type":"image/png", + "uri":"s3://test-bucket/path/to/image.png", + "endpoint":"cos.ap-beijing.myqcloud.com", + "region":"ap-beijing", + "ak":"test-ak", + "sk":"test-sk", + "role_arn":"test-role", + "external_id":"test-external-id" + })"}; + + auto col_resource = ColumnHelper::create_column(resources); + auto col_file = create_jsonb_column(file_json_rows); + + Block block; + block.insert({std::move(col_resource), std::make_shared(), "resource"}); + block.insert({std::move(col_file), std::make_shared(), "file"}); + block.insert( + {nullptr, + std::make_shared(make_nullable(std::make_shared())), + "result"}); + + ColumnNumbers arguments = {0, 1}; + size_t result_idx = 2; + + auto embed_func = FunctionEmbed::create(); + Status exec_status = embed_func->execute_impl(ctx.get(), block, arguments, result_idx, + file_json_rows.size()); + + S3ClientFactory::instance().clear_client_creator_for_test(); + + ASSERT_TRUE(exec_status.ok()) << exec_status.to_string(); + const auto& col_array = + assert_cast(*block.get_by_position(result_idx).column); + assert_mock_embedding_column(col_array, file_json_rows.size()); + + ASSERT_EQ(mock_client->last_opts.bucket, "test-bucket"); + ASSERT_EQ(mock_client->last_opts.key, "path/to/image.png"); + ASSERT_EQ(mock_client->last_expiration_secs, 123); + ASSERT_EQ(mock_client->last_conf.endpoint, "cos.ap-beijing.myqcloud.com"); + ASSERT_EQ(mock_client->last_conf.region, "ap-beijing"); + ASSERT_EQ(mock_client->last_conf.ak, "test-ak"); + ASSERT_EQ(mock_client->last_conf.sk, "test-sk"); + ASSERT_EQ(mock_client->last_conf.role_arn, "test-role"); + ASSERT_EQ(mock_client->last_conf.external_id, "test-external-id"); +} + +TEST(EMBED_TEST, embed_function_multimodal_s3_missing_endpoint) { + auto runtime_state = std::make_unique(); + auto ctx = FunctionContext::create_context(runtime_state.get(), {}, {}); + + std::vector resources = {"mock_resource"}; + std::vector file_json_rows = {R"({ + "content_type":"image/png", + "uri":"s3://test-bucket/path/to/image.png", + "region":"ap-beijing" + })"}; + auto col_resource = ColumnHelper::create_column(resources); + auto col_file = create_jsonb_column(file_json_rows); + + Block block; + block.insert({std::move(col_resource), std::make_shared(), "resource"}); + block.insert({std::move(col_file), std::make_shared(), "file"}); + block.insert( + {nullptr, + std::make_shared(make_nullable(std::make_shared())), + "result"}); + + ColumnNumbers arguments = {0, 1}; + size_t result_idx = 2; + + auto embed_func = FunctionEmbed::create(); + Status exec_status = embed_func->execute_impl(ctx.get(), block, arguments, result_idx, + file_json_rows.size()); + + ASSERT_FALSE(exec_status.ok()); + ASSERT_NE(exec_status.to_string().find("field 'endpoint' is required"), std::string::npos); +} + +TEST(EMBED_TEST, embed_function_multimodal_s3_missing_region) { + auto runtime_state = std::make_unique(); + auto ctx = FunctionContext::create_context(runtime_state.get(), {}, {}); + + std::vector resources = {"mock_resource"}; + std::vector file_json_rows = {R"({ + "content_type":"image/png", + "uri":"s3://test-bucket/path/to/image.png", + "endpoint":"cos.ap-beijing.myqcloud.com" + })"}; + auto col_resource = ColumnHelper::create_column(resources); + auto col_file = create_jsonb_column(file_json_rows); + + Block block; + block.insert({std::move(col_resource), std::make_shared(), "resource"}); + block.insert({std::move(col_file), std::make_shared(), "file"}); + block.insert( + {nullptr, + std::make_shared(make_nullable(std::make_shared())), + "result"}); + + ColumnNumbers arguments = {0, 1}; + size_t result_idx = 2; + + auto embed_func = FunctionEmbed::create(); + Status exec_status = embed_func->execute_impl(ctx.get(), block, arguments, result_idx, + file_json_rows.size()); + + ASSERT_FALSE(exec_status.ok()); + ASSERT_NE(exec_status.to_string().find("field 'region' is required"), std::string::npos); +} + +TEST(EMBED_TEST, embed_function_wrong_argument_count) { + auto runtime_state = std::make_unique(); + auto ctx = FunctionContext::create_context(runtime_state.get(), {}, {}); + + std::vector resources = {"mock_resource"}; + auto col_resource = ColumnHelper::create_column(resources); + + Block block; + block.insert({std::move(col_resource), std::make_shared(), "resource"}); + block.insert( + {nullptr, + std::make_shared(make_nullable(std::make_shared())), + "result"}); + + ColumnNumbers arguments = {0}; + size_t result_idx = 1; + + auto embed_func = FunctionEmbed::create(); + Status exec_status = + embed_func->execute_impl(ctx.get(), block, arguments, result_idx, resources.size()); + + ASSERT_FALSE(exec_status.ok()); + ASSERT_NE(exec_status.to_string().find("Function EMBED expects 2 arguments"), + std::string::npos); +} + +TEST(EMBED_TEST, embed_function_invalid_input_type) { + auto runtime_state = std::make_unique(); + auto ctx = FunctionContext::create_context(runtime_state.get(), {}, {}); + + std::vector resources = {"mock_resource"}; + std::vector ids = {1}; + auto col_resource = ColumnHelper::create_column(resources); + auto col_ids = ColumnHelper::create_column(ids); + + Block block; + block.insert({std::move(col_resource), std::make_shared(), "resource"}); + block.insert({std::move(col_ids), std::make_shared(), "id"}); + block.insert( + {nullptr, + std::make_shared(make_nullable(std::make_shared())), + "result"}); + + ColumnNumbers arguments = {0, 1}; + size_t result_idx = 2; + + auto embed_func = FunctionEmbed::create(); + Status exec_status = + embed_func->execute_impl(ctx.get(), block, arguments, result_idx, resources.size()); + + ASSERT_FALSE(exec_status.ok()); + ASSERT_NE(exec_status.to_string().find( + "Function EMBED expects the second argument to be STRING or JSON"), + std::string::npos); +} + +TEST(EMBED_TEST, embed_function_missing_required_json_field) { + auto runtime_state = std::make_unique(); + auto ctx = FunctionContext::create_context(runtime_state.get(), {}, {}); + + std::vector resources = {"mock_resource"}; + std::vector file_json_rows = {R"({"uri":"https://example.com/a.png"})"}; + auto col_resource = ColumnHelper::create_column(resources); + auto col_file = create_jsonb_column(file_json_rows); + + Block block; + block.insert({std::move(col_resource), std::make_shared(), "resource"}); + block.insert({std::move(col_file), std::make_shared(), "file"}); + block.insert( + {nullptr, + std::make_shared(make_nullable(std::make_shared())), + "result"}); + + ColumnNumbers arguments = {0, 1}; + size_t result_idx = 2; + + auto embed_func = FunctionEmbed::create(); + Status exec_status = embed_func->execute_impl(ctx.get(), block, arguments, result_idx, + file_json_rows.size()); + + ASSERT_FALSE(exec_status.ok()); + ASSERT_NE(exec_status.to_string().find("field 'content_type' is required"), std::string::npos); +} + +TEST(EMBED_TEST, embed_function_unsupported_content_type) { + auto runtime_state = std::make_unique(); + auto ctx = FunctionContext::create_context(runtime_state.get(), {}, {}); + + std::vector resources = {"mock_resource"}; + std::vector file_json_rows = { + R"({"content_type":"text/plain","uri":"https://example.com/a.txt"})"}; + auto col_resource = ColumnHelper::create_column(resources); + auto col_file = create_jsonb_column(file_json_rows); + + Block block; + block.insert({std::move(col_resource), std::make_shared(), "resource"}); + block.insert({std::move(col_file), std::make_shared(), "file"}); + block.insert( + {nullptr, + std::make_shared(make_nullable(std::make_shared())), + "result"}); + + ColumnNumbers arguments = {0, 1}; + size_t result_idx = 2; + + auto embed_func = FunctionEmbed::create(); + Status exec_status = embed_func->execute_impl(ctx.get(), block, arguments, result_idx, + file_json_rows.size()); + + ASSERT_FALSE(exec_status.ok()); + ASSERT_NE(exec_status.to_string().find("Unsupported content_type for EMBED"), + std::string::npos); } TEST(EMBED_TEST, local_adapter_embedding_request) { @@ -420,7 +1133,7 @@ TEST(EMBED_TEST, gemini_adapter_embedding_request) { EXPECT_STREQ(mock_client.get()->data, "x-goog-api-key: test_gemini_key"); EXPECT_STREQ(mock_client.get()->next->data, "Content-Type: application/json"); - std::vector inputs = {"embed with gemini"}; + std::vector inputs = {"embed with gemini", "embed batch with gemini"}; std::string request_body; Status st = adapter.build_embedding_request(inputs, request_body); ASSERT_TRUE(st.ok()); @@ -433,21 +1146,30 @@ TEST(EMBED_TEST, gemini_adapter_embedding_request) { ASSERT_TRUE(doc.HasMember("requests")) << "Missing requests field"; ASSERT_TRUE(doc["requests"].IsArray()) << request_body; - ASSERT_EQ(doc["requests"].Size(), 1); - const auto& request = doc["requests"][0]; - ASSERT_TRUE(request.HasMember("model")) << "Missing request model field"; - ASSERT_STREQ(request["model"].GetString(), "models/embedding-001"); - ASSERT_TRUE(request.HasMember("content")) << "Missing request content field"; - ASSERT_TRUE(request["content"].IsObject()) << request_body; - - auto& content = request["content"]; - ASSERT_TRUE(content.HasMember("parts")) << request_body; - ASSERT_TRUE(content["parts"].IsArray()); - ASSERT_TRUE(content["parts"][0].HasMember("text")) << request_body; - ASSERT_STREQ(content["parts"][0]["text"].GetString(), "embed with gemini"); - - // should not have dimension param; - ASSERT_FALSE(request.HasMember("outputDimensionality")); + ASSERT_EQ(doc["requests"].Size(), 2); + + const auto& request0 = doc["requests"][0]; + ASSERT_TRUE(request0.HasMember("model")) << request_body; + ASSERT_STREQ(request0["model"].GetString(), "models/embedding-001"); + ASSERT_TRUE(request0.HasMember("content")) << request_body; + ASSERT_TRUE(request0["content"].IsObject()) << request_body; + ASSERT_TRUE(request0["content"].HasMember("parts")) << request_body; + ASSERT_TRUE(request0["content"]["parts"].IsArray()) << request_body; + ASSERT_EQ(request0["content"]["parts"].Size(), 1); + ASSERT_TRUE(request0["content"]["parts"][0].HasMember("text")) << request_body; + ASSERT_STREQ(request0["content"]["parts"][0]["text"].GetString(), "embed with gemini"); + ASSERT_FALSE(request0.HasMember("outputDimensionality")); + + const auto& request1 = doc["requests"][1]; + ASSERT_TRUE(request1.HasMember("model")) << request_body; + ASSERT_STREQ(request1["model"].GetString(), "models/embedding-001"); + ASSERT_TRUE(request1.HasMember("content")) << request_body; + ASSERT_TRUE(request1["content"].IsObject()) << request_body; + ASSERT_TRUE(request1["content"].HasMember("parts")) << request_body; + ASSERT_TRUE(request1["content"]["parts"].IsArray()) << request_body; + ASSERT_EQ(request1["content"]["parts"].Size(), 1); + ASSERT_TRUE(request1["content"]["parts"][0].HasMember("text")) << request_body; + ASSERT_STREQ(request1["content"]["parts"][0]["text"].GetString(), "embed batch with gemini"); config.model_name = "gemini-embedding-001"; adapter.init(config); @@ -458,9 +1180,11 @@ TEST(EMBED_TEST, gemini_adapter_embedding_request) { ASSERT_TRUE(doc.IsObject()) << "JSON is not an object"; ASSERT_TRUE(doc.HasMember("requests")) << request_body; ASSERT_TRUE(doc["requests"].IsArray()) << request_body; - ASSERT_EQ(doc["requests"].Size(), 1); + ASSERT_EQ(doc["requests"].Size(), 2); ASSERT_TRUE(doc["requests"][0].HasMember("outputDimensionality")) << request_body; ASSERT_EQ(doc["requests"][0]["outputDimensionality"].GetInt(), 768) << request_body; + ASSERT_TRUE(doc["requests"][1].HasMember("outputDimensionality")) << request_body; + ASSERT_EQ(doc["requests"][1]["outputDimensionality"].GetInt(), 768) << request_body; } TEST(EMBED_TEST, gemini_adapter_parse_embedding_response) { @@ -484,6 +1208,36 @@ TEST(EMBED_TEST, gemini_adapter_parse_embedding_response) { ASSERT_FLOAT_EQ(results[0][0], 0.1F); ASSERT_FLOAT_EQ(results[0][1], 0.2F); ASSERT_FLOAT_EQ(results[0][2], 0.3F); + + resp = R"({ + "embeddings": [ + { + "values":[ + 1.1, + 1.2 + ] + }, + { + "values":[ + 2.1, + 2.2, + 2.3 + ] + } + ] + })"; + + results.clear(); + st = adapter.parse_embedding_response(resp, results); + ASSERT_TRUE(st.ok()) << st.to_string(); + ASSERT_EQ(results.size(), 2); + ASSERT_EQ(results[0].size(), 2); + ASSERT_EQ(results[1].size(), 3); + ASSERT_FLOAT_EQ(results[0][0], 1.1F); + ASSERT_FLOAT_EQ(results[0][1], 1.2F); + ASSERT_FLOAT_EQ(results[1][0], 2.1F); + ASSERT_FLOAT_EQ(results[1][1], 2.2F); + ASSERT_FLOAT_EQ(results[1][2], 2.3F); } TEST(EMBED_TEST, voyageai_adapter_embedding_request) { @@ -655,6 +1409,9 @@ TEST(EMBED_TEST, deepseek_adapter_embedding_request) { std::string request_body; Status st = adapter.build_embedding_request(inputs, request_body); ASSERT_FALSE(st.ok()); + ASSERT_THAT(st.to_string(), + ::testing::HasSubstr("Currently supported providers are OpenAI, Gemini, " + "Voyage, Jina, Qwen, and Minimax")); } TEST(EMBED_TEST, deepseek_adapter_parse_embedding_response) { @@ -679,6 +1436,9 @@ TEST(EMBED_TEST, deepseek_adapter_parse_embedding_response) { std::vector> results; Status st = adapter.parse_embedding_response(resp, results); ASSERT_FALSE(st.ok()); + ASSERT_THAT(st.to_string(), + ::testing::HasSubstr("Currently supported providers are OpenAI, Gemini, " + "Voyage, Jina, Qwen, and Minimax")); } TEST(EMBED_TEST, moonshot_adapter_embedding_request) { @@ -702,6 +1462,9 @@ TEST(EMBED_TEST, moonshot_adapter_embedding_request) { std::string request_body; Status st = adapter.build_embedding_request(inputs, request_body); ASSERT_FALSE(st.ok()); + ASSERT_THAT(st.to_string(), + ::testing::HasSubstr("Currently supported providers are OpenAI, Gemini, " + "Voyage, Jina, Qwen, and Minimax")); } TEST(EMBED_TEST, moonshot_adapter_parse_embedding_response) { @@ -727,6 +1490,9 @@ TEST(EMBED_TEST, moonshot_adapter_parse_embedding_response) { std::vector> results; Status st = adapter.parse_embedding_response(resp, results); ASSERT_FALSE(st.ok()); + ASSERT_THAT(st.to_string(), + ::testing::HasSubstr("Currently supported providers are OpenAI, Gemini, " + "Voyage, Jina, Qwen, and Minimax")); } TEST(EMBED_TEST, minimax_adapter_embedding_request) { diff --git a/fe/fe-core/src/main/java/org/apache/doris/catalog/Resource.java b/fe/fe-core/src/main/java/org/apache/doris/catalog/Resource.java index 8250cbca56291d..87152506178103 100644 --- a/fe/fe-core/src/main/java/org/apache/doris/catalog/Resource.java +++ b/fe/fe-core/src/main/java/org/apache/doris/catalog/Resource.java @@ -28,6 +28,7 @@ import org.apache.doris.nereids.trees.plans.commands.info.CreateResourceInfo; import org.apache.doris.persist.gson.GsonPostProcessable; import org.apache.doris.persist.gson.GsonUtils; +import org.apache.doris.qe.ConnectContext; import com.google.common.base.Strings; import com.google.common.collect.ImmutableMap; @@ -359,4 +360,11 @@ private void notifyUpdate(Map properties) { } public void applyDefaultProperties() {} + + public static void registerUsedAIResourceName(String resourceName) { + ConnectContext ctx = ConnectContext.get(); + if (ctx != null && ctx.getStatementContext() != null) { + ctx.getStatementContext().registerUsedAIResourceName(resourceName); + } + } } diff --git a/fe/fe-core/src/main/java/org/apache/doris/datasource/property/constants/AIProperties.java b/fe/fe-core/src/main/java/org/apache/doris/datasource/property/constants/AIProperties.java index 8bddd0949546c0..6019aaef303268 100644 --- a/fe/fe-core/src/main/java/org/apache/doris/datasource/property/constants/AIProperties.java +++ b/fe/fe-core/src/main/java/org/apache/doris/datasource/property/constants/AIProperties.java @@ -55,7 +55,7 @@ public class AIProperties extends BaseProperties { public static final List REQUIRED_FIELDS = Arrays.asList(ENDPOINT, PROVIDER_TYPE, MODEL_NAME); public static final List PROVIDERS = Arrays.asList("OPENAI", "LOCAL", "GEMINI", "DEEPSEEK", "ANTHROPIC", - "MOONSHOT", "QWEN", "MINIMAX", "ZHIPU", "BAICHUAN", "VOYAGEAI"); + "MOONSHOT", "QWEN", "MINIMAX", "ZHIPU", "BAICHUAN", "VOYAGEAI", "JINA"); public static void requiredAIProperties(Map properties) throws DdlException { // Check required field diff --git a/fe/fe-core/src/main/java/org/apache/doris/nereids/StatementContext.java b/fe/fe-core/src/main/java/org/apache/doris/nereids/StatementContext.java index 2dbec3729bb6f4..ba0a6d23261784 100644 --- a/fe/fe-core/src/main/java/org/apache/doris/nereids/StatementContext.java +++ b/fe/fe-core/src/main/java/org/apache/doris/nereids/StatementContext.java @@ -70,6 +70,7 @@ import com.google.common.annotations.VisibleForTesting; import com.google.common.base.Stopwatch; +import com.google.common.base.Strings; import com.google.common.base.Supplier; import com.google.common.base.Suppliers; import com.google.common.base.Throwables; @@ -90,6 +91,7 @@ import java.util.HashMap; import java.util.HashSet; import java.util.LinkedHashMap; +import java.util.LinkedHashSet; import java.util.List; import java.util.Map; import java.util.Optional; @@ -325,9 +327,11 @@ public enum TableFrom { private boolean useGatherForIcebergRewrite = false; private boolean hasNestedColumns; + private final Set mustInlineCTE = new HashSet<>(); + private final Set usedAIResourceNames = new LinkedHashSet<>(); + private final Map lowerCaseTableNamesCache = Maps.newHashMap(); private final Map lowerCaseDatabaseNamesCache = Maps.newHashMap(); - private final Set mustInlineCTE = new HashSet<>(); public StatementContext() { this(ConnectContext.get(), null, 0); @@ -471,6 +475,17 @@ public ConnectContext getConnectContext() { return connectContext; } + public Set getUsedAIResourceNames() { + return Collections.unmodifiableSet(usedAIResourceNames); + } + + public void registerUsedAIResourceName(String resourceName) { + if (Strings.isNullOrEmpty(resourceName)) { + throw new AnalysisException("AI resource name can not be empty"); + } + usedAIResourceNames.add(resourceName); + } + /** * Register an external relation that may preload metadata before internal table locks are acquired. * diff --git a/fe/fe-core/src/main/java/org/apache/doris/nereids/trees/expressions/functions/agg/AIAgg.java b/fe/fe-core/src/main/java/org/apache/doris/nereids/trees/expressions/functions/agg/AIAgg.java index a609019ab00348..833f1b7db4d364 100644 --- a/fe/fe-core/src/main/java/org/apache/doris/nereids/trees/expressions/functions/agg/AIAgg.java +++ b/fe/fe-core/src/main/java/org/apache/doris/nereids/trees/expressions/functions/agg/AIAgg.java @@ -90,6 +90,7 @@ public void checkLegalityAfterRewrite() { if (!(resource instanceof AIResource)) { throw new AnalysisException("AI resource '" + resourceName + "' does not exist"); } + Resource.registerUsedAIResourceName(resourceName); } } diff --git a/fe/fe-core/src/main/java/org/apache/doris/nereids/trees/expressions/functions/ai/AIFunction.java b/fe/fe-core/src/main/java/org/apache/doris/nereids/trees/expressions/functions/ai/AIFunction.java index 2826991208c84e..7399f2a348f68f 100644 --- a/fe/fe-core/src/main/java/org/apache/doris/nereids/trees/expressions/functions/ai/AIFunction.java +++ b/fe/fe-core/src/main/java/org/apache/doris/nereids/trees/expressions/functions/ai/AIFunction.java @@ -61,6 +61,7 @@ public void checkLegalityAfterRewrite() { if (!(resource instanceof AIResource)) { throw new AnalysisException("AI resource '" + resourceName + "' does not exist"); } + Resource.registerUsedAIResourceName(resourceName); } } diff --git a/fe/fe-core/src/main/java/org/apache/doris/nereids/trees/expressions/functions/ai/Embed.java b/fe/fe-core/src/main/java/org/apache/doris/nereids/trees/expressions/functions/ai/Embed.java index 3ce12e72043f03..e3e8725856ff5d 100644 --- a/fe/fe-core/src/main/java/org/apache/doris/nereids/trees/expressions/functions/ai/Embed.java +++ b/fe/fe-core/src/main/java/org/apache/doris/nereids/trees/expressions/functions/ai/Embed.java @@ -17,12 +17,18 @@ package org.apache.doris.nereids.trees.expressions.functions.ai; +import org.apache.doris.catalog.AIResource; +import org.apache.doris.catalog.Env; import org.apache.doris.catalog.FunctionSignature; +import org.apache.doris.catalog.Resource; +import org.apache.doris.nereids.exceptions.AnalysisException; import org.apache.doris.nereids.trees.expressions.Expression; +import org.apache.doris.nereids.trees.expressions.literal.StringLikeLiteral; import org.apache.doris.nereids.trees.expressions.literal.StringLiteral; import org.apache.doris.nereids.trees.expressions.visitor.ExpressionVisitor; import org.apache.doris.nereids.types.ArrayType; import org.apache.doris.nereids.types.FloatType; +import org.apache.doris.nereids.types.JsonType; import org.apache.doris.nereids.types.StringType; import org.apache.doris.nereids.types.VarcharType; @@ -41,7 +47,11 @@ public class Embed extends AIFunction { FunctionSignature.ret(ArrayType.of(FloatType.INSTANCE)) .args(StringType.INSTANCE, StringType.INSTANCE), FunctionSignature.ret(ArrayType.of(FloatType.INSTANCE)) - .args(VarcharType.SYSTEM_DEFAULT, VarcharType.SYSTEM_DEFAULT) + .args(VarcharType.SYSTEM_DEFAULT, VarcharType.SYSTEM_DEFAULT), + FunctionSignature.ret(ArrayType.of(FloatType.INSTANCE)) + .args(StringType.INSTANCE, JsonType.INSTANCE), + FunctionSignature.ret(ArrayType.of(FloatType.INSTANCE)) + .args(VarcharType.SYSTEM_DEFAULT, JsonType.INSTANCE) ); /** @@ -60,7 +70,7 @@ public Embed(Expression arg0, Expression arg1) { @Override public Embed withChildren(List children) { - Preconditions.checkArgument(children.size() == 1 || children.size() == 2, + Preconditions.checkArgument(children.size() >= 1 && children.size() <= 2, "Function EMBED only accepts 1 or 2 arguments"); if (children.size() == 1) { return new Embed(new StringLiteral(getResourceName()), @@ -83,4 +93,38 @@ public List getSignatures() { public R accept(ExpressionVisitor visitor, C context) { return visitor.visitEmbed(this, context); } + + @Override + public void checkLegalityBeforeTypeCoercion() { + if (arity() == 1) { + return; + } + if (arity() == 2) { + String aiResourceName = requireStringLiteral(child(0), "resource name", + "AI Function must accept literal for the resource name."); + validateAIResource(aiResourceName); + return; + } + throw new AnalysisException("Function EMBED only accepts 1 or 2 arguments"); + } + + private static String requireStringLiteral(Expression arg, String argName, String errorMsg) { + if (!(arg instanceof StringLikeLiteral)) { + throw new AnalysisException(errorMsg); + } + String value = ((StringLikeLiteral) arg).getStringValue(); + if (value == null || value.isEmpty()) { + throw new AnalysisException("EMBED " + argName + " can not be empty."); + } + return value; + } + + private static void validateAIResource(String resourceName) { + Resource resource = Env.getCurrentEnv().getResourceMgr().getResource(resourceName); + if (!(resource instanceof AIResource)) { + throw new AnalysisException("AI resource '" + resourceName + "' does not exist"); + } + Resource.registerUsedAIResourceName(resourceName); + } + } diff --git a/fe/fe-core/src/main/java/org/apache/doris/qe/Coordinator.java b/fe/fe-core/src/main/java/org/apache/doris/qe/Coordinator.java index a08a8de50bfaa9..13e14e6ff3eb48 100644 --- a/fe/fe-core/src/main/java/org/apache/doris/qe/Coordinator.java +++ b/fe/fe-core/src/main/java/org/apache/doris/qe/Coordinator.java @@ -1783,6 +1783,21 @@ private TNetworkAddress toArrowFlightHost(TNetworkAddress host) throws Exception return backend.getArrowFlightAddress(); } + private Map getNeededAiResources() { + Map aiResourceMap = Maps.newLinkedHashMap(); + if (context == null || context.getStatementContext() == null) { + return aiResourceMap; + } + for (String resourceName : context.getStatementContext().getUsedAIResourceNames()) { + Resource resource = Env.getCurrentEnv().getResourceMgr().getResource(resourceName); + if (!(resource instanceof AIResource)) { + throw new IllegalStateException("AI resource '" + resourceName + "' does not exist"); + } + aiResourceMap.put(resourceName, ((AIResource) resource).toThrift()); + } + return aiResourceMap; + } + // estimate if this fragment contains UnionNode private boolean containsUnionNode(PlanNode node) { if (node instanceof UnionNode) { @@ -3415,15 +3430,7 @@ Map toThrift(int backendNum) { } // Used for AI Functions - Map aiResourceMap = Maps.newLinkedHashMap(); - for (Resource resource : Env.getCurrentEnv().getResourceMgr() - .getResource(Resource.ResourceType.AI)) { - if (resource instanceof AIResource) { - aiResourceMap.put(resource.getName(), ((AIResource) resource).toThrift()); - } - } - - params.setAiResources(aiResourceMap); + params.setAiResources(getNeededAiResources()); res.put(instanceExecParam.host, params); res.get(instanceExecParam.host).setBucketSeqToInstanceIdx(new HashMap()); res.get(instanceExecParam.host).setShuffleIdxToInstanceIdx(new HashMap()); @@ -3671,4 +3678,3 @@ public void setIsProfileSafeStmt(boolean isSafe) { this.queryOptions.setEnableProfile(isSafe && queryOptions.isEnableProfile()); } } - diff --git a/fe/fe-core/src/main/java/org/apache/doris/qe/SessionVariable.java b/fe/fe-core/src/main/java/org/apache/doris/qe/SessionVariable.java index 28b33167980ca8..aa3e37d92cab0c 100644 --- a/fe/fe-core/src/main/java/org/apache/doris/qe/SessionVariable.java +++ b/fe/fe-core/src/main/java/org/apache/doris/qe/SessionVariable.java @@ -990,6 +990,7 @@ public static double getHotValueThreshold() { public static final String ENABLE_STRICT_CAST = "enable_strict_cast"; public static final String DEFAULT_AI_RESOURCE = "default_ai_resource"; + public static final String FILE_PRESIGNED_URL_TTL_SECONDS = "file_presigned_url_ttl_seconds"; public static final String EMBED_MAX_BATCH_SIZE = "embed_max_batch_size"; public static final String AI_CONTEXT_WINDOW_SIZE = "ai_context_window_size"; public static final String HNSW_EF_SEARCH = "hnsw_ef_search"; @@ -3557,6 +3558,13 @@ public void setDetailShapePlanNodes(String detailShapePlanNodes) { }) public String defaultAIResource = ""; + @VariableMgr.VarAttr(name = FILE_PRESIGNED_URL_TTL_SECONDS, needForward = true, + description = { + "EMBED 多模态场景中,S3 预签名 URL 的有效期(秒)。", + "Expiration time in seconds for S3 presigned URL used by multimodal EMBED." + }) + public long filePresignedUrlTtlSeconds = 3600; + public void setEnableEsParallelScroll(boolean enableESParallelScroll) { this.enableESParallelScroll = enableESParallelScroll; } @@ -5643,6 +5651,7 @@ public TQueryOptions toThrift() { tResult.setEnableOrcFilterByMinMax(enableOrcFilterByMinMax); tResult.setEnableExprZonemapFilter(enableExprZonemapFilter); tResult.setEnablePaimonCppReader(enablePaimonCppReader); + tResult.setFilePresignedUrlTtlSeconds(filePresignedUrlTtlSeconds); tResult.setCheckOrcInitSargsSuccess(checkOrcInitSargsSuccess); tResult.setTruncateCharOrVarcharColumns(truncateCharOrVarcharColumns); diff --git a/fe/fe-core/src/main/java/org/apache/doris/qe/runtime/ThriftPlansBuilder.java b/fe/fe-core/src/main/java/org/apache/doris/qe/runtime/ThriftPlansBuilder.java index 2465d2c6659b46..5b3d8b24f6df4d 100644 --- a/fe/fe-core/src/main/java/org/apache/doris/qe/runtime/ThriftPlansBuilder.java +++ b/fe/fe-core/src/main/java/org/apache/doris/qe/runtime/ThriftPlansBuilder.java @@ -23,6 +23,7 @@ import org.apache.doris.catalog.Resource; import org.apache.doris.common.Config; import org.apache.doris.datasource.FileQueryScanNode; +import org.apache.doris.nereids.StatementContext; import org.apache.doris.nereids.trees.plans.distribute.DistributedPlan; import org.apache.doris.nereids.trees.plans.distribute.PipelineDistributedPlan; import org.apache.doris.nereids.trees.plans.distribute.worker.DistributedPlanWorker; @@ -217,6 +218,27 @@ private static Supplier> topNFilterToThrift(List collectAiResources(ConnectContext connectContext) { + Map aiResourceMap = Maps.newLinkedHashMap(); + if (connectContext == null) { + return aiResourceMap; + } + + StatementContext statementContext = connectContext.getStatementContext(); + if (statementContext == null) { + return aiResourceMap; + } + + for (String resourceName : statementContext.getUsedAIResourceNames()) { + Resource resource = Env.getCurrentEnv().getResourceMgr().getResource(resourceName); + if (!(resource instanceof AIResource)) { + throw new IllegalStateException("AI resource '" + resourceName + "' does not exist"); + } + aiResourceMap.put(resourceName, ((AIResource) resource).toThrift()); + } + return aiResourceMap; + } + private static void setParamsForOlapTableSink(List distributedPlans, Map fragmentsGroupByWorker, CoordinatorContext coordinatorContext) { @@ -497,13 +519,7 @@ private static TPipelineFragmentParams fragmentToThriftIfAbsent( params.setShuffleIdxToInstanceIdx(computeDestIdToInstanceId(fragmentPlan, w, instanceToIndex)); // Only used for AI Functions - Map aiResourceMap = Maps.newLinkedHashMap(); - for (Resource resource : Env.getCurrentEnv().getResourceMgr().getResource(Resource.ResourceType.AI)) { - if (resource instanceof AIResource) { - aiResourceMap.put(resource.getName(), ((AIResource) resource).toThrift()); - } - } - params.setAiResources(aiResourceMap); + params.setAiResources(collectAiResources(connectContext)); return params; }); diff --git a/gensrc/thrift/PaloInternalService.thrift b/gensrc/thrift/PaloInternalService.thrift index 64a8331dd931fb..52a9ab406dbfd2 100644 --- a/gensrc/thrift/PaloInternalService.thrift +++ b/gensrc/thrift/PaloInternalService.thrift @@ -505,6 +505,9 @@ struct TQueryOptions { 225: optional i64 runtime_filter_tree_publish_max_send_bytes = 268435456 226: optional bool enable_local_exchange_before_streaming_agg = false + + 227: optional i64 file_presigned_url_ttl_seconds = 3600; + // For cloud, to control if the content would be written into file cache // In write path, to control if the content would be written into file cache. // In read path, read from file cache or remote storage when execute query. diff --git a/regression-test/suites/ai_p0/test_ai_functions.groovy b/regression-test/suites/ai_p0/test_ai_functions.groovy index 7ff9ff7777a257..57cc2040e14162 100644 --- a/regression-test/suites/ai_p0/test_ai_functions.groovy +++ b/regression-test/suites/ai_p0/test_ai_functions.groovy @@ -100,7 +100,10 @@ suite("test_ai_functions") { sql """${sql_text}""" assertTrue(false) } catch (Exception e) { - assertTrue(e.getMessage().contains("timeout") || e.getMessage().contains("requested URL returned error")) + assertTrue(e.getMessage().contains("timeout") + || e.getMessage().contains("requested URL returned error") + || e.getMessage().contains("http status code is not 200"), + "Unexpected exception message: " + e.getMessage()) } finally { sql """UNSET VARIABLE query_timeout;""" }