Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
20 changes: 20 additions & 0 deletions be/src/exprs/function/ai/ai_adapter.h
Original file line number Diff line number Diff line change
Expand Up @@ -1568,6 +1568,14 @@ class AnthropicAdapter : public VoyageAIAdapter {
// Mock adapter used only for UT to bypass real HTTP calls and return deterministic data.
class MockAdapter : public AIAdapter {
public:
#ifdef BE_TEST
static void clear_embedding_inputs_for_test() { _embedding_inputs_for_test().clear(); }

static const std::vector<std::string>& 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<std::string>& inputs,
Expand All @@ -1583,6 +1591,10 @@ class MockAdapter : public AIAdapter {

Status build_embedding_request(const std::vector<std::string>& 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();
}

Expand Down Expand Up @@ -1612,6 +1624,14 @@ class MockAdapter : public AIAdapter {
[](const auto& val) { return val.GetFloat(); });
return Status::OK();
}

private:
#ifdef BE_TEST
static std::vector<std::string>& _embedding_inputs_for_test() {
static thread_local std::vector<std::string> embedding_inputs;
return embedding_inputs;
}
#endif
};

class AIAdapterFactory {
Expand Down
4 changes: 2 additions & 2 deletions be/src/exprs/function/ai/ai_classify.h
Original file line number Diff line number Diff line change
Expand Up @@ -37,13 +37,13 @@ class FunctionAIClassify : public AIFunction<FunctionAIClassify> {

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<DataTypeString>();
}

static FunctionPtr create() { return std::make_shared<FunctionAIClassify>(); }

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
4 changes: 2 additions & 2 deletions be/src/exprs/function/ai/ai_extract.h
Original file line number Diff line number Diff line change
Expand Up @@ -38,13 +38,13 @@ class FunctionAIExtract : public AIFunction<FunctionAIExtract> {

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<DataTypeString>();
}

static FunctionPtr create() { return std::make_shared<FunctionAIExtract>(); }

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;
};

Expand Down
2 changes: 1 addition & 1 deletion be/src/exprs/function/ai/ai_filter.h
Original file line number Diff line number Diff line change
Expand Up @@ -38,7 +38,7 @@ class FunctionAIFilter : public AIFunction<FunctionAIFilter> {

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<DataTypeBool>();
}

Expand Down
2 changes: 1 addition & 1 deletion be/src/exprs/function/ai/ai_fix_grammar.h
Original file line number Diff line number Diff line change
Expand Up @@ -38,7 +38,7 @@ class FunctionAIFixGrammar : public AIFunction<FunctionAIFixGrammar> {

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<DataTypeString>();
}

Expand Down
57 changes: 23 additions & 34 deletions be/src/exprs/function/ai/ai_functions.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -30,17 +30,15 @@
#include "exprs/function/simple_function_factory.h"

namespace doris {
Status FunctionAIClassify::build_prompt(const Block& block, const ColumnNumbers& arguments,
size_t row_num, std::string& prompt) const {
Status FunctionAIClassify::build_prompt(const Columns& prompt_columns, size_t row_num,
std::string& prompt) const {
// Get the text column
const ColumnWithTypeAndName& text_column = block.get_by_position(arguments[1]);
StringRef text = text_column.column->get_data_at(row_num);
StringRef text = prompt_columns[0]->get_data_at(row_num);
std::string text_str = std::string(text.data, text.size);

// Get the labels array column
const ColumnWithTypeAndName& labels_column = block.get_by_position(arguments[2]);
const auto& [array_column, array_row_num] =
check_column_const_set_readability(*labels_column.column, row_num);
check_column_const_set_readability(*prompt_columns[1], row_num);
const auto* col_array = check_and_get_column<ColumnArray>(*array_column);
if (col_array == nullptr) {
return Status::InternalError(
Expand Down Expand Up @@ -72,17 +70,15 @@ Status FunctionAIClassify::build_prompt(const Block& block, const ColumnNumbers&
return Status::OK();
}

Status FunctionAIExtract::build_prompt(const Block& block, const ColumnNumbers& arguments,
size_t row_num, std::string& prompt) const {
Status FunctionAIExtract::build_prompt(const Columns& prompt_columns, size_t row_num,
std::string& prompt) const {
// Get the text column
const ColumnWithTypeAndName& text_column = block.get_by_position(arguments[1]);
StringRef text = text_column.column->get_data_at(row_num);
StringRef text = prompt_columns[0]->get_data_at(row_num);
std::string text_str = std::string(text.data, text.size);

// Get the labels array column
const ColumnWithTypeAndName& labels_column = block.get_by_position(arguments[2]);
const auto& [array_column, array_row_num] =
check_column_const_set_readability(*labels_column.column, row_num);
check_column_const_set_readability(*prompt_columns[1], row_num);
const auto* col_array = check_and_get_column<ColumnArray>(*array_column);
if (col_array == nullptr) {
return Status::InternalError(
Expand Down Expand Up @@ -114,26 +110,23 @@ Status FunctionAIExtract::build_prompt(const Block& block, const ColumnNumbers&
return Status::OK();
}

Status FunctionAIGenerate::build_prompt(const Block& block, const ColumnNumbers& arguments,
size_t row_num, std::string& prompt) const {
const ColumnWithTypeAndName& text_column = block.get_by_position(arguments[1]);
StringRef text_ref = text_column.column->get_data_at(row_num);
Status FunctionAIGenerate::build_prompt(const Columns& prompt_columns, size_t row_num,
std::string& prompt) const {
StringRef text_ref = prompt_columns[0]->get_data_at(row_num);
prompt = std::string(text_ref.data, text_ref.size);

return Status::OK();
}

Status FunctionAIMask::build_prompt(const Block& block, const ColumnNumbers& arguments,
size_t row_num, std::string& prompt) const {
Status FunctionAIMask::build_prompt(const Columns& prompt_columns, size_t row_num,
std::string& prompt) const {
// Get the text column
const ColumnWithTypeAndName& text_column = block.get_by_position(arguments[1]);
StringRef text = text_column.column->get_data_at(row_num);
StringRef text = prompt_columns[0]->get_data_at(row_num);
std::string text_str = std::string(text.data, text.size);

// Get the labels array column
const ColumnWithTypeAndName& labels_column = block.get_by_position(arguments[2]);
const auto& [array_column, array_row_num] =
check_column_const_set_readability(*labels_column.column, row_num);
check_column_const_set_readability(*prompt_columns[1], row_num);
const auto* col_array = check_and_get_column<ColumnArray>(*array_column);
if (col_array == nullptr) {
return Status::InternalError(
Expand Down Expand Up @@ -165,33 +158,29 @@ Status FunctionAIMask::build_prompt(const Block& block, const ColumnNumbers& arg
return Status::OK();
}

Status FunctionAISimilarity::build_prompt(const Block& block, const ColumnNumbers& arguments,
size_t row_num, std::string& prompt) const {
Status FunctionAISimilarity::build_prompt(const Columns& prompt_columns, size_t row_num,
std::string& prompt) const {
// text1
const ColumnWithTypeAndName& text_column_1 = block.get_by_position(arguments[1]);
StringRef text_1 = text_column_1.column.get()->get_data_at(row_num);
StringRef text_1 = prompt_columns[0]->get_data_at(row_num);
std::string text_str_1 = std::string(text_1.data, text_1.size);

// text2
const ColumnWithTypeAndName& text_column_2 = block.get_by_position(arguments[2]);
StringRef text_2 = text_column_2.column.get()->get_data_at(row_num);
StringRef text_2 = prompt_columns[1]->get_data_at(row_num);
std::string text_str_2 = std::string(text_2.data, text_2.size);

prompt = "Text 1: " + text_str_1 + "\nText 2: " + text_str_2;

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;
Expand Down
Loading
Loading