This is an automated email from the ASF dual-hosted git repository.

HappenLee pushed a commit to branch master
in repository https://gitbox.apache.org/repos/asf/doris.git


The following commit(s) were added to refs/heads/master by this push:
     new 2624e053a37 [fix](be) Prevent BE crashes on malformed AI adapter 
responses (#68247)
2624e053a37 is described below

commit 2624e053a37104bbdd7937d1fb807ab0e53fc461
Author: linrrarity <[email protected]>
AuthorDate: Mon Sep 21 17:40:31 2026 +0800

    [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 propag [...]
    
    ### 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 b6376c6c004..79bdd18b6c1 100644
--- a/be/src/exprs/function/ai/ai_adapter.h
+++ b/be/src/exprs/function/ai/ai_adapter.h
@@ -20,7 +20,6 @@
 #include <gen_cpp/PaloInternalService_types.h>
 #include <rapidjson/rapidjson.h>
 
-#include <algorithm>
 #include <cctype>
 #include <memory>
 #include <string>
@@ -210,6 +209,27 @@ protected:
         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; }
 
@@ -407,14 +427,12 @@ public:
         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();
@@ -481,6 +499,14 @@ public:
             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(
@@ -560,37 +586,31 @@ public:
         }
 
         // 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);
@@ -945,7 +965,8 @@ public:
             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: {}",
@@ -1129,14 +1150,12 @@ public:
             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();
         }
@@ -1323,10 +1342,12 @@ public:
         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",
@@ -1498,13 +1519,12 @@ public:
             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();
         }
@@ -1519,13 +1539,12 @@ public:
           }
         }*/
         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();
     }
@@ -1625,6 +1644,10 @@ public:
 
         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;
@@ -1696,10 +1719,7 @@ public:
         }
 
         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 da40ef217dc..1eac053d205 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<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
@@ -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";
@@ -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";
@@ -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;
diff --git a/be/test/ai/embed_test.cpp b/be/test/ai/embed_test.cpp
index cfb4c521890..8f7206fb24f 100644
--- a/be/test/ai/embed_test.cpp
+++ b/be/test/ai/embed_test.cpp
@@ -946,6 +946,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<std::vector<float>> 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<std::vector<float>> 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<std::vector<float>> 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<std::vector<float>> 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<std::vector<float>> 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<std::vector<float>> 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;
@@ -1112,6 +1155,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<std::vector<float>> 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<std::vector<float>> 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;
@@ -1235,6 +1293,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<std::vector<float>> 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<std::vector<float>> 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<std::vector<float>> 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;
@@ -1322,6 +1403,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<std::vector<float>> 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<std::vector<float>> 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;
 


---------------------------------------------------------------------
To unsubscribe, e-mail: [email protected]
For additional commands, e-mail: [email protected]

Reply via email to