From 5c0df0ef687d71b52d6695ab37e966cae2d6104a Mon Sep 17 00:00:00 2001 From: linrrarity Date: Tue, 14 Apr 2026 17:22:26 +0800 Subject: [PATCH 1/2] [Feature](multimodal_embed) Support multimodal(file) embed (#62147) Support image/video/audio embedding from json. Supported multimodal_embed AI provider : Gemini(image, video, audio), Qwen(image, video), Voyage(image, video), Jina(image) Supported authentication methods include IAM + EXTERNAL_ID / AK + SK ```text -- ima role, video mysql> SELECT ARRAY_SIZE( -> EMBED( -> 'qwen_mul_embed', -> CAST('{ '> "uri": "s3://selectdb-qa-test-3/lzq-multimodal-test/video/45944deac7d96c872f559f2ef94ea0a9.mp4", '> "content_type": "video/mp4", '> "provider": "AWS", '> "endpoint": "s3.us-east-1.amazonaws.com", '> "region": "us-east-1", '> "role_arn": "arn:aws:iam::447051187841:role/lzq-role-test", '> "external_id": "1001" '> }' AS JSON) -> ) -> ) AS video_embed_size; +------------------+ | video_embed_size | +------------------+ | 2560 | +------------------+ -- ima role, video mysql> SELECT ARRAY_SIZE( -> EMBED( -> 'qwen_mul_embed', -> CAST('{ '> "uri": "s3://selectdb-qa-test-3/lzq-multimodal-test/image/ai-agg.png", '> "content_type": "image/png", '> "provider": "AWS", '> "endpoint": "s3.us-east-1.amazonaws.com", '> "region": "us-east-1", '> "role_arn": "arn:aws:iam::447051187841:role/lzq-role-test", '> "external_id": "1001" '> }' AS JSON) -> ) -> ) AS img_embed_size; +----------------+ | img_embed_size | +----------------+ | 2560 | +----------------+ -- gemini multimodal embed mysql> SELECT ARRAY_SIZE( -> EMBED( -> 'gemini_mul_embed', -> CAST('{ '> "uri": "s3://selectdb-qa-test-3/lzq-multimodal-test/image/ai-agg.png", '> "content_type": "image/png", '> "provider": "AWS", '> "endpoint": "s3.us-east-1.amazonaws.com", '> "region": "us-east-1", '> "role_arn": "arn:aws:iam::447051187841:role/lzq-role-test", '> "external_id": "1001" '> }' AS JSON) -> ) -> ) AS img_embed_size; +----------------+ | img_embed_size | +----------------+ | 3072 | +----------------+ -- ak, sk Doris> SELECT ARRAY_SIZE( -> EMBED( -> 'qwen_mul_embed', -> CAST('{ '> "uri": "s3:////test_img.png", '> "content_type": "image/png", '> "endpoint": "cos.ap-hongkong.myqcloud.com", '> "region": "ap-hongkong", '> "ak": "", '> "sk": "" '> }' AS JSON) -> ) -> ) AS img_embed_size_ak_sk; +----------------------+ | img_embed_size_ak_sk | +----------------------+ | 2560 | +----------------------+ ``` --- be/src/exprs/function/ai/ai_adapter.h | 413 ++++++++- be/src/exprs/function/ai/ai_filter.h | 2 + be/src/exprs/function/ai/ai_functions.h | 106 ++- be/src/exprs/function/ai/embed.h | 302 +++++- be/test/ai/ai_adapter_test.cpp | 202 +++- be/test/ai/ai_function_test.cpp | 875 ++++++++++++++++-- be/test/ai/embed_test.cpp | 653 ++++++++++++- .../org/apache/doris/catalog/Resource.java | 8 + .../property/constants/AIProperties.java | 2 +- .../doris/nereids/StatementContext.java | 17 +- .../expressions/functions/agg/AIAgg.java | 1 + .../expressions/functions/ai/AIFunction.java | 1 + .../trees/expressions/functions/ai/Embed.java | 48 +- .../java/org/apache/doris/qe/Coordinator.java | 26 +- .../org/apache/doris/qe/SessionVariable.java | 9 + .../doris/qe/runtime/ThriftPlansBuilder.java | 30 +- gensrc/thrift/PaloInternalService.thrift | 3 + .../suites/ai_p0/test_ai_functions.groovy | 5 +- 18 files changed, 2470 insertions(+), 233 deletions(-) diff --git a/be/src/exprs/function/ai/ai_adapter.h b/be/src/exprs/function/ai/ai_adapter.h index a50fa7123eb240..4e2ba693933202 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++) { @@ -1211,6 +1587,14 @@ class MockAdapter : 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::OK(); + } + Status parse_embedding_response(const std::string& response_body, std::vector>& results) const override { rapidjson::Document doc; @@ -1242,6 +1626,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_filter.h b/be/src/exprs/function/ai/ai_filter.h index cf66ed0b835f97..6d6962e81dd62d 100644 --- a/be/src/exprs/function/ai/ai_filter.h +++ b/be/src/exprs/function/ai/ai_filter.h @@ -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_functions.h b/be/src/exprs/function/ai/ai_functions.h index 528499992f5cf2..8458a1d8fe0133 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 @@ -45,6 +44,7 @@ #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" @@ -80,7 +80,7 @@ class AIFunction : public IFunction { 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; } @@ -92,8 +92,10 @@ class AIFunction : public IFunction { // 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; } @@ -123,11 +125,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 +146,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 +168,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 +182,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 +247,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 +255,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 +265,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"); @@ -278,7 +291,7 @@ class AIFunction : public IFunction { const TAIResource& config, std::shared_ptr& adapter, IColumn& col_result) const { 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)); @@ -290,23 +303,34 @@ class AIFunction : public IFunction { 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,8 +340,11 @@ 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)); } return Status::OK(); } @@ -333,38 +360,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/embed.h b/be/src/exprs/function/ai/embed.h index 4df5b732063369..2367e4b9459540 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 { @@ -33,8 +42,6 @@ class FunctionEmbed : public AIFunction { return std::make_shared(make_nullable(std::make_shared())); } - static FunctionPtr create() { return std::make_shared(); } - Status execute_with_adapter(FunctionContext* context, Block& block, const ColumnNumbers& arguments, uint32_t result, size_t input_rows_count, const TAIResource& config, @@ -44,11 +51,41 @@ class FunctionEmbed : public AIFunction { arguments.size()); } + PrimitiveType input_type = + remove_nullable(block.get_by_position(arguments[1]).type)->get_primitive_type(); + if (input_type == PrimitiveType::TYPE_JSONB) { + return _execute_multimodal_embed(context, block, arguments, result, input_rows_count, + config, adapter); + } + if (input_type == PrimitiveType::TYPE_STRING || input_type == PrimitiveType::TYPE_VARCHAR || + input_type == PrimitiveType::TYPE_CHAR) { + return _execute_text_embed(context, block, arguments, result, input_rows_count, config, + adapter); + } + return Status::InvalidArgument( + "Function EMBED expects the second argument to be STRING or JSON, but got type {}", + block.get_by_position(arguments[1]).type->get_name()); + } + + 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, + const ColumnNumbers& arguments, uint32_t result, + size_t input_rows_count, const TAIResource& config, + std::shared_ptr& adapter) 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)); @@ -57,22 +94,24 @@ class FunctionEmbed : public AIFunction { RETURN_IF_ERROR(build_prompt(block, arguments, 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 +120,73 @@ 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)); 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, + const ColumnNumbers& arguments, uint32_t result, + size_t input_rows_count, const TAIResource& config, + std::shared_ptr& adapter) 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; - 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); + const ColumnWithTypeAndName& file_column = block.get_by_position(arguments[1]); + for (size_t i = 0; i < input_rows_count; ++i) { + rapidjson::Document file_input; + RETURN_IF_ERROR(_parse_file_input(file_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, std::move(col_result)); 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 +198,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 +213,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 +278,116 @@ class FunctionEmbed : public AIFunction { auto& null_map = nested_nullable_col.get_null_map_column(); null_map.insert_many_vals(0, float_result.size()); } + + 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 ColumnWithTypeAndName& file_column, size_t row_num, + rapidjson::Document& file_input) { + std::string file_json = + JsonbToJson::jsonb_to_json_string(file_column.column->get_data_at(row_num).data, + file_column.column->get_data_at(row_num).size); + 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..aaaf64e7eaa37c 100644 --- a/be/test/ai/ai_adapter_test.cpp +++ b/be/test/ai/ai_adapter_test.cpp @@ -561,9 +561,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 +688,200 @@ 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_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, 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, 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.png", "image/png"}, + {MultimodalType::AUDIO, "https://a/b/c.mp3", "audio/mpeg"}, + {MultimodalType::VIDEO, "https://a/b/c.mp4", "video/mp4"}, + }; + + 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, 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(), "models/gemini-embedding-2-preview"); + ASSERT_TRUE(doc.HasMember("outputDimensionality")); + ASSERT_EQ(doc["outputDimensionality"].GetInt(), 768); + ASSERT_TRUE(doc.HasMember("content")); + ASSERT_TRUE(doc["content"].HasMember("parts")); + ASSERT_TRUE(doc["content"]["parts"].IsArray()); + ASSERT_EQ(doc["content"]["parts"].Size(), 1); + ASSERT_TRUE(doc["content"]["parts"][0].HasMember("file_data")); + ASSERT_TRUE(doc["content"]["parts"][0]["file_data"].IsObject()); + ASSERT_STREQ(doc["content"]["parts"][0]["file_data"]["mime_type"].GetString(), + test_case.mime_type); + ASSERT_STREQ(doc["content"]["parts"][0]["file_data"]["file_uri"].GetString(), + test_case.media_url); + } +} + } // namespace doris diff --git a/be/test/ai/ai_function_test.cpp b/be/test/ai/ai_function_test.cpp index 5a8777f7377bbd..23855611861733 100644 --- a/be/test/ai/ai_function_test.cpp +++ b/be/test/ai/ai_function_test.cpp @@ -15,10 +15,15 @@ // 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" @@ -26,6 +31,7 @@ #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" @@ -42,6 +48,140 @@ 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_return_type_impl(const DataTypes& arguments) const override { + 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) { auto nested_column = ColumnString::create(); @@ -63,6 +203,7 @@ MutableColumnPtr create_string_array_column(const std::vector(); 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 +446,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 +454,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 +504,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; @@ -376,6 +602,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 +621,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 +668,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 +719,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 +755,236 @@ 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); + + unsetenv("AI_TEST_RESULT"); } -TEST(AIFunctionTest, MockResourceBatchStringResult) { - setenv("AI_TEST_RESULT", R"(["first result","second result"])", 1); +TEST(AIFunctionTest, AIFilterBatchInvalidJson) { + auto runtime_state = std::make_unique(); + auto ctx = FunctionContext::create_context(runtime_state.get(), {}, {}); + + std::vector invalid_cases = {"1,0", "{}", "", " "}; + + 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 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 = {"first input", "second 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); 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()); - 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_at(0).to_string(), "first result"); - ASSERT_EQ(res_col.get_data_at(1).to_string(), "second result"); + 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], 1); + EXPECT_EQ(res_col.get_data()[2], 1); + + unsetenv("AI_TEST_RESULT"); } -TEST(AIFunctionTest, MockResourceBatchBoolResult) { - setenv("AI_TEST_RESULT", R"(["1","0"])", 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(), {}, {}); - auto runtime_state = std::make_unique(); + 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); + + 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 texts = {"valid input", "invalid input"}; + std::vector texts = {"small row", std::string(130 * 1024, 'x')}; auto col_resource = ColumnHelper::create_column(resources); auto col_text = ColumnHelper::create_column(texts); @@ -588,53 +997,150 @@ TEST(AIFunctionTest, MockResourceBatchBoolResult) { size_t result_idx = 2; auto filter_func = FunctionAIFilter::create(); + setenv("AI_TEST_RESULT", R"(["1"])", 1); 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_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_EQ(res_col.size(), 2); - ASSERT_EQ(res_col.get_data()[0], 1); - ASSERT_EQ(res_col.get_data()[1], 0); + EXPECT_EQ(res_col.get_data()[0], 1); + EXPECT_EQ(res_col.get_data()[1], 1); + + unsetenv("AI_TEST_RESULT"); } -TEST(AIFunctionTest, MockResourceBatchFloatResult) { - setenv("AI_TEST_RESULT", R"(["0.5","1.25"])", 1); +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", "mock_resource"}; - std::vector text1 = {"first text", "second text"}; - std::vector text2 = {"first compare", "second compare"}; + std::vector resources = {"mock_resource"}; + std::vector texts = {"test input"}; 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 sentiment_func = FunctionAISentiment::create(); Status exec_status = - similarity_func->execute_impl(ctx.get(), block, arguments, result_idx, text1.size()); - - unsetenv("AI_TEST_RESULT"); + 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); - 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); + 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, MissingAIResourcesMetadataTest) { @@ -658,12 +1164,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 +1237,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 +1260,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..2074697614252a 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,11 @@ #include #include +#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 "io/fs/obj_storage_client.h" #include "testutil/column_helper.h" #include "testutil/mock/mock_runtime_state.h" @@ -40,6 +45,129 @@ 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 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)); + } + } +} + TEST(EMBED_TEST, embed_function_build_test) { FunctionEmbed function; @@ -73,7 +201,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 +227,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 +254,417 @@ 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_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_with_adapter(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_with_adapter(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_with_adapter(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 +952,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 +965,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 +999,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 +1027,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 +1228,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 +1255,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 +1281,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 +1309,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;""" } From 66f178355c7035dee5085a7a5499acbdd36024a1 Mon Sep 17 00:00:00 2001 From: linrrarity Date: Wed, 5 Aug 2026 11:05:43 +0800 Subject: [PATCH 2/2] [Enhancement](ai_func) Skip Null inputs in AI functions (#66242) Problem Summary: The framework's default NULL implementation unwraps Nullable arguments and executes AI functions for every input row. For partially NULL inputs, the nested placeholder values of NULL rows are still included in prompts and sent to external AI providers. This causes unnecessary remote requests and token consumption. It also requires special handling for embedding results to preserve the original row order without duplicating large embedding vectors. This PR: - Disables the framework's default NULL implementation for all `AIFunction` subclasses. - Determines the Nullable return type in the `AIFunction` base class. - Extracts nested prompt columns and merges argument null maps in the common AI execution path. - Skips NULL rows before building prompts or sending requests. - Restores NULL rows in the final result while preserving the original row order. - Handles both text and multimodal Nullable inputs for `EMBED`. - Expands embedding array offsets in place, avoiding a copy of the nested Float32 embedding data. - Returns a constant NULL column when all input rows are NULL. ### Release note Fix AI scalar functions to skip NULL input rows instead of sending their placeholder values to external AI providers. --- be/src/exprs/function/ai/ai_adapter.h | 20 ++ be/src/exprs/function/ai/ai_classify.h | 4 +- be/src/exprs/function/ai/ai_extract.h | 4 +- be/src/exprs/function/ai/ai_filter.h | 2 +- be/src/exprs/function/ai/ai_fix_grammar.h | 2 +- be/src/exprs/function/ai/ai_functions.cpp | 164 ++++------- be/src/exprs/function/ai/ai_functions.h | 117 ++++++-- be/src/exprs/function/ai/ai_generate.h | 4 +- be/src/exprs/function/ai/ai_mask.h | 4 +- be/src/exprs/function/ai/ai_sentiment.h | 2 +- be/src/exprs/function/ai/ai_similarity.h | 4 +- be/src/exprs/function/ai/ai_summarize.h | 2 +- be/src/exprs/function/ai/ai_translate.h | 4 +- be/src/exprs/function/ai/embed.h | 111 +++++-- be/test/ai/ai_adapter_test.cpp | 275 +++++++++++++++-- be/test/ai/ai_function_test.cpp | 341 +++++++++++++++++++++- be/test/ai/embed_test.cpp | 195 ++++++++++++- 17 files changed, 1025 insertions(+), 230 deletions(-) diff --git a/be/src/exprs/function/ai/ai_adapter.h b/be/src/exprs/function/ai/ai_adapter.h index 4e2ba693933202..ba3a0039eb8beb 100644 --- a/be/src/exprs/function/ai/ai_adapter.h +++ b/be/src/exprs/function/ai/ai_adapter.h @@ -1569,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, @@ -1584,6 +1592,10 @@ class MockAdapter : public AIAdapter { Status build_embedding_request(const std::vector& inputs, std::string& request_body) const override { +#ifdef BE_TEST + auto& embedding_inputs = _embedding_inputs_for_test(); + embedding_inputs.insert(embedding_inputs.end(), inputs.begin(), inputs.end()); +#endif return Status::OK(); } @@ -1613,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 { 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 6d6962e81dd62d..e92c5991405f49 100644 --- a/be/src/exprs/function/ai/ai_filter.h +++ b/be/src/exprs/function/ai/ai_filter.h @@ -38,7 +38,7 @@ class FunctionAIFilter : public AIFunction { static constexpr size_t number_of_arguments = 2; - DataTypePtr get_return_type_impl(const DataTypes& arguments) const override { + DataTypePtr get_nested_return_type_impl(const DataTypes& /*arguments*/) const { return std::make_shared(); } diff --git a/be/src/exprs/function/ai/ai_fix_grammar.h b/be/src/exprs/function/ai/ai_fix_grammar.h index 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 8458a1d8fe0133..6fdd0414e47907 100644 --- a/be/src/exprs/function/ai/ai_functions.h +++ b/be/src/exprs/function/ai/ai_functions.h @@ -36,9 +36,11 @@ #include "core/column/column_nullable.h" #include "core/cow.h" #include "core/data_type/data_type_array.h" +#include "core/data_type/data_type_nullable.h" #include "core/data_type/data_type_number.h" #include "core/data_type/define_primitive_type.h" #include "core/data_type/primitive_type.h" +#include "exec/common/util.hpp" #include "exprs/function/ai/ai_adapter.h" #include "exprs/function/function.h" #include "runtime/query_context.h" @@ -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,6 +90,13 @@ class AIFunction : public IFunction { Status execute_impl(FunctionContext* context, Block& block, const ColumnNumbers& arguments, uint32_t result, size_t input_rows_count) const override { + if (block.get_by_position(arguments[0]).column->only_null()) { + block.get_by_position(result).column = + block.get_by_position(result).type->create_column_const(input_rows_count, + Field()); + return Status::OK(); + } + TAIResource config; std::shared_ptr adapter; if (Status status = this->_init_from_resource(context, block, arguments, config, adapter); @@ -84,8 +104,8 @@ class AIFunction : public IFunction { return status; } - return assert_cast(*this).execute_with_adapter( - context, block, arguments, result, input_rows_count, config, adapter); + return assert_cast(*this).execute(context, block, arguments, result, + input_rows_count, config, adapter); } protected: @@ -99,20 +119,6 @@ class AIFunction : public IFunction { return query_ctx->query_options().ai_context_window_size; } - // Derived classes can override this method for non-text/default behavior. - // The base implementation handles all string-input/string-output batchable functions. - Status execute_with_adapter(FunctionContext* context, Block& block, - const ColumnNumbers& arguments, uint32_t result, - size_t input_rows_count, const TAIResource& config, - std::shared_ptr& adapter) const { - auto col_result = assert_cast(*this).create_result_column(); - RETURN_IF_ERROR(execute_batched_prompts(context, block, arguments, input_rows_count, config, - adapter, *col_result)); - - block.replace_by_position(result, std::move(col_result)); - return Status::OK(); - } - MutableColumnPtr create_result_column() const { return ColumnString::create(); } // Provider-reusable hook for AI functions(string) -> string. @@ -286,19 +292,50 @@ class AIFunction : public IFunction { // Provider-reusable helper for string-returning functions. // Runs the common batch execution flow; derived classes only need to define how one batch of // string results is inserted into the final output column. - Status execute_batched_prompts(FunctionContext* context, Block& block, - const ColumnNumbers& arguments, size_t input_rows_count, - const TAIResource& config, std::shared_ptr& adapter, - IColumn& col_result) const { + Status execute(FunctionContext* context, Block& block, const ColumnNumbers& arguments, + uint32_t result, size_t input_rows_count, const TAIResource& config, + std::shared_ptr& adapter) const { + Columns prompt_columns; + prompt_columns.reserve(arguments.size() - 1); + ColumnUInt8::MutablePtr result_null_map; + for (size_t i = 1; i < arguments.size(); ++i) { + const auto& argument = block.get_by_position(arguments[i]); + if (argument.type->is_nullable()) { + const auto& [column, is_const] = unpack_if_const(argument.column); + const auto& nullable = + assert_cast(*column); + if (!result_null_map) { + result_null_map = ColumnUInt8::create(input_rows_count, 0); + } + VectorizedUtils::update_null_map(result_null_map->get_data(), + nullable.get_null_map_data(), is_const); + } + prompt_columns.emplace_back(argument.unnest_nullable().column); + } + + if (result_null_map && + !simd::contain_zero(result_null_map->get_data().data(), input_rows_count)) { + block.get_by_position(result).column = + block.get_by_position(result).type->create_column_const(input_rows_count, + Field()); + return Status::OK(); + } + + auto col_result = assert_cast(*this).create_result_column(); std::vector batch_prompts; size_t current_batch_size = 2; // [] const size_t max_batch_prompt_size = static_cast(get_ai_context_window_size(context)); + const NullMap* null_map = result_null_map ? &result_null_map->get_data() : nullptr; for (size_t i = 0; i < input_rows_count; ++i) { + if (null_map && (*null_map)[i]) { + continue; + } + std::string prompt; RETURN_IF_ERROR( - assert_cast(*this).build_prompt(block, arguments, i, prompt)); + assert_cast(*this).build_prompt(prompt_columns, i, prompt)); size_t entry_size = estimate_batch_entry_size(batch_prompts.size(), prompt); if (entry_size > max_batch_prompt_size) { @@ -307,7 +344,7 @@ class AIFunction : public IFunction { RETURN_IF_ERROR(this->execute_batch_request(batch_prompts, batch_results, config, adapter, context)); RETURN_IF_ERROR(assert_cast(*this).append_batch_results( - batch_results, col_result)); + batch_results, *col_result)); batch_prompts.clear(); current_batch_size = 2; } @@ -318,7 +355,7 @@ class AIFunction : public IFunction { RETURN_IF_ERROR(this->execute_batch_request(single_prompts, single_results, config, adapter, context)); RETURN_IF_ERROR(assert_cast(*this).append_batch_results( - single_results, col_result)); + single_results, *col_result)); continue; } @@ -329,7 +366,7 @@ class AIFunction : public IFunction { RETURN_IF_ERROR(this->execute_batch_request(batch_prompts, batch_results, config, adapter, context)); RETURN_IF_ERROR(assert_cast(*this).append_batch_results( - batch_results, col_result)); + batch_results, *col_result)); batch_prompts.clear(); current_batch_size = 2; additional_size = entry_size; @@ -344,8 +381,32 @@ class AIFunction : public IFunction { RETURN_IF_ERROR(this->execute_batch_request(batch_prompts, batch_results, config, adapter, context)); RETURN_IF_ERROR(assert_cast(*this).append_batch_results(batch_results, - col_result)); + *col_result)); + } + + if (!result_null_map) { + block.replace_by_position(result, std::move(col_result)); + return Status::OK(); + } + + if (!simd::contain_one(result_null_map->get_data().data(), input_rows_count)) { + block.replace_by_position(result, ColumnNullable::create(std::move(col_result), + std::move(result_null_map))); + return Status::OK(); } + + auto nested_result = col_result->clone_empty(); + size_t result_row = 0; + for (UInt8 is_null : result_null_map->get_data()) { + if (is_null) { + nested_result->insert_default(); + } else { + nested_result->insert_from(*col_result, result_row++); + } + } + + block.replace_by_position(result, ColumnNullable::create(std::move(nested_result), + std::move(result_null_map))); return Status::OK(); } diff --git a/be/src/exprs/function/ai/ai_generate.h b/be/src/exprs/function/ai/ai_generate.h index 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 2367e4b9459540..f193a1c171c1b5 100644 --- a/be/src/exprs/function/ai/embed.h +++ b/be/src/exprs/function/ai/embed.h @@ -38,33 +38,53 @@ class FunctionEmbed : public AIFunction { static constexpr auto system_prompt = ""; - DataTypePtr get_return_type_impl(const DataTypes& arguments) const override { + DataTypePtr get_nested_return_type_impl(const DataTypes& /*arguments*/) const { return std::make_shared(make_nullable(std::make_shared())); } - Status execute_with_adapter(FunctionContext* context, Block& block, - const ColumnNumbers& arguments, uint32_t result, - size_t input_rows_count, const TAIResource& config, - std::shared_ptr& adapter) const { + using PreparedFunctionImpl::execute; + + Status execute(FunctionContext* context, Block& block, const ColumnNumbers& arguments, + uint32_t result, size_t input_rows_count, const TAIResource& config, + std::shared_ptr& adapter) const { if (arguments.size() != 2) { return Status::InvalidArgument("Function EMBED expects 2 arguments, but got {}", arguments.size()); } - PrimitiveType input_type = - remove_nullable(block.get_by_position(arguments[1]).type)->get_primitive_type(); + const auto& input = block.get_by_position(arguments[1]); + ColumnUInt8::MutablePtr result_null_map; + if (input.type->is_nullable()) { + const auto& [column, is_const] = unpack_if_const(input.column); + const auto& nullable = + assert_cast(*column); + result_null_map = ColumnUInt8::create(input_rows_count, 0); + VectorizedUtils::update_null_map(result_null_map->get_data(), + nullable.get_null_map_data(), is_const); + } + + if (result_null_map && + !simd::contain_zero(result_null_map->get_data().data(), input_rows_count)) { + block.get_by_position(result).column = + block.get_by_position(result).type->create_column_const(input_rows_count, + Field()); + return Status::OK(); + } + + ColumnPtr input_column = input.unnest_nullable().column; + PrimitiveType input_type = remove_nullable(input.type)->get_primitive_type(); if (input_type == PrimitiveType::TYPE_JSONB) { - return _execute_multimodal_embed(context, block, arguments, result, input_rows_count, - config, adapter); + return _execute_multimodal_embed(context, block, result, input_rows_count, config, + adapter, input_column, std::move(result_null_map)); } if (input_type == PrimitiveType::TYPE_STRING || input_type == PrimitiveType::TYPE_VARCHAR || input_type == PrimitiveType::TYPE_CHAR) { - return _execute_text_embed(context, block, arguments, result, input_rows_count, config, - adapter); + return _execute_text_embed(context, block, result, input_rows_count, config, adapter, + input_column, std::move(result_null_map)); } return Status::InvalidArgument( "Function EMBED expects the second argument to be STRING or JSON, but got type {}", - block.get_by_position(arguments[1]).type->get_name()); + input.type->get_name()); } static FunctionPtr create() { return std::make_shared(); } @@ -77,10 +97,10 @@ class FunctionEmbed : public AIFunction { return query_ctx->query_options().embed_max_batch_size; } - Status _execute_text_embed(FunctionContext* context, Block& block, - const ColumnNumbers& arguments, uint32_t result, + Status _execute_text_embed(FunctionContext* context, Block& block, uint32_t result, size_t input_rows_count, const TAIResource& config, - std::shared_ptr& adapter) const { + std::shared_ptr& adapter, const ColumnPtr& input_column, + ColumnUInt8::MutablePtr result_null_map) const { auto col_result = ColumnArray::create( ColumnNullable::create(ColumnFloat32::create(), ColumnUInt8::create())); std::vector batch_prompts; @@ -88,10 +108,16 @@ class FunctionEmbed : public AIFunction { const int32_t max_batch_size = _get_embed_max_batch_size(context); const size_t max_context_window_size = static_cast(get_ai_context_window_size(context)); + const NullMap* null_map = result_null_map ? &result_null_map->get_data() : nullptr; + const Columns prompt_columns {input_column}; for (size_t i = 0; i < input_rows_count; ++i) { + if (null_map && (*null_map)[i]) { + continue; + } + std::string prompt; - RETURN_IF_ERROR(build_prompt(block, arguments, i, prompt)); + RETURN_IF_ERROR(build_prompt(prompt_columns, i, prompt)); const size_t prompt_size = prompt.size(); @@ -122,19 +148,23 @@ class FunctionEmbed : public AIFunction { RETURN_IF_ERROR( _flush_text_embedding_batch(batch_prompts, *col_result, config, adapter, context)); - block.replace_by_position(result, std::move(col_result)); + block.replace_by_position(result, _expand_and_wrap_nullable_result( + std::move(col_result), std::move(result_null_map), + input_rows_count)); return Status::OK(); } - Status _execute_multimodal_embed(FunctionContext* context, Block& block, - const ColumnNumbers& arguments, uint32_t result, + Status _execute_multimodal_embed(FunctionContext* context, Block& block, uint32_t result, size_t input_rows_count, const TAIResource& config, - std::shared_ptr& adapter) const { + std::shared_ptr& adapter, + const ColumnPtr& input_column, + ColumnUInt8::MutablePtr result_null_map) const { auto col_result = ColumnArray::create( ColumnNullable::create(ColumnFloat32::create(), ColumnUInt8::create())); std::vector batch_media_types; std::vector batch_media_content_types; std::vector batch_media_urls; + const NullMap* null_map = result_null_map ? &result_null_map->get_data() : nullptr; int64_t ttl_seconds = 3600; QueryContext* query_ctx = context->state()->get_query_ctx(); @@ -147,10 +177,13 @@ class FunctionEmbed : public AIFunction { const int32_t max_batch_size = _get_embed_max_batch_size(context); - const ColumnWithTypeAndName& file_column = block.get_by_position(arguments[1]); for (size_t i = 0; i < input_rows_count; ++i) { + if (null_map && (*null_map)[i]) { + continue; + } + rapidjson::Document file_input; - RETURN_IF_ERROR(_parse_file_input(file_column, i, file_input)); + RETURN_IF_ERROR(_parse_file_input(*input_column, i, file_input)); std::string content_type; MultimodalType media_type; @@ -175,7 +208,9 @@ class FunctionEmbed : public AIFunction { batch_media_types, batch_media_content_types, batch_media_urls, *col_result, config, adapter, context)); - block.replace_by_position(result, std::move(col_result)); + block.replace_by_position(result, _expand_and_wrap_nullable_result( + std::move(col_result), std::move(result_null_map), + input_rows_count)); return Status::OK(); } @@ -279,6 +314,29 @@ class FunctionEmbed : public AIFunction { null_map.insert_many_vals(0, float_result.size()); } + static ColumnPtr _expand_and_wrap_nullable_result(ColumnArray::MutablePtr result, + ColumnUInt8::MutablePtr result_null_map, + size_t input_rows_count) { + if (!result_null_map) { + return result; + } + + auto& offsets = result->get_offsets(); + size_t compact_row = offsets.size(); + offsets.resize(input_rows_count); + // For example, embedding rows 1 and 3 produces compact offsets [5, 10]. Given + // result_null_map [1, 0, 1, 0, 1], expand them to [0, 5, 5, 10, 10], where NULL rows + // reuse the previous offset. Fill backwards to avoid overwriting unread compact offsets. + for (size_t row = input_rows_count; row-- > 0;) { + if (result_null_map->get_data()[row]) { + offsets[row] = compact_row == 0 ? 0 : offsets[compact_row - 1]; + } else { + offsets[row] = offsets[--compact_row]; + } + } + return ColumnNullable::create(std::move(result), std::move(result_null_map)); + } + static bool _starts_with_ignore_case(std::string_view s, std::string_view prefix) { if (s.size() < prefix.size()) { return false; @@ -308,11 +366,10 @@ class FunctionEmbed : public AIFunction { } // Parse the FILE-like JSONB argument into a JSON object for downstream field reads. - static Status _parse_file_input(const ColumnWithTypeAndName& file_column, size_t row_num, + static Status _parse_file_input(const IColumn& file_column, size_t row_num, rapidjson::Document& file_input) { - std::string file_json = - JsonbToJson::jsonb_to_json_string(file_column.column->get_data_at(row_num).data, - file_column.column->get_data_at(row_num).size); + StringRef file_ref = file_column.get_data_at(row_num); + std::string file_json = JsonbToJson::jsonb_to_json_string(file_ref.data, file_ref.size); file_input.Parse(file_json.c_str()); DORIS_CHECK(!file_input.HasParseError() && file_input.IsObject()); return Status::OK(); diff --git a/be/test/ai/ai_adapter_test.cpp b/be/test/ai/ai_adapter_test.cpp index aaaf64e7eaa37c..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; @@ -696,8 +713,8 @@ TEST(AI_ADAPTER_TEST, qwen_multimodal_embedding_request_image) { adapter.init(config); std::string request_body; - Status st = adapter.build_multimodal_embedding_request(MultimodalType::IMAGE, - "https://a/b/c.png", 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; @@ -727,8 +744,8 @@ TEST(AI_ADAPTER_TEST, qwen_multimodal_embedding_request_video) { adapter.init(config); std::string request_body; - Status st = adapter.build_multimodal_embedding_request(MultimodalType::VIDEO, - "https://a/b/c.mp4", 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; @@ -750,6 +767,35 @@ TEST(AI_ADAPTER_TEST, qwen_multimodal_embedding_request_video) { 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; @@ -757,8 +803,8 @@ TEST(AI_ADAPTER_TEST, qwen_multimodal_embedding_request_audio_not_supported) { adapter.init(config); std::string request_body; - Status st = adapter.build_multimodal_embedding_request(MultimodalType::AUDIO, - "https://a/b/c.mp3", 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")); @@ -773,8 +819,8 @@ TEST(AI_ADAPTER_TEST, voyage_multimodal_embedding_request) { adapter.init(config); std::string request_body; - Status st = adapter.build_multimodal_embedding_request(MultimodalType::VIDEO, - "https://a/b/c.mp4", 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; @@ -797,6 +843,35 @@ TEST(AI_ADAPTER_TEST, voyage_multimodal_embedding_request) { 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; @@ -805,8 +880,8 @@ TEST(AI_ADAPTER_TEST, jina_multimodal_embedding_request) { adapter.init(config); std::string request_body; - Status st = adapter.build_multimodal_embedding_request(MultimodalType::IMAGE, - "https://a/b/c.jpg", 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; @@ -825,6 +900,34 @@ TEST(AI_ADAPTER_TEST, jina_multimodal_embedding_request) { 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; @@ -833,7 +936,7 @@ TEST(AI_ADAPTER_TEST, multimodal_provider_support) { std::string request_body; Status st = openai_adapter.build_multimodal_embedding_request( - MultimodalType::IMAGE, "https://a/b/c.png", request_body); + {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")); } @@ -852,36 +955,154 @@ TEST(AI_ADAPTER_TEST, gemini_multimodal_embedding_request) { const char* mime_type; }; const std::vector test_cases = { - {MultimodalType::IMAGE, "https://a/b/c.png", "image/png"}, - {MultimodalType::AUDIO, "https://a/b/c.mp3", "audio/mpeg"}, - {MultimodalType::VIDEO, "https://a/b/c.mp4", "video/mp4"}, + {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, request_body); + {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("model")); - ASSERT_STREQ(doc["model"].GetString(), "models/gemini-embedding-2-preview"); - ASSERT_TRUE(doc.HasMember("outputDimensionality")); - ASSERT_EQ(doc["outputDimensionality"].GetInt(), 768); - ASSERT_TRUE(doc.HasMember("content")); - ASSERT_TRUE(doc["content"].HasMember("parts")); - ASSERT_TRUE(doc["content"]["parts"].IsArray()); - ASSERT_EQ(doc["content"]["parts"].Size(), 1); - ASSERT_TRUE(doc["content"]["parts"][0].HasMember("file_data")); - ASSERT_TRUE(doc["content"]["parts"][0]["file_data"].IsObject()); - ASSERT_STREQ(doc["content"]["parts"][0]["file_data"]["mime_type"].GetString(), + 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(doc["content"]["parts"][0]["file_data"]["file_uri"].GetString(), + 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 23855611861733..976eaddd5bcb40 100644 --- a/be/test/ai/ai_function_test.cpp +++ b/be/test/ai/ai_function_test.cpp @@ -27,6 +27,7 @@ #include "core/block/block.h" #include "core/column/column_array.h" +#include "core/column/column_const.h" #include "core/column/column_nullable.h" #include "core/column/column_string.h" #include "core/column/column_vector.h" @@ -43,6 +44,7 @@ #include "exprs/function/ai/ai_summarize.h" #include "exprs/function/ai/ai_translate.h" #include "exprs/function/ai/embed.h" +#include "exprs/function/simple_function_factory.h" #include "testutil/column_helper.h" #include "testutil/mock/mock_runtime_state.h" @@ -63,7 +65,7 @@ class FunctionAIFilterBatchTestHelper : public AIFunction::execute_batch_request; - DataTypePtr get_return_type_impl(const DataTypes& arguments) const override { + DataTypePtr get_nested_return_type_impl(const DataTypes& /*arguments*/) const { return std::make_shared(); } @@ -183,16 +185,19 @@ class OneShotHttpServer { }; 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); @@ -202,6 +207,21 @@ 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; @@ -388,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, @@ -414,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"); @@ -593,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."); @@ -1143,6 +1195,269 @@ TEST(AIFunctionTest, AIStringFunctionBatchExecuteTest) { 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) { auto query_ctx = MockQueryContext::create(); TQueryOptions query_options; diff --git a/be/test/ai/embed_test.cpp b/be/test/ai/embed_test.cpp index 2074697614252a..c9bd32ed17cb7a 100644 --- a/be/test/ai/embed_test.cpp +++ b/be/test/ai/embed_test.cpp @@ -26,10 +26,12 @@ #include #include +#include "core/column/column_const.h" #include "core/data_type/data_type_jsonb.h" #include "core/data_type/data_type_number.h" #include "core/value/jsonb_value.h" #include "exprs/function/ai/ai_adapter.h" +#include "exprs/function/simple_function_factory.h" #include "io/fs/obj_storage_client.h" #include "testutil/column_helper.h" #include "testutil/mock/mock_runtime_state.h" @@ -151,6 +153,25 @@ static ColumnString::MutablePtr create_jsonb_column(const std::vector& json_rows, + const std::vector& null_map) { + EXPECT_EQ(json_rows.size(), null_map.size()); + auto column = ColumnString::create(); + auto null_column = ColumnUInt8::create(); + for (size_t i = 0; i < json_rows.size(); ++i) { + if (null_map[i]) { + column->insert_default(); + } else { + JsonBinaryValue jsonb_value; + Status st = jsonb_value.from_json_string(json_rows[i]); + EXPECT_TRUE(st.ok()) << st.to_string(); + column->insert_data(jsonb_value.value(), jsonb_value.size()); + } + null_column->insert_value(null_map[i]); + } + return ColumnNullable::create(std::move(column), std::move(null_column)); +} + static void assert_mock_embedding_column(const ColumnArray& col_array, size_t row_count) { const auto& offsets = col_array.get_offsets(); ASSERT_EQ(offsets.size(), row_count); @@ -168,6 +189,39 @@ static void assert_mock_embedding_column(const ColumnArray& col_array, size_t ro } } +static void assert_mock_nullable_embedding_column(const IColumn& column, + const std::vector& expected_null_map) { + const auto& nullable_column = assert_cast(column); + ASSERT_EQ(nullable_column.size(), expected_null_map.size()); + + const auto& col_array = assert_cast(nullable_column.get_nested_column()); + const auto& offsets = col_array.get_offsets(); + const auto& nested_nullable_col = assert_cast(col_array.get_data()); + const auto& nested_col = + assert_cast(*nested_nullable_col.get_nested_column_ptr()); + + size_t expected_offset = 0; + for (size_t row = 0; row < expected_null_map.size(); ++row) { + ASSERT_EQ(nullable_column.is_null_at(row), expected_null_map[row] != 0); + if (expected_null_map[row]) { + ASSERT_EQ(offsets[row], expected_offset); + continue; + } + + expected_offset += 5; + ASSERT_EQ(offsets[row], expected_offset); + for (size_t i = 0; i < 5; ++i) { + ASSERT_FLOAT_EQ(nested_col.get_element(expected_offset - 5 + i), static_cast(i)); + } + } + ASSERT_EQ(nested_col.size(), expected_offset); +} + +static FunctionBasePtr get_embed_function(const Block& block, const DataTypePtr& return_type) { + return SimpleFunctionFactory::instance().get_function( + "embed", block.get_columns_with_type_and_name(), return_type); +} + TEST(EMBED_TEST, embed_function_build_test) { FunctionEmbed function; @@ -183,7 +237,7 @@ TEST(EMBED_TEST, embed_function_build_test) { ColumnNumbers arguments = {0, 1}; std::string prompt; - Status status = function.build_prompt(block, arguments, 0, prompt); + Status status = function.build_prompt({block.get_by_position(arguments[1]).column}, 0, prompt); ASSERT_TRUE(status.ok()); ASSERT_EQ(prompt, "this is a test prompt"); @@ -297,6 +351,133 @@ TEST(EMBED_TEST, embed_function_multimodal_direct_url) { assert_mock_embedding_column(col_array, file_json_rows.size()); } +TEST(EMBED_TEST, embed_function_partial_null_through_framework) { + auto runtime_state = std::make_unique(); + auto ctx = FunctionContext::create_context(runtime_state.get(), {}, {}); + + std::vector resources(5, "mock_resource"); + std::vector texts = {"", "text-a", "", "text-c", ""}; + std::vector null_map = {1, 0, 1, 0, 1}; + auto col_resource = ColumnHelper::create_column(resources); + auto col_text = ColumnHelper::create_nullable_column(texts, null_map); + auto return_type = make_nullable( + std::make_shared(make_nullable(std::make_shared()))); + + Block block; + block.insert({std::move(col_resource), std::make_shared(), "resource"}); + block.insert({std::move(col_text), make_nullable(std::make_shared()), "text"}); + + auto function = get_embed_function(block, return_type); + ASSERT_NE(function, nullptr); + EXPECT_TRUE(function->get_return_type()->equals(*return_type)); + + block.insert({nullptr, return_type, "result"}); + const size_t result_idx = 2; + MockAdapter::clear_embedding_inputs_for_test(); + Status exec_status = function->execute(ctx.get(), block, {0, 1}, result_idx, texts.size()); + + ASSERT_TRUE(exec_status.ok()) << exec_status.to_string(); + EXPECT_THAT(MockAdapter::get_embedding_inputs_for_test(), + ::testing::ElementsAre("text-a", "text-c")); + assert_mock_nullable_embedding_column(*block.get_by_position(result_idx).column, null_map); +} + +TEST(EMBED_TEST, embed_function_multimodal_partial_null_through_framework) { + auto runtime_state = std::make_unique(); + auto ctx = FunctionContext::create_context(runtime_state.get(), {}, {}); + + std::vector resources(5, "mock_resource"); + std::vector file_json_rows = { + "", R"({"content_type":"image/png","uri":"https://example.com/a.png"})", "", + R"({"content_type":"video/mp4","uri":"https://example.com/b.mp4"})", ""}; + std::vector null_map = {1, 0, 1, 0, 1}; + auto col_resource = ColumnHelper::create_column(resources); + auto col_file = create_nullable_jsonb_column(file_json_rows, null_map); + auto return_type = make_nullable( + std::make_shared(make_nullable(std::make_shared()))); + + Block block; + block.insert({std::move(col_resource), std::make_shared(), "resource"}); + block.insert({std::move(col_file), make_nullable(std::make_shared()), "file"}); + + auto function = get_embed_function(block, return_type); + ASSERT_NE(function, nullptr); + + block.insert({nullptr, return_type, "result"}); + const size_t result_idx = 2; + Status exec_status = + function->execute(ctx.get(), block, {0, 1}, result_idx, file_json_rows.size()); + + ASSERT_TRUE(exec_status.ok()) << exec_status.to_string(); + assert_mock_nullable_embedding_column(*block.get_by_position(result_idx).column, null_map); +} + +TEST(EMBED_TEST, embed_function_all_null_const_nullable_through_framework) { + auto runtime_state = std::make_unique(); + auto ctx = FunctionContext::create_context(runtime_state.get(), {}, {}); + + constexpr size_t row_count = 5; + std::vector resources(row_count, "mock_resource"); + auto col_resource = ColumnHelper::create_column(resources); + auto nullable_text = + ColumnHelper::create_nullable_column({""}, std::vector {1}); + auto col_text = ColumnConst::create(std::move(nullable_text), row_count); + auto return_type = make_nullable( + std::make_shared(make_nullable(std::make_shared()))); + + Block block; + block.insert({std::move(col_resource), std::make_shared(), "resource"}); + block.insert({std::move(col_text), make_nullable(std::make_shared()), "text"}); + + auto function = get_embed_function(block, return_type); + ASSERT_NE(function, nullptr); + + block.insert({nullptr, return_type, "result"}); + const size_t result_idx = 2; + Status exec_status = function->execute(ctx.get(), block, {0, 1}, result_idx, row_count); + + ASSERT_TRUE(exec_status.ok()) << exec_status.to_string(); + const auto& result_column = block.get_by_position(result_idx).column; + ASSERT_TRUE(is_column_const(*result_column)); + EXPECT_TRUE(result_column->only_null()); + ColumnPtr full_result = result_column->convert_to_full_column_if_const(); + assert_mock_nullable_embedding_column(*full_result, std::vector(row_count, 1)); +} + +TEST(EMBED_TEST, embed_function_null_rows_across_batches_through_framework) { + TQueryOptions query_options = create_fake_query_options(); + query_options.__set_embed_max_batch_size(2); + auto query_ctx = MockQueryContext::create(TUniqueId(), ExecEnv::GetInstance(), query_options); + query_ctx->set_mock_ai_resource(); + TQueryGlobals query_globals; + RuntimeState runtime_state(TUniqueId(), 0, query_options, query_globals, nullptr, + query_ctx.get()); + auto ctx = FunctionContext::create_context(&runtime_state, {}, {}); + + std::vector texts = {"", "text-a", "", "text-b", "", "text-c", + "", "text-d", "", "text-e", ""}; + std::vector null_map = {1, 0, 1, 0, 1, 0, 1, 0, 1, 0, 1}; + std::vector resources(texts.size(), "mock_resource"); + auto col_resource = ColumnHelper::create_column(resources); + auto col_text = ColumnHelper::create_nullable_column(texts, null_map); + auto return_type = make_nullable( + std::make_shared(make_nullable(std::make_shared()))); + + Block block; + block.insert({std::move(col_resource), std::make_shared(), "resource"}); + block.insert({std::move(col_text), make_nullable(std::make_shared()), "text"}); + + auto function = get_embed_function(block, return_type); + ASSERT_NE(function, nullptr); + + block.insert({nullptr, return_type, "result"}); + const size_t result_idx = 2; + Status exec_status = function->execute(ctx.get(), block, {0, 1}, result_idx, texts.size()); + + ASSERT_TRUE(exec_status.ok()) << exec_status.to_string(); + assert_mock_nullable_embedding_column(*block.get_by_position(result_idx).column, null_map); +} + TEST(EMBED_TEST, embed_function_multimodal_batch_request) { auto runtime_state = std::make_unique(); auto ctx = FunctionContext::create_context(runtime_state.get(), {}, {}); @@ -327,8 +508,8 @@ TEST(EMBED_TEST, embed_function_multimodal_batch_request) { ColumnNumbers arguments = {0, 1}; size_t result_idx = 2; FunctionEmbed embed_func; - Status exec_status = embed_func.execute_with_adapter(ctx.get(), block, arguments, result_idx, - file_json_rows.size(), config, adapter); + Status exec_status = embed_func.execute(ctx.get(), block, arguments, result_idx, + file_json_rows.size(), config, adapter); ASSERT_TRUE(exec_status.ok()) << exec_status.to_string(); EXPECT_THAT(counting_adapter->batch_sizes, ::testing::ElementsAre(3)); @@ -373,8 +554,8 @@ TEST(EMBED_TEST, embed_function_multimodal_batch_split_by_session_variable) { ColumnNumbers arguments = {0, 1}; size_t result_idx = 2; FunctionEmbed embed_func; - Status exec_status = embed_func.execute_with_adapter(ctx.get(), block, arguments, result_idx, - file_json_rows.size(), config, adapter); + Status exec_status = embed_func.execute(ctx.get(), block, arguments, result_idx, + file_json_rows.size(), config, adapter); ASSERT_TRUE(exec_status.ok()) << exec_status.to_string(); EXPECT_THAT(counting_adapter->batch_sizes, ::testing::ElementsAre(2, 1)); @@ -416,8 +597,8 @@ TEST(EMBED_TEST, embed_function_text_batch_split_by_session_variable) { ColumnNumbers arguments = {0, 1}; size_t result_idx = 2; FunctionEmbed embed_func; - Status exec_status = embed_func.execute_with_adapter(ctx.get(), block, arguments, result_idx, - texts.size(), config, adapter); + Status exec_status = embed_func.execute(ctx.get(), block, arguments, result_idx, texts.size(), + config, adapter); ASSERT_TRUE(exec_status.ok()) << exec_status.to_string(); EXPECT_THAT(counting_adapter->batch_sizes, ::testing::ElementsAre(2, 1));