diff --git a/be/src/exprs/function/ai/ai_adapter.h b/be/src/exprs/function/ai/ai_adapter.h index bda8e6798b5f27..eff48ef83cde78 100644 --- a/be/src/exprs/function/ai/ai_adapter.h +++ b/be/src/exprs/function/ai/ai_adapter.h @@ -20,7 +20,6 @@ #include #include -#include #include #include #include @@ -211,6 +210,27 @@ class AIAdapter { return Status::OK(); } + Status append_parsed_embedding_result(const rapidjson::Value& embedding, + std::vector>& results, + const std::string& response_body) const { + if (!embedding.IsArray()) { + return Status::InternalError("Invalid {} response format: {}", _config.provider_type, + response_body); + } + + std::vector 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; } @@ -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(); @@ -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( @@ -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); @@ -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: {}", @@ -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(); } @@ -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", @@ -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(); } @@ -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(); } @@ -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; @@ -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: diff --git a/be/test/ai/ai_adapter_test.cpp b/be/test/ai/ai_adapter_test.cpp index da40ef217dc94a..1eac053d205cf5 100644 --- a/be/test/ai/ai_adapter_test.cpp +++ b/be/test/ai/ai_adapter_test.cpp @@ -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 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 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 @@ -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 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 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"; @@ -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 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 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 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"; @@ -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 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; diff --git a/be/test/ai/embed_test.cpp b/be/test/ai/embed_test.cpp index c9bd32ed17cb7a..3ed13921d199ce 100644 --- a/be/test/ai/embed_test.cpp +++ b/be/test/ai/embed_test.cpp @@ -951,6 +951,49 @@ TEST(EMBED_TEST, local_adapter_parse_embedding_response) { ASSERT_FLOAT_EQ(results[0][1], 0.7F); } +TEST(EMBED_TEST, local_adapter_rejects_non_object_data_item) { + LocalAdapter adapter; + std::vector> results; + Status st = adapter.parse_embedding_response(R"({"data":[1]})", results); + ASSERT_FALSE(st.ok()); +} + +TEST(EMBED_TEST, local_adapter_rejects_non_numeric_data_embedding) { + LocalAdapter adapter; + std::vector> results; + Status st = + adapter.parse_embedding_response(R"({"data":[{"embedding":[0.1,"bad"]}]})", results); + ASSERT_FALSE(st.ok()); +} + +TEST(EMBED_TEST, local_adapter_rejects_non_array_embeddings_item) { + LocalAdapter adapter; + std::vector> results; + Status st = adapter.parse_embedding_response(R"({"embeddings":[0.1,0.2]})", results); + ASSERT_FALSE(st.ok()); +} + +TEST(EMBED_TEST, local_adapter_rejects_non_numeric_embeddings_item) { + LocalAdapter adapter; + std::vector> results; + Status st = adapter.parse_embedding_response(R"({"embeddings":[[0.1,"bad"]]})", results); + ASSERT_FALSE(st.ok()); +} + +TEST(EMBED_TEST, local_adapter_rejects_non_numeric_embedding) { + LocalAdapter adapter; + std::vector> results; + Status st = adapter.parse_embedding_response(R"({"embedding":[0.1,"bad"]})", results); + ASSERT_FALSE(st.ok()); +} + +TEST(EMBED_TEST, mock_adapter_rejects_non_numeric_embedding) { + MockAdapter adapter; + std::vector> results; + Status st = adapter.parse_embedding_response(R"({"embedding":[0.1,"bad"]})", results); + ASSERT_FALSE(st.ok()); +} + TEST(EMBED_TEST, openai_adapter_embedding_request) { OpenAIAdapter adapter; TAIResource config; @@ -1117,6 +1160,21 @@ TEST(EMBED_TEST, qwen_embedding_request) { ASSERT_EQ(doc["dimension"].GetInt(), config.dimensions); } +TEST(EMBED_TEST, qwen_adapter_rejects_non_object_embedding_item) { + QwenAdapter adapter; + std::vector> results; + Status st = adapter.parse_embedding_response(R"({"output":{"embeddings":[1]}})", results); + ASSERT_FALSE(st.ok()); +} + +TEST(EMBED_TEST, qwen_adapter_rejects_non_numeric_embedding) { + QwenAdapter adapter; + std::vector> results; + Status st = adapter.parse_embedding_response( + R"({"output":{"embeddings":[{"embedding":[0.1,"bad"]}]}})", results); + ASSERT_FALSE(st.ok()); +} + TEST(EMBED_TEST, gemini_adapter_embedding_request) { GeminiAdapter adapter; TAIResource config; @@ -1240,6 +1298,29 @@ TEST(EMBED_TEST, gemini_adapter_parse_embedding_response) { ASSERT_FLOAT_EQ(results[1][2], 2.3F); } +TEST(EMBED_TEST, gemini_adapter_rejects_non_object_embedding_item) { + GeminiAdapter adapter; + std::vector> results; + Status st = adapter.parse_embedding_response(R"({"embeddings":[1]})", results); + ASSERT_FALSE(st.ok()); +} + +TEST(EMBED_TEST, gemini_adapter_rejects_non_numeric_batch_embedding) { + GeminiAdapter adapter; + std::vector> results; + Status st = + adapter.parse_embedding_response(R"({"embeddings":[{"values":[0.1,"bad"]}]})", results); + ASSERT_FALSE(st.ok()); +} + +TEST(EMBED_TEST, gemini_adapter_rejects_non_numeric_single_embedding) { + GeminiAdapter adapter; + std::vector> results; + Status st = + adapter.parse_embedding_response(R"({"embedding":{"values":[0.1,"bad"]}})", results); + ASSERT_FALSE(st.ok()); +} + TEST(EMBED_TEST, voyageai_adapter_embedding_request) { VoyageAIAdapter adapter; TAIResource config; @@ -1327,6 +1408,21 @@ TEST(EMBED_TEST, voyageai_adapter_parse_embedding_response) { ASSERT_FLOAT_EQ(results[1][1], 0.5F); } +TEST(EMBED_TEST, voyageai_adapter_rejects_non_object_data_item) { + VoyageAIAdapter adapter; + std::vector> results; + Status st = adapter.parse_embedding_response(R"({"data":[1]})", results); + ASSERT_FALSE(st.ok()); +} + +TEST(EMBED_TEST, voyageai_adapter_rejects_non_numeric_embedding) { + VoyageAIAdapter adapter; + std::vector> results; + Status st = + adapter.parse_embedding_response(R"({"data":[{"embedding":[0.1,"bad"]}]})", results); + ASSERT_FALSE(st.ok()); +} + TEST(EMBED_TEST, voyageai_adapter_parse_error_test) { VoyageAIAdapter adapter;