Coverage Report

Created: 2026-09-29 12:58

next uncovered line (L), next uncovered region (R), next uncovered branch (B)
be/src/exprs/function/ai/embed.h
Line
Count
Source
1
// Licensed to the Apache Software Foundation (ASF) under one
2
// or more contributor license agreements.  See the NOTICE file
3
// distributed with this work for additional information
4
// regarding copyright ownership.  The ASF licenses this file
5
// to you under the Apache License, Version 2.0 (the
6
// "License"); you may not use this file except in compliance
7
// with the License.  You may obtain a copy of the License at
8
//
9
//   http://www.apache.org/licenses/LICENSE-2.0
10
//
11
// Unless required by applicable law or agreed to in writing,
12
// software distributed under the License is distributed on an
13
// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
14
// KIND, either express or implied.  See the License for the
15
// specific language governing permissions and limitations
16
// under the License.
17
18
#pragma once
19
20
#include <glog/logging.h>
21
#include <rapidjson/document.h>
22
23
#include <string_view>
24
25
#include "core/data_type/data_type_nullable.h"
26
#include "core/data_type/primitive_type.h"
27
#include "exprs/function/ai/ai_functions.h"
28
#include "util/jsonb_utils.h"
29
#include "util/s3_uri.h"
30
#include "util/s3_util.h"
31
32
namespace doris {
33
class FunctionEmbed : public AIFunction<FunctionEmbed> {
34
public:
35
    static constexpr auto name = "embed";
36
37
    static constexpr size_t number_of_arguments = 2;
38
39
    static constexpr auto system_prompt = "";
40
41
10
    DataTypePtr get_nested_return_type_impl(const DataTypes& /*arguments*/) const {
42
10
        return std::make_shared<DataTypeArray>(make_nullable(std::make_shared<DataTypeFloat32>()));
43
10
    }
44
45
    using PreparedFunctionImpl::execute;
46
47
46
    AIResource select_ai_resource(const TAIResource& resource, PrimitiveType input_type) const {
48
46
        bool has_complete_multimodal_embed_properties = _has_complete_resource_properties(
49
46
                resource.embed_mm_endpoint, resource.embed_mm_provider_type,
50
46
                resource.embed_mm_model_name, resource.embed_mm_api_key);
51
46
        if (input_type == PrimitiveType::TYPE_JSONB && has_complete_multimodal_embed_properties) {
52
2
            return AIResource::from_multimodal_embed(resource);
53
2
        }
54
44
        bool has_complete_embed_properties = _has_complete_resource_properties(
55
44
                resource.embed_endpoint, resource.embed_provider_type, resource.embed_model_name,
56
44
                resource.embed_api_key);
57
44
        return has_complete_embed_properties ? AIResource::from_embed(resource)
58
44
                                             : AIResource(resource);
59
46
    }
60
61
    Status execute(FunctionContext* context, Block& block, const ColumnNumbers& arguments,
62
                   uint32_t result, size_t input_rows_count, const AIResource& config,
63
36
                   std::shared_ptr<AIAdapter>& adapter) const {
64
36
        if (arguments.size() != 2) {
65
2
            return Status::InvalidArgument("Function EMBED expects 2 arguments, but got {}",
66
2
                                           arguments.size());
67
2
        }
68
69
34
        const auto& input = block.get_by_position(arguments[1]);
70
34
        ColumnUInt8::MutablePtr result_null_map;
71
34
        if (input.type->is_nullable()) {
72
8
            const auto& [column, is_const] = unpack_if_const(input.column);
73
8
            const auto& nullable =
74
8
                    assert_cast<const ColumnNullable&, TypeCheckOnRelease::DISABLE>(*column);
75
8
            result_null_map = ColumnUInt8::create(input_rows_count, 0);
76
8
            VectorizedUtils::update_null_map(result_null_map->get_data(),
77
8
                                             nullable.get_null_map_data(), is_const);
78
8
        }
79
80
34
        if (result_null_map &&
81
34
            !simd::contain_zero(result_null_map->get_data().data(), input_rows_count)) {
82
2
            block.get_by_position(result).column =
83
2
                    block.get_by_position(result).type->create_column_const(input_rows_count,
84
2
                                                                            Field());
85
2
            return Status::OK();
86
2
        }
87
88
32
        ColumnPtr input_column =
89
32
                input.unnest_nullable(input.type->is_nullable() ? input.get_nullable_column_info()
90
32
                                                                : NullableColumnInfo {},
91
32
                                      false)
92
32
                        .column;
93
32
        PrimitiveType input_type = remove_nullable(input.type)->get_primitive_type();
94
32
        if (input_type == PrimitiveType::TYPE_JSONB) {
95
18
            return _execute_multimodal_embed(context, block, result, input_rows_count, config,
96
18
                                             adapter, input_column, std::move(result_null_map));
97
18
        }
98
14
        if (input_type == PrimitiveType::TYPE_STRING || input_type == PrimitiveType::TYPE_VARCHAR ||
99
14
            input_type == PrimitiveType::TYPE_CHAR) {
100
12
            return _execute_text_embed(context, block, result, input_rows_count, config, adapter,
101
12
                                       input_column, std::move(result_null_map));
102
12
        }
103
2
        return Status::InvalidArgument(
104
2
                "Function EMBED expects the second argument to be STRING or JSON, but got type {}",
105
2
                input.type->get_name());
106
14
    }
107
108
34
    static FunctionPtr create() { return std::make_shared<FunctionEmbed>(); }
109
110
private:
111
    static bool _has_complete_resource_properties(std::string_view endpoint,
112
                                                  std::string_view provider_type,
113
                                                  std::string_view model_name,
114
90
                                                  std::string_view api_key) {
115
90
        return !endpoint.empty() && !provider_type.empty() && !model_name.empty() &&
116
90
               (provider_type == "LOCAL" || !api_key.empty());
117
90
    }
118
119
30
    static int32_t _get_embed_max_batch_size(FunctionContext* context) {
120
30
        QueryContext* query_ctx = context->state()->get_query_ctx();
121
30
        DORIS_CHECK(query_ctx != nullptr);
122
123
30
        return query_ctx->query_options().embed_max_batch_size;
124
30
    }
125
126
    Status _execute_text_embed(FunctionContext* context, Block& block, uint32_t result,
127
                               size_t input_rows_count, const AIResource& config,
128
                               std::shared_ptr<AIAdapter>& adapter, const ColumnPtr& input_column,
129
12
                               ColumnUInt8::MutablePtr result_null_map) const {
130
12
        auto col_result = ColumnArray::create(
131
12
                ColumnNullable::create(ColumnFloat32::create(), ColumnUInt8::create()));
132
12
        std::vector<std::string> batch_prompts;
133
12
        size_t current_batch_size = 0;
134
12
        const int32_t max_batch_size = _get_embed_max_batch_size(context);
135
12
        const size_t max_context_window_size =
136
12
                static_cast<size_t>(get_ai_context_window_size(context));
137
12
        const NullMap* null_map = result_null_map ? &result_null_map->get_data() : nullptr;
138
12
        const Columns prompt_columns {input_column};
139
140
58
        for (size_t i = 0; i < input_rows_count; ++i) {
141
46
            if (null_map && (*null_map)[i]) {
142
18
                continue;
143
18
            }
144
145
28
            std::string prompt;
146
28
            RETURN_IF_ERROR(build_prompt(prompt_columns, i, prompt));
147
148
28
            const size_t prompt_size = prompt.size();
149
150
28
            if (prompt_size > max_context_window_size) {
151
                // flush history batch
152
0
                RETURN_IF_ERROR(_flush_text_embedding_batch(batch_prompts, *col_result, config,
153
0
                                                            adapter, context));
154
0
                current_batch_size = 0;
155
156
0
                batch_prompts.emplace_back(std::move(prompt));
157
0
                RETURN_IF_ERROR(_flush_text_embedding_batch(batch_prompts, *col_result, config,
158
0
                                                            adapter, context));
159
0
                continue;
160
0
            }
161
162
28
            if (!batch_prompts.empty() &&
163
28
                (current_batch_size + prompt_size > max_context_window_size ||
164
16
                 batch_prompts.size() >= static_cast<size_t>(max_batch_size))) {
165
6
                RETURN_IF_ERROR(_flush_text_embedding_batch(batch_prompts, *col_result, config,
166
6
                                                            adapter, context));
167
6
                current_batch_size = 0;
168
6
            }
169
170
28
            batch_prompts.emplace_back(std::move(prompt));
171
28
            current_batch_size += prompt_size;
172
28
        }
173
174
12
        RETURN_IF_ERROR(
175
12
                _flush_text_embedding_batch(batch_prompts, *col_result, config, adapter, context));
176
177
12
        block.replace_by_position(result, _expand_and_wrap_nullable_result(
178
12
                                                  std::move(col_result), std::move(result_null_map),
179
12
                                                  input_rows_count));
180
12
        return Status::OK();
181
12
    }
182
183
    Status _execute_multimodal_embed(FunctionContext* context, Block& block, uint32_t result,
184
                                     size_t input_rows_count, const AIResource& config,
185
                                     std::shared_ptr<AIAdapter>& adapter,
186
                                     const ColumnPtr& input_column,
187
18
                                     ColumnUInt8::MutablePtr result_null_map) const {
188
18
        auto col_result = ColumnArray::create(
189
18
                ColumnNullable::create(ColumnFloat32::create(), ColumnUInt8::create()));
190
18
        std::vector<MultimodalType> batch_media_types;
191
18
        std::vector<std::string> batch_media_content_types;
192
18
        std::vector<std::string> batch_media_urls;
193
18
        const NullMap* null_map = result_null_map ? &result_null_map->get_data() : nullptr;
194
195
18
        int64_t ttl_seconds = 3600;
196
18
        QueryContext* query_ctx = context->state()->get_query_ctx();
197
18
        if (query_ctx && query_ctx->query_options().__isset.file_presigned_url_ttl_seconds) {
198
18
            ttl_seconds = query_ctx->query_options().file_presigned_url_ttl_seconds;
199
18
            if (ttl_seconds <= 0) {
200
2
                ttl_seconds = 3600;
201
2
            }
202
18
        }
203
204
18
        const int32_t max_batch_size = _get_embed_max_batch_size(context);
205
206
48
        for (size_t i = 0; i < input_rows_count; ++i) {
207
38
            if (null_map && (*null_map)[i]) {
208
6
                continue;
209
6
            }
210
211
32
            rapidjson::Document file_input;
212
32
            RETURN_IF_ERROR(_parse_file_input(*input_column, i, file_input));
213
214
32
            std::string content_type;
215
32
            MultimodalType media_type;
216
32
            RETURN_IF_ERROR(_infer_media_type(file_input, content_type, media_type));
217
218
28
            std::string media_url;
219
28
            RETURN_IF_ERROR(_resolve_media_url(file_input, ttl_seconds, media_url));
220
221
24
            if (!batch_media_urls.empty() &&
222
24
                batch_media_urls.size() >= static_cast<size_t>(max_batch_size)) {
223
2
                RETURN_IF_ERROR(_flush_multimodal_embedding_batch(
224
2
                        batch_media_types, batch_media_content_types, batch_media_urls, *col_result,
225
2
                        config, adapter, context));
226
2
            }
227
228
24
            batch_media_types.emplace_back(media_type);
229
24
            batch_media_content_types.emplace_back(std::move(content_type));
230
24
            batch_media_urls.emplace_back(std::move(media_url));
231
24
        }
232
233
10
        RETURN_IF_ERROR(_flush_multimodal_embedding_batch(
234
10
                batch_media_types, batch_media_content_types, batch_media_urls, *col_result, config,
235
10
                adapter, context));
236
237
10
        block.replace_by_position(result, _expand_and_wrap_nullable_result(
238
10
                                                  std::move(col_result), std::move(result_null_map),
239
10
                                                  input_rows_count));
240
10
        return Status::OK();
241
10
    }
242
243
    // EMBED-private helper.
244
    // Sends one embedding request with a prebuilt request body and validates returned row count.
245
    Status _execute_prebuilt_embedding_request(const std::string& request_body,
246
                                               std::vector<std::vector<float>>& results,
247
                                               size_t expected_size, const AIResource& config,
248
                                               std::shared_ptr<AIAdapter>& adapter,
249
30
                                               FunctionContext* context) const {
250
30
        std::string response;
251
30
#ifdef BE_TEST
252
30
        if (config.provider_type == "MOCK") {
253
30
            results.clear();
254
30
            results.reserve(expected_size);
255
82
            for (size_t i = 0; i < expected_size; ++i) {
256
52
                results.emplace_back(std::initializer_list<float> {0, 1, 2, 3, 4});
257
52
            }
258
30
            return Status::OK();
259
30
        }
260
0
#endif
261
262
0
        RETURN_IF_ERROR(
263
0
                this->send_request_to_llm(request_body, response, config, adapter, context));
264
265
0
        RETURN_IF_ERROR(adapter->parse_embedding_response(response, results));
266
0
        if (results.empty()) {
267
0
            return Status::InternalError("AI returned empty result");
268
0
        }
269
0
        if (results.size() != expected_size) [[unlikely]] {
270
0
            return Status::InternalError(
271
0
                    "AI embedding returned {} results, but {} inputs were sent", results.size(),
272
0
                    expected_size);
273
0
        }
274
0
        return Status::OK();
275
0
    }
276
277
    // EMBED-private helper.
278
    // Flushes one accumulated text embedding batch into the output array column.
279
    Status _flush_text_embedding_batch(std::vector<std::string>& batch_prompts,
280
                                       ColumnArray& col_result, const AIResource& config,
281
                                       std::shared_ptr<AIAdapter>& adapter,
282
18
                                       FunctionContext* context) const {
283
18
        if (batch_prompts.empty()) {
284
0
            return Status::OK();
285
0
        }
286
287
18
        std::string request_body;
288
18
        RETURN_IF_ERROR(adapter->build_embedding_request(batch_prompts, request_body));
289
18
        std::vector<std::vector<float>> batch_results;
290
18
        RETURN_IF_ERROR(_execute_prebuilt_embedding_request(
291
18
                request_body, batch_results, batch_prompts.size(), config, adapter, context));
292
28
        for (const auto& batch_result : batch_results) {
293
28
            _insert_embedding_result(col_result, batch_result);
294
28
        }
295
18
        batch_prompts.clear();
296
18
        return Status::OK();
297
18
    }
298
299
    // EMBED-private helper.
300
    // Flushes one accumulated multimodal embedding batch into the output array column.
301
    Status _flush_multimodal_embedding_batch(std::vector<MultimodalType>& batch_media_types,
302
                                             std::vector<std::string>& batch_media_content_types,
303
                                             std::vector<std::string>& batch_media_urls,
304
                                             ColumnArray& col_result, const AIResource& config,
305
                                             std::shared_ptr<AIAdapter>& adapter,
306
12
                                             FunctionContext* context) const {
307
12
        if (batch_media_urls.empty()) {
308
0
            return Status::OK();
309
0
        }
310
311
12
        std::string request_body;
312
12
        RETURN_IF_ERROR(adapter->build_multimodal_embedding_request(
313
12
                batch_media_types, batch_media_urls, batch_media_content_types, request_body));
314
315
12
        std::vector<std::vector<float>> batch_results;
316
12
        RETURN_IF_ERROR(_execute_prebuilt_embedding_request(
317
12
                request_body, batch_results, batch_media_urls.size(), config, adapter, context));
318
24
        for (const auto& batch_result : batch_results) {
319
24
            _insert_embedding_result(col_result, batch_result);
320
24
        }
321
12
        batch_media_types.clear();
322
12
        batch_media_content_types.clear();
323
12
        batch_media_urls.clear();
324
12
        return Status::OK();
325
12
    }
326
327
    static void _insert_embedding_result(ColumnArray& col_array,
328
52
                                         const std::vector<float>& float_result) {
329
52
        auto& offsets = col_array.get_offsets();
330
52
        auto& nested_nullable_col = assert_cast<ColumnNullable&>(col_array.get_data());
331
52
        auto& nested_col =
332
52
                assert_cast<ColumnFloat32&>(*(nested_nullable_col.get_nested_column_ptr()));
333
52
        nested_col.reserve(nested_col.size() + float_result.size());
334
335
52
        size_t current_offset = nested_col.size();
336
52
        nested_col.insert_many_raw_data(reinterpret_cast<const char*>(float_result.data()),
337
52
                                        float_result.size());
338
52
        offsets.push_back(current_offset + float_result.size());
339
52
        auto& null_map = nested_nullable_col.get_null_map_column();
340
52
        null_map.insert_many_vals(0, float_result.size());
341
52
    }
342
343
    static ColumnPtr _expand_and_wrap_nullable_result(ColumnArray::MutablePtr result,
344
                                                      ColumnUInt8::MutablePtr result_null_map,
345
22
                                                      size_t input_rows_count) {
346
22
        if (!result_null_map) {
347
16
            return result;
348
16
        }
349
350
6
        auto& offsets = result->get_offsets();
351
6
        size_t compact_row = offsets.size();
352
6
        offsets.resize(input_rows_count);
353
        // For example, embedding rows 1 and 3 produces compact offsets [5, 10]. Given
354
        // result_null_map [1, 0, 1, 0, 1], expand them to [0, 5, 5, 10, 10], where NULL rows
355
        // reuse the previous offset. Fill backwards to avoid overwriting unread compact offsets.
356
48
        for (size_t row = input_rows_count; row-- > 0;) {
357
42
            if (result_null_map->get_data()[row]) {
358
24
                offsets[row] = compact_row == 0 ? 0 : offsets[compact_row - 1];
359
24
            } else {
360
18
                offsets[row] = offsets[--compact_row];
361
18
            }
362
42
        }
363
6
        return ColumnNullable::create(std::move(result), std::move(result_null_map));
364
22
    }
365
366
104
    static bool _starts_with_ignore_case(std::string_view s, std::string_view prefix) {
367
104
        if (s.size() < prefix.size()) {
368
0
            return false;
369
0
        }
370
486
        return std::equal(prefix.begin(), prefix.end(), s.begin(), [](char a, char b) {
371
486
            return std::tolower(static_cast<unsigned char>(a)) ==
372
486
                   std::tolower(static_cast<unsigned char>(b));
373
486
        });
374
104
    }
375
376
    static Status _infer_media_type(const rapidjson::Value& file_input, std::string& content_type,
377
32
                                    MultimodalType& media_type) {
378
32
        RETURN_IF_ERROR(_get_required_string_field(file_input, "content_type", content_type));
379
380
30
        if (_starts_with_ignore_case(content_type, "image/")) {
381
18
            media_type = MultimodalType::IMAGE;
382
18
            return Status::OK();
383
18
        } else if (_starts_with_ignore_case(content_type, "video/")) {
384
6
            media_type = MultimodalType::VIDEO;
385
6
            return Status::OK();
386
6
        } else if (_starts_with_ignore_case(content_type, "audio/")) {
387
4
            media_type = MultimodalType::AUDIO;
388
4
            return Status::OK();
389
4
        }
390
391
2
        return Status::InvalidArgument("Unsupported content_type for EMBED: {}", content_type);
392
30
    }
393
394
    // Parse the FILE-like JSONB argument into a JSON object for downstream field reads.
395
    static Status _parse_file_input(const IColumn& file_column, size_t row_num,
396
32
                                    rapidjson::Document& file_input) {
397
32
        StringRef file_ref = file_column.get_data_at(row_num);
398
32
        std::string file_json = JsonbToJson::jsonb_to_json_string(file_ref.data, file_ref.size);
399
32
        file_input.Parse(file_json.c_str());
400
32
        DORIS_CHECK(!file_input.HasParseError() && file_input.IsObject());
401
32
        return Status::OK();
402
32
    }
403
404
    // TODO(lzq): After support FILE type, We should use the interface provided by FILE to get the fields
405
    // replacing this function
406
    static Status _get_required_string_field(const rapidjson::Value& obj, const char* field_name,
407
70
                                             std::string& value) {
408
70
        auto iter = obj.FindMember(field_name);
409
70
        if (iter == obj.MemberEnd() || !iter->value.IsString()) {
410
6
            return Status::InvalidArgument(
411
6
                    "EMBED file json field '{}' is required and must be a string", field_name);
412
6
        }
413
64
        value = iter->value.GetString();
414
64
        if (value.empty()) {
415
0
            return Status::InvalidArgument("EMBED file json field '{}' can not be empty",
416
0
                                           field_name);
417
0
        }
418
64
        return Status::OK();
419
64
    }
420
421
    static Status init_s3_client_conf_from_json(const rapidjson::Value& file_input,
422
6
                                                S3ClientConf& s3_client_conf) {
423
6
        std::string endpoint;
424
6
        RETURN_IF_ERROR(_get_required_string_field(file_input, "endpoint", endpoint));
425
4
        std::string region;
426
4
        RETURN_IF_ERROR(_get_required_string_field(file_input, "region", region));
427
428
8
        auto get_optional_string_field = [&](const char* field_name, std::string& value) {
429
8
            auto iter = file_input.FindMember(field_name);
430
8
            if (iter == file_input.MemberEnd() || iter->value.IsNull()) {
431
0
                return;
432
0
            }
433
8
            DORIS_CHECK(iter->value.IsString());
434
8
            value = iter->value.GetString();
435
8
        };
436
437
2
        get_optional_string_field("ak", s3_client_conf.ak);
438
2
        get_optional_string_field("sk", s3_client_conf.sk);
439
2
        get_optional_string_field("role_arn", s3_client_conf.role_arn);
440
2
        get_optional_string_field("external_id", s3_client_conf.external_id);
441
2
        s3_client_conf.endpoint = endpoint;
442
2
        s3_client_conf.region = region;
443
444
2
        return Status::OK();
445
4
    }
446
447
    Status _resolve_media_url(const rapidjson::Value& file_input, int64_t ttl_seconds,
448
28
                              std::string& media_url) const {
449
28
        std::string uri;
450
28
        RETURN_IF_ERROR(_get_required_string_field(file_input, "uri", uri));
451
452
        // If it's a direct http/https URL, use it as-is
453
28
        if (_starts_with_ignore_case(uri, "http://") || _starts_with_ignore_case(uri, "https://")) {
454
22
            media_url = uri;
455
22
            return Status::OK();
456
22
        }
457
458
6
        S3ClientConf s3_client_conf;
459
6
        RETURN_IF_ERROR(init_s3_client_conf_from_json(file_input, s3_client_conf));
460
2
        auto s3_client = DORIS_TRY(S3ClientFactory::instance().create(s3_client_conf));
461
462
2
        S3URI s3_uri(uri);
463
2
        RETURN_IF_ERROR(s3_uri.parse());
464
2
        std::string bucket = s3_uri.get_bucket();
465
2
        std::string key = s3_uri.get_key();
466
2
        DORIS_CHECK(!bucket.empty() && !key.empty());
467
2
        media_url = s3_client->generate_presigned_url({.bucket = bucket, .key = key, .prefix = ""},
468
2
                                                      ttl_seconds);
469
2
        return Status::OK();
470
2
    }
471
};
472
473
}; // namespace doris