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
102 changes: 61 additions & 41 deletions be/src/exprs/function/ai/ai_adapter.h
Original file line number Diff line number Diff line change
Expand Up @@ -20,7 +20,6 @@
#include <gen_cpp/PaloInternalService_types.h>
#include <rapidjson/rapidjson.h>

#include <algorithm>
#include <cctype>
#include <memory>
#include <string>
Expand Down Expand Up @@ -211,6 +210,27 @@ class AIAdapter {
return Status::OK();
}

Status append_parsed_embedding_result(const rapidjson::Value& embedding,
std::vector<std::vector<float>>& results,
const std::string& response_body) const {
if (!embedding.IsArray()) {
return Status::InternalError("Invalid {} response format: {}", _config.provider_type,
response_body);
}

std::vector<float> parsed_embedding;
parsed_embedding.reserve(embedding.Size());
for (const auto& value : embedding.GetArray()) {
if (!value.IsNumber()) {
return Status::InternalError("Invalid {} response format: {}",
_config.provider_type, response_body);
}
parsed_embedding.emplace_back(value.GetFloat());
}
results.emplace_back(std::move(parsed_embedding));
return Status::OK();
}

// return true if the model support dimension parameter
virtual bool supports_dimension_param(const std::string& model_name) const { return false; }

Expand Down Expand Up @@ -408,14 +428,12 @@ class VoyageAIAdapter : public AIAdapter {
const auto& data = doc["data"];
results.reserve(data.Size());
for (rapidjson::SizeType i = 0; i < data.Size(); i++) {
if (!data[i].HasMember("embedding") || !data[i]["embedding"].IsArray()) {
if (!data[i].IsObject() || !data[i].HasMember("embedding")) {
return Status::InternalError("Invalid {} response format: {}",
_config.provider_type, response_body);
}

std::transform(data[i]["embedding"].Begin(), data[i]["embedding"].End(),
std::back_inserter(results.emplace_back()),
[](const auto& val) { return val.GetFloat(); });
RETURN_IF_ERROR(
append_parsed_embedding_result(data[i]["embedding"], results, response_body));
}

return Status::OK();
Expand Down Expand Up @@ -482,6 +500,14 @@ class LocalAdapter : public AIAdapter {
results.reserve(choices.Size());

for (rapidjson::SizeType i = 0; i < choices.Size(); i++) {
if (!choices[i].IsObject()) {
return Status::InternalError("Invalid {} response format: {}",
_config.provider_type, response_body);
}
if (choices[i].HasMember("message") && !choices[i]["message"].IsObject()) {
return Status::InternalError("Invalid {} response format: {}",
_config.provider_type, response_body);
}
if (choices[i].HasMember("message") && choices[i]["message"].HasMember("content") &&
choices[i]["message"]["content"].IsString()) {
RETURN_IF_ERROR(append_parsed_text_result(
Expand Down Expand Up @@ -561,37 +587,31 @@ class LocalAdapter : public AIAdapter {
}

// parse different response format
rapidjson::Value embedding;
if (doc.HasMember("data") && doc["data"].IsArray()) {
// "data":["object":"embedding", "embedding":[0.1, 0.2...], "index":0]
const auto& data = doc["data"];
results.reserve(data.Size());
for (rapidjson::SizeType i = 0; i < data.Size(); i++) {
if (!data[i].HasMember("embedding") || !data[i]["embedding"].IsArray()) {
if (!data[i].IsObject() || !data[i].HasMember("embedding")) {
return Status::InternalError("Invalid {} response format",
_config.provider_type);
}

std::transform(data[i]["embedding"].Begin(), data[i]["embedding"].End(),
std::back_inserter(results.emplace_back()),
[](const auto& val) { return val.GetFloat(); });
RETURN_IF_ERROR(append_parsed_embedding_result(data[i]["embedding"], results,
response_body));
}
} else if (doc.HasMember("embeddings") && doc["embeddings"].IsArray()) {
// "embeddings":[[0.1, 0.2, ...]]
results.reserve(1);
for (int i = 0; i < doc["embeddings"].Size(); i++) {
embedding = doc["embeddings"][i];
std::transform(embedding.Begin(), embedding.End(),
std::back_inserter(results.emplace_back()),
[](const auto& val) { return val.GetFloat(); });
const auto& embeddings = doc["embeddings"];
results.reserve(embeddings.Size());
for (rapidjson::SizeType i = 0; i < embeddings.Size(); i++) {
RETURN_IF_ERROR(
append_parsed_embedding_result(embeddings[i], results, response_body));
}
} else if (doc.HasMember("embedding") && doc["embedding"].IsArray()) {
// "embedding":[0.1, 0.2, ...]
results.reserve(1);
embedding = doc["embedding"];
std::transform(embedding.Begin(), embedding.End(),
std::back_inserter(results.emplace_back()),
[](const auto& val) { return val.GetFloat(); });
RETURN_IF_ERROR(
append_parsed_embedding_result(doc["embedding"], results, response_body));
} else {
return Status::InternalError("Invalid {} response format: {}", _config.provider_type,
response_body);
Expand Down Expand Up @@ -946,7 +966,8 @@ class OpenAIAdapter : public VoyageAIAdapter {
results.reserve(choices.Size());

for (rapidjson::SizeType i = 0; i < choices.Size(); i++) {
if (!choices[i].HasMember("message") ||
if (!choices[i].IsObject() || !choices[i].HasMember("message") ||
!choices[i]["message"].IsObject() ||
!choices[i]["message"].HasMember("content") ||
!choices[i]["message"]["content"].IsString()) {
return Status::InternalError("Invalid choice format in {} response: {}",
Expand Down Expand Up @@ -1130,14 +1151,12 @@ class QwenAdapter : public OpenAIAdapter {
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()) {
if (!embeddings[i].IsObject() || !embeddings[i].HasMember("embedding")) {
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_IF_ERROR(append_parsed_embedding_result(embeddings[i]["embedding"], results,
response_body));
}
return Status::OK();
}
Expand Down Expand Up @@ -1324,10 +1343,12 @@ class GeminiAdapter : public AIAdapter {
results.reserve(candidates.Size());

for (rapidjson::SizeType i = 0; i < candidates.Size(); i++) {
if (!candidates[i].HasMember("content") ||
if (!candidates[i].IsObject() || !candidates[i].HasMember("content") ||
!candidates[i]["content"].IsObject() ||
!candidates[i]["content"].HasMember("parts") ||
!candidates[i]["content"]["parts"].IsArray() ||
candidates[i]["content"]["parts"].Empty() ||
!candidates[i]["content"]["parts"][0].IsObject() ||
!candidates[i]["content"]["parts"][0].HasMember("text") ||
!candidates[i]["content"]["parts"][0]["text"].IsString()) {
return Status::InternalError("Invalid candidate format in {} response",
Expand Down Expand Up @@ -1499,13 +1520,12 @@ class GeminiAdapter : public AIAdapter {
const auto& embeddings = doc["embeddings"];
results.reserve(embeddings.Size());
for (rapidjson::SizeType i = 0; i < embeddings.Size(); i++) {
if (!embeddings[i].HasMember("values") || !embeddings[i]["values"].IsArray()) {
if (!embeddings[i].IsObject() || !embeddings[i].HasMember("values")) {
return Status::InternalError("Invalid {} response format: {}",
_config.provider_type, response_body);
}
std::transform(embeddings[i]["values"].Begin(), embeddings[i]["values"].End(),
std::back_inserter(results.emplace_back()),
[](const auto& val) { return val.GetFloat(); });
RETURN_IF_ERROR(append_parsed_embedding_result(embeddings[i]["values"], results,
response_body));
}
return Status::OK();
}
Expand All @@ -1520,13 +1540,12 @@ class GeminiAdapter : public AIAdapter {
}
}*/
const auto& embedding = doc["embedding"];
if (!embedding.HasMember("values") || !embedding["values"].IsArray()) {
if (!embedding.HasMember("values")) {
return Status::InternalError("Invalid {} response format: {}", _config.provider_type,
response_body);
}
std::transform(embedding["values"].Begin(), embedding["values"].End(),
std::back_inserter(results.emplace_back()),
[](const auto& val) { return val.GetFloat(); });
RETURN_IF_ERROR(
append_parsed_embedding_result(embedding["values"], results, response_body));

return Status::OK();
}
Expand Down Expand Up @@ -1626,6 +1645,10 @@ class AnthropicAdapter : public VoyageAIAdapter {

std::string result;
for (rapidjson::SizeType i = 0; i < content.Size(); i++) {
if (!content[i].IsObject()) {
return Status::InternalError("Invalid {} response format: {}",
_config.provider_type, response_body);
}
if (!content[i].HasMember("type") || !content[i]["type"].IsString() ||
!content[i].HasMember("text") || !content[i]["text"].IsString()) {
continue;
Expand Down Expand Up @@ -1697,10 +1720,7 @@ class MockAdapter : public AIAdapter {
}

results.reserve(1);
std::transform(doc["embedding"].Begin(), doc["embedding"].End(),
std::back_inserter(results.emplace_back()),
[](const auto& val) { return val.GetFloat(); });
return Status::OK();
return append_parsed_embedding_result(doc["embedding"], results, response_body);
}

private:
Expand Down
56 changes: 56 additions & 0 deletions be/test/ai/ai_adapter_test.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -863,6 +863,20 @@ TEST(AI_ADAPTER_TEST, parse_response_wrong_type) {
::testing::HasSubstr("Unsupported response format from local AI."));
}

TEST(AI_ADAPTER_TEST, local_adapter_rejects_non_object_choice) {
LocalAdapter adapter;
std::vector<std::string> results;
Status st = adapter.parse_response(R"({"choices":[1]})", results);
ASSERT_FALSE(st.ok());
}

TEST(AI_ADAPTER_TEST, local_adapter_rejects_non_object_message) {
LocalAdapter adapter;
std::vector<std::string> results;
Status st = adapter.parse_response(R"({"choices":[{"message":1}]})", results);
ASSERT_FALSE(st.ok());
}

TEST(AI_ADAPTER_TEST, openai_adapter_parse_response_choice_format_error) {
OpenAIAdapter adapter;
// message field missing
Expand All @@ -880,6 +894,20 @@ TEST(AI_ADAPTER_TEST, openai_adapter_parse_response_choice_format_error) {
EXPECT_THAT(st.to_string().c_str(), ::testing::HasSubstr("Invalid choice format in response"));
}

TEST(AI_ADAPTER_TEST, openai_adapter_rejects_non_object_choice) {
OpenAIAdapter adapter;
std::vector<std::string> results;
Status st = adapter.parse_response(R"({"choices":[1]})", results);
ASSERT_FALSE(st.ok());
}

TEST(AI_ADAPTER_TEST, openai_adapter_rejects_non_object_message) {
OpenAIAdapter adapter;
std::vector<std::string> results;
Status st = adapter.parse_response(R"({"choices":[{"message":1}]})", results);
ASSERT_FALSE(st.ok());
}

TEST(AI_ADAPTER_TEST, openai_adapter_parse_response_parse_error) {
OpenAIAdapter adapter;
std::string resp = "not a json";
Expand Down Expand Up @@ -916,6 +944,27 @@ TEST(AI_ADAPTER_TEST, gemini_parse_response_missing_candidates) {
EXPECT_THAT(st.to_string().c_str(), ::testing::HasSubstr("Invalid response format"));
}

TEST(AI_ADAPTER_TEST, gemini_adapter_rejects_non_object_candidate) {
GeminiAdapter adapter;
std::vector<std::string> results;
Status st = adapter.parse_response(R"({"candidates":[1]})", results);
ASSERT_FALSE(st.ok());
}

TEST(AI_ADAPTER_TEST, gemini_adapter_rejects_non_object_content) {
GeminiAdapter adapter;
std::vector<std::string> results;
Status st = adapter.parse_response(R"({"candidates":[{"content":1}]})", results);
ASSERT_FALSE(st.ok());
}

TEST(AI_ADAPTER_TEST, gemini_adapter_rejects_non_object_part) {
GeminiAdapter adapter;
std::vector<std::string> results;
Status st = adapter.parse_response(R"({"candidates":[{"content":{"parts":[1]}}]})", results);
ASSERT_FALSE(st.ok());
}

TEST(AI_ADAPTER_TEST, anthropic_adapter_parse_response_parse_error) {
AnthropicAdapter adapter;
std::string resp = "not a json";
Expand All @@ -934,6 +983,13 @@ TEST(AI_ADAPTER_TEST, anthropic_adapter_parse_response_content_not_array) {
EXPECT_THAT(st.to_string().c_str(), ::testing::HasSubstr("Invalid response format"));
}

TEST(AI_ADAPTER_TEST, anthropic_adapter_rejects_non_object_content_item) {
AnthropicAdapter adapter;
std::vector<std::string> results;
Status st = adapter.parse_response(R"({"content":[1]})", results);
ASSERT_FALSE(st.ok());
}

TEST(AI_ADAPTER_TEST, voyage_adapter_chat_test) {
VoyageAIAdapter adapter;
TAIResource config;
Expand Down
Loading
Loading