From 88a61133ec4554554edc59aff47facf54cb1522e Mon Sep 17 00:00:00 2001 From: linrrarity Date: Mon, 21 Sep 2026 17:40:31 +0800 Subject: [PATCH] [fix](be) Prevent BE crashes on malformed AI adapter responses (#68247) ### What problem does this PR solve? Issue Number: N/A Related PR: #68247 Problem Summary: AI adapters can receive valid JSON with unexpected nested value types, such as a numeric choice, a non-array embedding, or a string inside an embedding array. Calling RapidJSON object, array, or numeric accessors on those values can trigger assertions or invalid accesses in the BE. Validate each required type before accessing it, and share embedding-array validation and conversion across provider adapters. Malformed responses return a non-OK Status that callers propagate as a query error. ### Release note Return query errors for malformed AI provider response types instead of risking a BE crash during response parsing. ### Check List (For Author) - Test: Unit Test / Manual test - Add 21 malformed-response unit tests; BE UT and regression CI passed for e004a8123a80625ac3e00892fc3cd598cda3ff03. - Isolated parser comparison with ASAN/UBSAN: all 21 malformed examples return errors after the fix, and 25 normal examples produce unchanged results. This comparison does not run the full Doris engine. - Behavior changed: Yes. Invalid response types return errors before invoking incompatible RapidJSON accessors. - Does this need documentation: No. --- be/src/exprs/function/ai/ai_adapter.h | 102 +++++++++++++++----------- be/test/ai/ai_adapter_test.cpp | 56 ++++++++++++++ be/test/ai/embed_test.cpp | 96 ++++++++++++++++++++++++ 3 files changed, 213 insertions(+), 41 deletions(-) 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;