From 6dc84c51bbc84982d5b94c07ee6030a7acd95b8c Mon Sep 17 00:00:00 2001 From: Spade A <71589810+SpadeA-Tang@users.noreply.github.com> Date: Wed, 24 Jun 2026 14:44:26 +0800 Subject: [PATCH] enhance: support range search and group-by for element-level hybrid search (#50383) issue: https://github.com/milvus-io/milvus/issues/42148 1. Support range search in hybrid search Enable element-level vector fields to participate in hybrid range search while preserving element identity and offsets. 2. Support group-by in hybrid search, with constraints Enable element-level hybrid search with PK group-by and preserve element offsets. For group-by, sub-searches that need element identity merging must stay within the same struct array semantic scope. Also, it fixes 2 related bugs: 1. Fix group-by related issues Fix lost element indices and incorrect merging behavior in query/search group-by paths, so multiple matching elements under the same PK can still be returned with correct offsets. 2. Fix aggregation output issues Fix aggregation/count output corruption when a collection contains struct array fields. Proxy struct-field reconstruction now preserves virtual aggregation fields such as `count(*)` instead of treating `field_id=0` as a struct sub-field. As a result, `id + count(*)`, element-level `count(*)`, and group-by `count(*)` return correctly named fields. --------- Signed-off-by: SpadeA --- internal/core/src/common/Consts.h | 3 + .../core/src/exec/operator/ProjectNode.cpp | 123 +++++++++-- internal/core/src/exec/operator/ProjectNode.h | 3 + .../src/exec/operator/SearchGroupByNode.cpp | 20 +- .../search-groupby/SearchGroupByOperator.cpp | 51 +++-- .../search-groupby/SearchGroupByOperator.h | 3 +- internal/core/src/query/PlanProto.cpp | 19 +- .../core/src/segcore/SegmentInterface.cpp | 59 ++++- .../core/unittest/test_element_filter.cpp | 25 ++- .../core/unittest/test_query_group_by.cpp | 165 ++++++++++++++ .../core/unittest/test_query_order_by.cpp | 152 +++++++++++++ .../core/unittest/test_search_group_by.cpp | 203 +++++++++++++++++- internal/proxy/query_pipeline_test.go | 80 +++++++ internal/proxy/search_pipeline.go | 6 + internal/proxy/search_pipeline_test.go | 20 ++ internal/proxy/task_search.go | 57 ++++- internal/proxy/task_search_test.go | 125 ++++++++--- internal/proxy/util.go | 6 +- internal/proxy/util_test.go | 65 ++++++ 19 files changed, 1089 insertions(+), 96 deletions(-) diff --git a/internal/core/src/common/Consts.h b/internal/core/src/common/Consts.h index a082ee49b6..0012ea5fe3 100644 --- a/internal/core/src/common/Consts.h +++ b/internal/core/src/common/Consts.h @@ -36,6 +36,9 @@ const milvus::FieldId TimestampFieldID = milvus::FieldId(1); // Virtual field ID for two-project mode: carries segment offsets through // the pipeline so that deferred fields can be fetched after TopK. const milvus::FieldId SegmentOffsetFieldID = milvus::FieldId(-100); +// Virtual field ID for element-level ORDER BY: carries the matched element +// index alongside SegmentOffsetFieldID after TopK. +const milvus::FieldId ElementIndexFieldID = milvus::FieldId(-101); // fill followed extra info to binlog file const char ORIGIN_SIZE_KEY[] = "original_size"; diff --git a/internal/core/src/exec/operator/ProjectNode.cpp b/internal/core/src/exec/operator/ProjectNode.cpp index 4a8d1b5530..95404c5f04 100644 --- a/internal/core/src/exec/operator/ProjectNode.cpp +++ b/internal/core/src/exec/operator/ProjectNode.cpp @@ -17,6 +17,7 @@ #include "ProjectNode.h" #include +#include #include #include "common/Consts.h" @@ -27,11 +28,86 @@ #include "exec/expression/Utils.h" #include "exec/operator/Operator.h" #include "plan/PlanNode.h" +#include "segcore/InsertRecord.h" #include "segcore/SegmentInterface.h" #include "segcore/Utils.h" namespace milvus { namespace exec { +namespace { + +struct SelectedOffsets { + std::vector row_offsets; + std::vector element_indices; +}; + +SelectedOffsets +SelectOffsets(const TargetBitmapView& raw_data_view, + QueryContext* query_context, + const segcore::SegmentInternalInterface* segment) { + if (!query_context->bitset_is_element_level()) { + auto result_pair = + segment->find_first_n(segcore::Unlimited, raw_data_view); + return {std::move(result_pair.first), {}}; + } + + auto array_offsets = query_context->get_array_offsets(); + AssertInfo(array_offsets != nullptr, + "element-level ProjectNode requires array offsets"); + + auto [doc_offsets, element_indices_by_doc, _] = + segment->find_first_n_element(segcore::Unlimited, + raw_data_view, + array_offsets.get(), + std::nullopt); + + size_t selected_count = 0; + for (const auto& element_indices : element_indices_by_doc) { + selected_count += element_indices.size(); + } + + SelectedOffsets selected; + selected.row_offsets.reserve(selected_count); + selected.element_indices.reserve(selected_count); + + for (size_t i = 0; i < doc_offsets.size(); i++) { + for (auto element_index : element_indices_by_doc[i]) { + selected.row_offsets.emplace_back(doc_offsets[i]); + selected.element_indices.emplace_back(element_index); + } + } + return selected; +} + +ColumnVectorPtr +MakeInt64Column(const std::vector& values) { + auto selected_count = values.size(); + FixedVector output(selected_count); + milvus::fastmem::FastMemcpy( + output.data(), values.data(), selected_count * sizeof(int64_t)); + auto field_data = std::make_shared>( + DataType::INT64, false, std::move(output)); + auto valid_map = TargetBitmap(selected_count, true); + return std::make_shared( + std::move(field_data), std::move(valid_map), 0); +} + +ColumnVectorPtr +MakeElementIndexColumn(const std::vector& element_indices) { + auto selected_count = element_indices.size(); + FixedVector output(selected_count); + for (size_t i = 0; i < selected_count; i++) { + output[i] = element_indices[i]; + } + auto field_data = std::make_shared>( + DataType::INT64, false, std::move(output)); + auto valid_map = TargetBitmap(selected_count, true); + return std::make_shared( + std::move(field_data), std::move(valid_map), 0); +} + +} // namespace + PhyProjectNode::PhyProjectNode( int32_t operator_id, milvus::exec::DriverContext* ctx, @@ -41,10 +117,13 @@ PhyProjectNode::PhyProjectNode( operator_id, projectNode->id(), "Project"), - fields_to_project_(projectNode->FieldsToProject()) { + fields_to_project_(projectNode->FieldsToProject()), + query_context_(nullptr), + op_context_(nullptr) { auto exec_context = operator_context_->get_exec_context(); - segment_ = exec_context->get_query_context()->get_segment(); - op_context_ = exec_context->get_query_context()->get_op_context(); + query_context_ = exec_context->get_query_context(); + segment_ = query_context_->get_segment(); + op_context_ = query_context_->get_op_context(); AssertInfo(op_context_, "op_context_ cannot be nullptr for ProjectNode"); AssertInfo(segment_, "segment_ cannot be nullptr for ProjectNode"); } @@ -63,10 +142,11 @@ PhyProjectNode::GetOutput() { // raw data view TargetBitmapView raw_data_view(col_input->GetRawData(), col_input->size()); - // When no fields need to be projected (e.g., count(*) only), skip - // find_first and count valid rows directly from the bitmap. - // find_first deduplicates by PK in growing segments (OffsetOrderedMap), - // which would undercount rows with duplicate PKs. + // When no fields need to be projected (e.g., count(*) only), count valid + // logical rows directly from the bitmap. For element-level bitmaps this is + // the matching element count. find_first deduplicates by PK in growing + // segments (OffsetOrderedMap), which would undercount rows with duplicate + // PKs. if (fields_to_project_.empty()) { auto valid_count = static_cast(col_input->size()) - raw_data_view.count(); @@ -79,8 +159,9 @@ PhyProjectNode::GetOutput() { return row_vector; } - auto result_pair = segment_->find_first_n(-1, raw_data_view); - auto& selected_offsets = result_pair.first; + auto selected = SelectOffsets(raw_data_view, query_context_, segment_); + auto& selected_offsets = selected.row_offsets; + auto& selected_element_indices = selected.element_indices; auto selected_count = selected_offsets.size(); // When all rows are filtered out, return nullptr. // Driver requires GetOutput to return nullptr or a non-empty vector. @@ -96,16 +177,18 @@ PhyProjectNode::GetOutput() { auto field_id = fields_to_project_.at(i); if (field_id == SegmentOffsetFieldID) { - FixedVector offsets(selected_count); - milvus::fastmem::FastMemcpy(offsets.data(), - selected_offsets.data(), - selected_count * sizeof(int64_t)); - auto field_data = std::make_shared>( - DataType::INT64, false, std::move(offsets)); - auto valid_map = TargetBitmap(selected_count, true); - auto offset_col = std::make_shared( - std::move(field_data), std::move(valid_map), 0); - column_vectors.emplace_back(std::move(offset_col)); + column_vectors.emplace_back(MakeInt64Column(selected_offsets)); + continue; + } + + if (field_id == ElementIndexFieldID) { + AssertInfo(selected_element_indices.size() == selected_count, + "element indices size ({}) must match selected row " + "count ({})", + selected_element_indices.size(), + selected_count); + column_vectors.emplace_back( + MakeElementIndexColumn(selected_element_indices)); continue; } @@ -139,4 +222,4 @@ PhyProjectNode::GetOutput() { } }; // namespace exec -}; // namespace milvus \ No newline at end of file +}; // namespace milvus diff --git a/internal/core/src/exec/operator/ProjectNode.h b/internal/core/src/exec/operator/ProjectNode.h index a2280ede02..b4b42135d3 100644 --- a/internal/core/src/exec/operator/ProjectNode.h +++ b/internal/core/src/exec/operator/ProjectNode.h @@ -20,6 +20,8 @@ namespace milvus { namespace exec { +class QueryContext; + class PhyProjectNode : public Operator { public: PhyProjectNode(int32_t operator_id, @@ -70,6 +72,7 @@ class PhyProjectNode : public Operator { const segcore::SegmentInternalInterface* segment_; bool is_finished_{false}; const std::vector fields_to_project_; + QueryContext* query_context_; OpContext* op_context_; }; } // namespace exec diff --git a/internal/core/src/exec/operator/SearchGroupByNode.cpp b/internal/core/src/exec/operator/SearchGroupByNode.cpp index a20f3937dd..564a28453e 100644 --- a/internal/core/src/exec/operator/SearchGroupByNode.cpp +++ b/internal/core/src/exec/operator/SearchGroupByNode.cpp @@ -100,7 +100,10 @@ PhySearchGroupByNode::GetOutput() { *segment_, search_result.seg_offsets_, search_result.distances_, - search_result.topk_per_nq_prefix_sum_); + search_result.topk_per_nq_prefix_sum_, + search_result.element_level_ + ? &search_result.element_indices_ + : nullptr); search_result.composite_group_by_values_ = std::move(composite_group_by_values); search_result.group_size_ = search_info_.group_size_; @@ -110,15 +113,18 @@ PhySearchGroupByNode::GetOutput() { "size:{} is not equal to search_result.seg_offsets.size:{}", search_result.composite_group_by_values_.value().size(), search_result.seg_offsets_.size()); + if (search_result.element_level_) { + AssertInfo(search_result.seg_offsets_.size() == + search_result.element_indices_.size(), + "Wrong state! search_result element_indices size:{} is " + "not equal to search_result.seg_offsets size:{}", + search_result.element_indices_.size(), + search_result.seg_offsets_.size()); + } } tracer::AddEvent( fmt::format("grouped_results: {}", search_result.seg_offsets_.size())); - // When group by is used with element-level search, the result is row-level - // because group by operates on rows, not elements. - search_result.element_level_ = false; - search_result.element_indices_.clear(); - query_context_->set_search_result(std::move(search_result)); std::chrono::high_resolution_clock::time_point vector_end = std::chrono::high_resolution_clock::now(); @@ -136,4 +142,4 @@ PhySearchGroupByNode::IsFinished() { } } // namespace exec -} // namespace milvus \ No newline at end of file +} // namespace milvus diff --git a/internal/core/src/exec/operator/search-groupby/SearchGroupByOperator.cpp b/internal/core/src/exec/operator/search-groupby/SearchGroupByOperator.cpp index b20d398d29..0a17db9fb4 100644 --- a/internal/core/src/exec/operator/search-groupby/SearchGroupByOperator.cpp +++ b/internal/core/src/exec/operator/search-groupby/SearchGroupByOperator.cpp @@ -30,6 +30,17 @@ namespace milvus { namespace exec { +namespace { + +struct GroupedResult { + int64_t row_offset; + int32_t element_index; + float distance; + CompositeGroupKey group_key; +}; + +} // namespace + // Helper to create a single-field getter that returns GroupByValueType template static std::function @@ -189,18 +200,21 @@ GroupIteratorResult(const std::shared_ptr& iterator, std::vector& composite_group_by_values, std::vector& offsets, std::vector& distances, - const SearchInfo& search_info) { + const SearchInfo& search_info, + std::vector* element_indices) { // 1. Create group map for composite keys CompositeGroupByMap groupMap(search_info.topk_, search_info.group_size_, search_info.strict_group_size_); auto is_element_id = search_info.element_level(); + AssertInfo(element_indices == nullptr || is_element_id, + "element_indices output requires element-level search"); //2. do iteration until fill the whole map or run out of all data //note it may enumerate all data inside a segment and can block following //query and search possibly - std::vector> res; + std::vector res; CompositeGroupKey scratch_key; while (iterator->HasNext() && !groupMap.IsGroupResEnough()) { auto offset_dis_pair = iterator->Next(); @@ -212,35 +226,41 @@ GroupIteratorResult(const std::shared_ptr& iterator, // For element-level search, the offset is the element_id, we need to convert it to the row_id. int64_t row_offset = raw_offset; + int32_t element_index = -1; if (is_element_id) { AssertInfo(search_info.array_offsets_ != nullptr, "Array offsets not available for element-level search"); - row_offset = - search_info.array_offsets_ - ->ElementIDToRowID(static_cast(raw_offset)) - .first; + auto [doc_id, elem_idx] = + search_info.array_offsets_->ElementIDToRowID( + static_cast(raw_offset)); + row_offset = doc_id; + element_index = elem_idx; } data_getter->GetInto(row_offset, scratch_key); if (groupMap.Push(scratch_key)) { // Safe to move: next iteration's GetInto() will Clear+Reserve+Add // on the moved-from small_vector, which is guaranteed empty-inline. - res.emplace_back(row_offset, dis, std::move(scratch_key)); + res.emplace_back(GroupedResult{ + row_offset, element_index, dis, std::move(scratch_key)}); } } // 3. Sort based on distances and metrics auto customComparator = [&](const auto& lhs, const auto& rhs) { return milvus::query::dis_closer( - std::get<1>(lhs), std::get<1>(rhs), search_info.metric_type_); + lhs.distance, rhs.distance, search_info.metric_type_); }; std::sort(res.begin(), res.end(), customComparator); // 4. Save results for (auto iter = res.begin(); iter != res.end(); ++iter) { - offsets.emplace_back(std::get<0>(*iter)); - distances.emplace_back(std::get<1>(*iter)); - composite_group_by_values.emplace_back(std::move(std::get<2>(*iter))); + offsets.emplace_back(iter->row_offset); + distances.emplace_back(iter->distance); + if (element_indices != nullptr) { + element_indices->emplace_back(iter->element_index); + } + composite_group_by_values.emplace_back(std::move(iter->group_key)); } } @@ -252,7 +272,8 @@ SearchGroupBy(milvus::OpContext* op_ctx, const segcore::SegmentInternalInterface& segment, std::vector& seg_offsets, std::vector& distances, - std::vector& topk_per_nq_prefix_sum) { + std::vector& topk_per_nq_prefix_sum, + std::vector* element_indices) { // Get field IDs for group by AssertInfo(!search_info.group_by_field_ids_.empty(), "group_by_field_ids must be set for group by search"); @@ -263,6 +284,9 @@ SearchGroupBy(milvus::OpContext* op_ctx, seg_offsets.reserve(max_total_size); distances.reserve(max_total_size); composite_group_by_values.reserve(max_total_size); + if (element_indices != nullptr) { + element_indices->reserve(max_total_size); + } topk_per_nq_prefix_sum.reserve(iterators.size() + 1); // Create data getter for all fields @@ -281,7 +305,8 @@ SearchGroupBy(milvus::OpContext* op_ctx, composite_group_by_values, seg_offsets, distances, - search_info); + search_info, + element_indices); topk_per_nq_prefix_sum.push_back(seg_offsets.size()); } } diff --git a/internal/core/src/exec/operator/search-groupby/SearchGroupByOperator.h b/internal/core/src/exec/operator/search-groupby/SearchGroupByOperator.h index 781d6c8033..17cd0ef2f2 100644 --- a/internal/core/src/exec/operator/search-groupby/SearchGroupByOperator.h +++ b/internal/core/src/exec/operator/search-groupby/SearchGroupByOperator.h @@ -432,7 +432,8 @@ SearchGroupBy(milvus::OpContext* op_ctx, const segcore::SegmentInternalInterface& segment, std::vector& seg_offsets, std::vector& distances, - std::vector& topk_per_nq_prefix_sum); + std::vector& topk_per_nq_prefix_sum, + std::vector* element_indices = nullptr); } // namespace exec } // namespace milvus diff --git a/internal/core/src/query/PlanProto.cpp b/internal/core/src/query/PlanProto.cpp index 6179cde134..fff4689b51 100644 --- a/internal/core/src/query/PlanProto.cpp +++ b/internal/core/src/query/PlanProto.cpp @@ -390,7 +390,8 @@ std::tuple, std::vector> BuildOrderByProjectNode(const proto::plan::QueryPlanNode& query, const planpb::PlanNode& plan_node_proto, const SchemaPtr& schema, - const std::vector& sources) { + const std::vector& sources, + bool is_element_level) { std::vector project_ids; std::vector project_names; std::vector project_types; @@ -452,6 +453,15 @@ BuildOrderByProjectNode(const proto::plan::QueryPlanNode& query, } } + // Element-level ORDER BY sorts logical element rows. Keep the matched + // element index in a hidden column so result assembly can restore SDK + // offsets after TopK reorders rows. + if (is_element_level) { + project_ids.push_back(ElementIndexFieldID); + project_names.push_back("ElementIndex"); + project_types.push_back(DataType::INT64); + } + // Always append SegmentOffsetFieldID as the last pipeline column. // FillOrderByResult uses these offsets to populate system fields // (e.g., TimestampField for QN-side pk+ts dedup) and, in two-project @@ -824,8 +834,11 @@ ProtoParser::RetrievePlanNodeFromProto( (group_by_field_count > 0 || agg_functions_count > 0); if (!has_aggregation) { auto [project, deferred, pipeline_ids] = - BuildOrderByProjectNode( - query, plan_node_proto, schema, sources); + BuildOrderByProjectNode(query, + plan_node_proto, + schema, + sources, + is_element_level); plannode = project; sources = std::vector{plannode}; plan_node->deferred_field_ids_ = std::move(deferred); diff --git a/internal/core/src/segcore/SegmentInterface.cpp b/internal/core/src/segcore/SegmentInterface.cpp index 43a868133a..5e7d58ce2c 100644 --- a/internal/core/src/segcore/SegmentInterface.cpp +++ b/internal/core/src/segcore/SegmentInterface.cpp @@ -15,6 +15,7 @@ #include #include #include +#include #include #include #include @@ -267,10 +268,11 @@ SegmentInternalInterface::Retrieve(tracer::TraceContext* trace_ctx, &fte_op_ctx); } else if (!plan->plan_node_->pipeline_field_ids_.empty()) { // Non-aggregation ORDER BY (single-project or two-project mode): - // Pipeline output contains [pk, sort_cols, ..., SegmentOffsetFieldID]. - // FillOrderByResult strips the offset column, sets field_id on each + // Pipeline output contains [pk, sort_cols, ..., hidden columns]. + // FillOrderByResult strips hidden columns, sets field_id on each // DataArray, bulk-fetches deferred fields (if any), populates system - // fields, and fills PK-based IDs for proxy reduce. + // fields, restores element indices when present, and fills PK-based + // IDs for proxy reduce. // // Aggregation + ORDER BY does NOT set pipeline_field_ids_ and produces // final columns directly, falling through to FillTargetEntryDirectly. @@ -311,7 +313,10 @@ SegmentInternalInterface::FillOrderByResult( auto fields_data = results->mutable_fields_data(); auto& deferred = plan->plan_node_->deferred_field_ids_; - // Pipeline layout: [...user_columns..., SegmentOffsetFieldID]. + // Pipeline layout: + // row-level: [...user_columns..., SegmentOffsetFieldID] + // element-level: [...user_columns..., ElementIndexFieldID, + // SegmentOffsetFieldID] // The last column is always SegmentOffsetFieldID carrying segment offsets. auto total_cols = retrieveResult.field_data_.size(); AssertInfo(total_cols >= 2, @@ -319,7 +324,7 @@ SegmentInternalInterface::FillOrderByResult( "(pk + SegmentOffsetFieldID), got: {}", total_cols); - // Move all columns except the last (SegmentOffsetFieldID) to results. + // Move all non-hidden columns to results. // Set field_id on each DataArray so QN-side AppendFieldData can match // fields correctly (pipeline-produced DataArrays have field_id=0 by default). auto& pipeline_ids = plan->plan_node_->pipeline_field_ids_; @@ -328,14 +333,30 @@ SegmentInternalInterface::FillOrderByResult( "count ({})", pipeline_ids.size(), total_cols); - for (size_t i = 0; i + 1 < total_cols; i++) { + + size_t segment_offset_col_idx = total_cols; + size_t element_index_col_idx = total_cols; + for (size_t i = 0; i < pipeline_ids.size(); i++) { + if (pipeline_ids[i] == SegmentOffsetFieldID) { + segment_offset_col_idx = i; + } else if (pipeline_ids[i] == ElementIndexFieldID) { + element_index_col_idx = i; + } + } + AssertInfo(segment_offset_col_idx != total_cols, + "ORDER BY pipeline must contain SegmentOffsetFieldID"); + + for (size_t i = 0; i < total_cols; i++) { + if (i == segment_offset_col_idx || i == element_index_col_idx) { + continue; + } auto* data = new DataArray(std::move(retrieveResult.field_data_[i])); data->set_field_id(pipeline_ids[i].get()); fields_data->AddAllocated(data); } - // Extract segment offsets from the last column. - auto& offset_col = retrieveResult.field_data_.back(); + // Extract segment offsets from the hidden SegmentOffsetFieldID column. + auto& offset_col = retrieveResult.field_data_[segment_offset_col_idx]; auto& offset_data = offset_col.scalars().long_data().data(); auto topk_count = offset_data.size(); @@ -343,6 +364,28 @@ SegmentInternalInterface::FillOrderByResult( // won't filter out this result (it checks len(r.GetOffset()) == 0). results->mutable_offset()->Add(offset_data.begin(), offset_data.end()); + if (element_index_col_idx != total_cols) { + auto& element_index_col = + retrieveResult.field_data_[element_index_col_idx]; + auto& element_index_data = + element_index_col.scalars().long_data().data(); + AssertInfo(element_index_data.size() == topk_count, + "element index column size ({}) must match offset column " + "size ({})", + element_index_data.size(), + topk_count); + results->set_element_level(true); + results->clear_element_indices(); + for (auto element_index : element_index_data) { + AssertInfo(element_index >= 0 && + element_index <= std::numeric_limits::max(), + "invalid element index: {}", + element_index); + auto* elem_indices = results->add_element_indices(); + elem_indices->add_indices(static_cast(element_index)); + } + } + milvus::OpContext op_ctx; // Two-project mode: bulk-fetch deferred fields using segment offsets. diff --git a/internal/core/unittest/test_element_filter.cpp b/internal/core/unittest/test_element_filter.cpp index 8b1d64f161..655e6a732b 100644 --- a/internal/core/unittest/test_element_filter.cpp +++ b/internal/core/unittest/test_element_filter.cpp @@ -2781,9 +2781,12 @@ TEST(ElementFilterGroupBy, SealedWithIndex) { ASSERT_TRUE(search_result->composite_group_by_values_.has_value()) << "Group by values should be present"; - // Group by returns row-level results even for element-level search - ASSERT_FALSE(search_result->element_level_) - << "Group by should return row-level results"; + // Group by preserves element-level metadata for element-level search. + ASSERT_TRUE(search_result->element_level_) + << "Group by should preserve element-level results"; + ASSERT_EQ(search_result->element_indices_.size(), + search_result->seg_offsets_.size()) + << "Element indices should align with grouped row offsets"; auto& group_by_values = search_result->composite_group_by_values_.value(); @@ -2904,9 +2907,12 @@ TEST(ElementFilterGroupBy, GrowingSegment) { ASSERT_TRUE(search_result->composite_group_by_values_.has_value()) << "Group by values should be present"; - // Verify element_level_ is false (group by returns row-level results) - ASSERT_FALSE(search_result->element_level_) - << "Group by should return row-level results"; + // Verify element_level_ is preserved. + ASSERT_TRUE(search_result->element_level_) + << "Group by should preserve element-level results"; + ASSERT_EQ(search_result->element_indices_.size(), + search_result->seg_offsets_.size()) + << "Element indices should align with grouped row offsets"; auto& group_by_values = search_result->composite_group_by_values_.value(); @@ -3151,8 +3157,11 @@ TEST(ElementFilterGroupBy, DeduplicateRowsInGroup) { ASSERT_NE(search_result, nullptr); ASSERT_TRUE(search_result->composite_group_by_values_.has_value()) << "Group by values should be present"; - ASSERT_FALSE(search_result->element_level_) - << "Group by should return row-level results"; + ASSERT_TRUE(search_result->element_level_) + << "Group by should preserve element-level results"; + ASSERT_EQ(search_result->element_indices_.size(), + search_result->seg_offsets_.size()) + << "Element indices should align with grouped row offsets"; auto& group_by_values = search_result->composite_group_by_values_.value(); diff --git a/internal/core/unittest/test_query_group_by.cpp b/internal/core/unittest/test_query_group_by.cpp index 79e53a53b1..cab3740fff 100644 --- a/internal/core/unittest/test_query_group_by.cpp +++ b/internal/core/unittest/test_query_group_by.cpp @@ -10,7 +10,9 @@ // or implied. See the License for the specific language governing permissions and limitations under the License #include +#include #include +#include #include "test_utils/DataGen.h" #include "segcore/SegmentSealed.h" #include "plan/PlanNode.h" @@ -20,8 +22,10 @@ #include "exec/operator/query-agg/CountAggregateBase.h" #include "exec/HashTable.h" #include "exec/VectorHasher.h" +#include "pb/plan.pb.h" #include "query/PlanImpl.h" #include "query/PlanNode.h" +#include "query/PlanProto.h" using namespace milvus; using namespace milvus::segcore; @@ -113,6 +117,92 @@ createRetrievePlan(SchemaPtr schema, return retrieve_plan; } +namespace { + +void +SetInt64FieldData(GeneratedData& raw_data, + FieldId field_id, + const std::vector& values) { + AssertInfo(raw_data.raw_->num_rows() == values.size(), + "values size must match row count"); + for (int i = 0; i < raw_data.raw_->fields_data_size(); i++) { + auto* field_data = raw_data.raw_->mutable_fields_data(i); + if (field_data->field_id() != field_id.get()) { + continue; + } + auto* data = + field_data->mutable_scalars()->mutable_long_data()->mutable_data(); + data->Clear(); + data->Add(values.data(), values.data() + values.size()); + return; + } + ThrowInfo(FieldIDInvalid, "field id not found"); +} + +void +SetIntArrayFieldData(GeneratedData& raw_data, + FieldId field_id, + const std::vector>& rows) { + AssertInfo(raw_data.raw_->num_rows() == rows.size(), + "array rows size must match row count"); + for (int i = 0; i < raw_data.raw_->fields_data_size(); i++) { + auto* field_data = raw_data.raw_->mutable_fields_data(i); + if (field_data->field_id() != field_id.get()) { + continue; + } + auto* arrays = + field_data->mutable_scalars()->mutable_array_data()->mutable_data(); + arrays->Clear(); + for (const auto& row : rows) { + auto* array_data = arrays->Add(); + array_data->mutable_int_data()->mutable_data()->Add( + row.data(), row.data() + row.size()); + } + return; + } + ThrowInfo(FieldIDInvalid, "field id not found"); +} + +void +AddPositiveElementFilter(proto::plan::QueryPlanNode* query, + FieldId array_field_id) { + auto* element_filter = + query->mutable_predicates()->mutable_element_filter_expr(); + element_filter->set_struct_name("structA"); + + auto* unary_range = + element_filter->mutable_element_expr()->mutable_unary_range_expr(); + auto* column_info = unary_range->mutable_column_info(); + column_info->set_field_id(array_field_id.get()); + column_info->set_data_type(proto::schema::DataType::Int32); + column_info->set_element_type(proto::schema::DataType::Int32); + column_info->set_is_element_level(true); + + unary_range->set_op(proto::plan::OpType::GreaterThan); + unary_range->mutable_value()->set_int64_val(0); +} + +SegmentSealedSPtr +CreateElementLevelQuerySegment(SchemaPtr schema, + FieldId pk_fid, + FieldId array_fid) { + constexpr size_t N = 3; + constexpr int array_len = 3; + auto raw_data = DataGen(schema, N, 42, 0, 1, array_len); + SetInt64FieldData(raw_data, pk_fid, {10, 20, 30}); + SetIntArrayFieldData(raw_data, + array_fid, + { + {1, 0, 3}, + {0, 5, 0}, + {7, 8, 0}, + }); + return SegmentSealedSPtr( + CreateSealedWithFieldDataLoaded(schema, raw_data).release()); +} + +} // namespace + TEST_P(QueryAggTest, GroupFixedLengthType) { std::vector sources; auto nullable = GetParam(); @@ -713,6 +803,81 @@ TEST_P(QueryAggTest, CountStarOnlyGlobalWithProjectNode) { EXPECT_EQ(num_rows_, actual_count); } +TEST(QueryAggElementLevel, CountStarUsesMatchingElements) { + auto schema = std::make_shared(); + schema->AddDebugVectorArrayField( + "structA[array_vec]", DataType::VECTOR_FLOAT, 4, knowhere::metric::L2); + auto array_fid = schema->AddDebugArrayField( + "structA[price_array]", DataType::INT32, false); + auto pk_fid = schema->AddDebugField("id", DataType::INT64); + schema->set_primary_field_id(pk_fid); + + auto segment = CreateElementLevelQuerySegment(schema, pk_fid, array_fid); + + proto::plan::PlanNode plan_node; + auto* query = plan_node.mutable_query(); + query->set_limit(100); + AddPositiveElementFilter(query, array_fid); + + auto* aggregate = query->add_aggregates(); + aggregate->set_op(proto::plan::count); + aggregate->set_field_id(0); + + auto parser = milvus::query::ProtoParser(schema); + auto plan = parser.CreateRetrievePlan(plan_node); + auto retrieve_results = segment->Retrieve( + nullptr, plan.get(), MAX_TIMESTAMP, DEFAULT_MAX_OUTPUT_SIZE, false); + + ASSERT_EQ(retrieve_results->fields_data_size(), 1); + const auto& count_data = retrieve_results->fields_data(0); + ASSERT_EQ(count_data.scalars().long_data().data_size(), 1); + EXPECT_EQ(count_data.scalars().long_data().data(0), 5); +} + +TEST(QueryAggElementLevel, GroupByPkCountStarUsesLogicalRows) { + auto schema = std::make_shared(); + schema->AddDebugVectorArrayField( + "structA[array_vec]", DataType::VECTOR_FLOAT, 4, knowhere::metric::L2); + auto array_fid = schema->AddDebugArrayField( + "structA[price_array]", DataType::INT32, false); + auto pk_fid = schema->AddDebugField("id", DataType::INT64); + schema->set_primary_field_id(pk_fid); + + auto segment = CreateElementLevelQuerySegment(schema, pk_fid, array_fid); + + proto::plan::PlanNode plan_node; + auto* query = plan_node.mutable_query(); + query->set_limit(100); + query->add_group_by_field_ids(pk_fid.get()); + AddPositiveElementFilter(query, array_fid); + + auto* aggregate = query->add_aggregates(); + aggregate->set_op(proto::plan::count); + aggregate->set_field_id(0); + + auto parser = milvus::query::ProtoParser(schema); + auto plan = parser.CreateRetrievePlan(plan_node); + auto retrieve_results = segment->Retrieve( + nullptr, plan.get(), MAX_TIMESTAMP, DEFAULT_MAX_OUTPUT_SIZE, false); + + ASSERT_EQ(retrieve_results->fields_data_size(), 2); + + const auto& pk_data = + retrieve_results->fields_data(0).scalars().long_data().data(); + const auto& count_data = + retrieve_results->fields_data(1).scalars().long_data().data(); + ASSERT_EQ(pk_data.size(), 3); + ASSERT_EQ(count_data.size(), 3); + + std::map counts_by_pk; + for (int i = 0; i < pk_data.size(); i++) { + counts_by_pk[pk_data.Get(i)] = count_data.Get(i); + } + EXPECT_EQ(counts_by_pk.at(10), 2); + EXPECT_EQ(counts_by_pk.at(20), 1); + EXPECT_EQ(counts_by_pk.at(30), 2); +} + // Test aggregation through segment->Retrieve() API to cover // fillDataArrayFromColumnVector and bitmap unpacking logic TEST_P(QueryAggTest, RetrieveAggregationWithValidityBitmap) { diff --git a/internal/core/unittest/test_query_order_by.cpp b/internal/core/unittest/test_query_order_by.cpp index f9d7511b0a..79ceeb3e0f 100644 --- a/internal/core/unittest/test_query_order_by.cpp +++ b/internal/core/unittest/test_query_order_by.cpp @@ -12,6 +12,7 @@ #include #include #include +#include #include "test_utils/DataGen.h" #include "segcore/SegmentSealed.h" #include "plan/PlanNode.h" @@ -21,14 +22,83 @@ #include "exec/QueryContext.h" #include "common/Consts.h" #include "common/FieldData.h" +#include "pb/plan.pb.h" #include "query/ExecPlanNodeVisitor.h" #include "query/PlanImpl.h" #include "query/PlanNode.h" +#include "query/PlanProto.h" using namespace milvus; using namespace milvus::segcore; using namespace milvus::plan; +namespace { + +void +SetInt64FieldData(GeneratedData& raw_data, + FieldId field_id, + const std::vector& values) { + AssertInfo(raw_data.raw_->num_rows() == values.size(), + "values size must match row count"); + for (int i = 0; i < raw_data.raw_->fields_data_size(); i++) { + auto* field_data = raw_data.raw_->mutable_fields_data(i); + if (field_data->field_id() != field_id.get()) { + continue; + } + auto* data = + field_data->mutable_scalars()->mutable_long_data()->mutable_data(); + data->Clear(); + data->Add(values.data(), values.data() + values.size()); + return; + } + ThrowInfo(FieldIDInvalid, "field id not found"); +} + +void +SetIntArrayFieldData(GeneratedData& raw_data, + FieldId field_id, + const std::vector>& rows) { + AssertInfo(raw_data.raw_->num_rows() == rows.size(), + "array rows size must match row count"); + for (int i = 0; i < raw_data.raw_->fields_data_size(); i++) { + auto* field_data = raw_data.raw_->mutable_fields_data(i); + if (field_data->field_id() != field_id.get()) { + continue; + } + auto* arrays = + field_data->mutable_scalars()->mutable_array_data()->mutable_data(); + arrays->Clear(); + for (const auto& row : rows) { + auto* array_data = arrays->Add(); + array_data->mutable_int_data()->mutable_data()->Add( + row.data(), row.data() + row.size()); + } + return; + } + ThrowInfo(FieldIDInvalid, "field id not found"); +} + +void +AddPositiveElementFilter(proto::plan::QueryPlanNode* query, + FieldId array_field_id) { + auto* element_filter = + query->mutable_predicates()->mutable_element_filter_expr(); + element_filter->set_struct_name("structA"); + + auto* unary_range = + element_filter->mutable_element_expr()->mutable_unary_range_expr(); + auto* column_info = unary_range->mutable_column_info(); + column_info->set_field_id(array_field_id.get()); + column_info->set_data_type(proto::schema::DataType::Int32); + column_info->set_element_type(proto::schema::DataType::Int32); + column_info->set_is_element_level(true); + + unary_range->set_op(proto::plan::OpType::GreaterThan); + unary_range->mutable_value()->set_int64_val(0); +} + +} // namespace + class QueryOrderByTest : public testing::TestWithParam { public: constexpr static const char int8_field[] = "int8"; @@ -261,6 +331,88 @@ TEST_P(QueryOrderByTest, EXPECT_FALSE(col->ValidAt(num_rows_)); } +TEST(QueryOrderByElementLevel, RestoresElementIndicesAfterSorting) { + auto schema = std::make_shared(); + schema->AddDebugVectorArrayField( + "structA[array_vec]", DataType::VECTOR_FLOAT, 4, knowhere::metric::L2); + auto array_fid = schema->AddDebugArrayField( + "structA[price_array]", DataType::INT32, false); + auto rank_fid = schema->AddDebugField("rank", DataType::INT64); + auto pk_fid = schema->AddDebugField("id", DataType::INT64); + schema->set_primary_field_id(pk_fid); + + constexpr size_t N = 3; + constexpr int array_len = 3; + auto raw_data = DataGen(schema, N, 42, 0, 1, array_len); + SetInt64FieldData(raw_data, pk_fid, {10, 20, 30}); + SetInt64FieldData(raw_data, rank_fid, {20, 10, 30}); + SetIntArrayFieldData(raw_data, + array_fid, + { + {1, 0, 3}, + {0, 5, 0}, + {7, 8, 0}, + }); + auto segment = SegmentSealedSPtr( + CreateSealedWithFieldDataLoaded(schema, raw_data).release()); + + proto::plan::PlanNode plan_node; + plan_node.add_output_field_ids(pk_fid.get()); + plan_node.add_output_field_ids(rank_fid.get()); + + auto* query = plan_node.mutable_query(); + query->set_limit(5); + AddPositiveElementFilter(query, array_fid); + auto* order_by = query->add_order_by_fields(); + order_by->set_field_id(rank_fid.get()); + order_by->set_ascending(true); + order_by->set_nulls_first(false); + + auto parser = milvus::query::ProtoParser(schema); + auto plan = parser.CreateRetrievePlan(plan_node); + auto results = segment->Retrieve( + nullptr, plan.get(), MAX_TIMESTAMP, DEFAULT_MAX_OUTPUT_SIZE, false); + + ASSERT_EQ(results->fields_data_size(), 2); + ASSERT_TRUE(results->element_level()); + ASSERT_EQ(results->element_indices_size(), 5); + ASSERT_EQ(results->offset_size(), 5); + + const auto& pk_data = results->fields_data(0).scalars().long_data().data(); + const auto& rank_data = + results->fields_data(1).scalars().long_data().data(); + ASSERT_EQ(pk_data.size(), 5); + ASSERT_EQ(rank_data.size(), 5); + + EXPECT_EQ(rank_data.Get(0), 10); + EXPECT_EQ(pk_data.Get(0), 20); + EXPECT_EQ(results->offset(0), 1); + ASSERT_EQ(results->element_indices(0).indices_size(), 1); + EXPECT_EQ(results->element_indices(0).indices(0), 1); + + for (int i : {1, 2}) { + EXPECT_EQ(rank_data.Get(i), 20); + EXPECT_EQ(pk_data.Get(i), 10); + EXPECT_EQ(results->offset(i), 0); + ASSERT_EQ(results->element_indices(i).indices_size(), 1); + } + std::set rank20_indices; + rank20_indices.insert(results->element_indices(1).indices(0)); + rank20_indices.insert(results->element_indices(2).indices(0)); + EXPECT_EQ(rank20_indices, (std::set{0, 2})); + + for (int i : {3, 4}) { + EXPECT_EQ(rank_data.Get(i), 30); + EXPECT_EQ(pk_data.Get(i), 30); + EXPECT_EQ(results->offset(i), 2); + ASSERT_EQ(results->element_indices(i).indices_size(), 1); + } + std::set rank30_indices; + rank30_indices.insert(results->element_indices(3).indices(0)); + rank30_indices.insert(results->element_indices(4).indices(0)); + EXPECT_EQ(rank30_indices, (std::set{0, 1})); +} + TEST_P(QueryOrderByTest, OrderBySingleInt64Asc) { auto nullable = GetParam(); std::vector pipeline_ids; diff --git a/internal/core/unittest/test_search_group_by.cpp b/internal/core/unittest/test_search_group_by.cpp index 3d260d090d..bda63ab520 100644 --- a/internal/core/unittest/test_search_group_by.cpp +++ b/internal/core/unittest/test_search_group_by.cpp @@ -32,6 +32,9 @@ #include "common/TracerBase.h" #include "common/Types.h" #include "common/Utils.h" +#include "exec/Task.h" +#include "exec/operator/SearchGroupByNode.h" +#include "exec/operator/search-groupby/SearchGroupByOperator.h" #include "common/common_type_c.h" #include "common/protobuf_utils.h" #include "common/type_c.h" @@ -63,6 +66,35 @@ using namespace milvus::query; using namespace milvus::segcore; using namespace milvus::storage; using namespace milvus::tracer; +using namespace milvus::exec; + +namespace { + +class FixedVectorIterator : public VectorIterator { + public: + explicit FixedVectorIterator(std::vector> results) + : results_(std::move(results)) { + } + + bool + HasNext() override { + return pos_ < results_.size(); + } + + std::optional> + Next() override { + if (!HasNext()) { + return std::nullopt; + } + return results_[pos_++]; + } + + private: + std::vector> results_; + size_t pos_{0}; +}; + +} // namespace const char* METRICS_TYPE = "metric_type"; @@ -503,6 +535,175 @@ TEST(GroupBY, SealedData) { } } +TEST(GroupBY, ElementLevelKeepsElementIndices) { + int dim = 4; + auto schema = std::make_shared(); + auto vec_fid = schema->AddDebugVectorArrayField("structA[array_vec]", + DataType::VECTOR_FLOAT, + dim, + knowhere::metric::L2); + auto pk_fid = schema->AddDebugField("id", DataType::INT64); + schema->set_primary_field_id(pk_fid); + + size_t N = 3; + int array_len = 3; + auto raw_data = DataGen(schema, N, 42, 0, 1, array_len); + + auto segment = CreateGrowingSegment(schema, empty_index_meta); + segment->PreInsert(N); + segment->Insert(0, + N, + raw_data.row_ids_.data(), + raw_data.timestamps_.data(), + raw_data.raw_); + + auto* growing = dynamic_cast(segment.get()); + ASSERT_NE(growing, nullptr); + auto array_offsets = growing->GetArrayOffsets(vec_fid); + ASSERT_NE(array_offsets, nullptr); + + SearchInfo search_info; + search_info.topk_ = 10; + search_info.group_size_ = 2; + search_info.metric_type_ = knowhere::metric::L2; + search_info.group_by_field_ids_.push_back(pk_fid); + search_info.array_offsets_ = array_offsets; + + // Element IDs map to rows as: + // row 0: element IDs 0, 1, 2 + // row 1: element IDs 3, 4, 5 + // row 2: element IDs 6, 7, 8 + // group_size=2 should keep only two hits from row 0 while preserving + // element indices for all accepted hits. + std::vector> iterators{ + std::make_shared( + std::vector>{ + {0, 0.1F}, + {1, 0.2F}, + {2, 0.3F}, + {3, 0.4F}, + {4, 0.5F}, + {6, 0.6F}, + })}; + + OpContext op_context; + std::vector group_by_values; + std::vector seg_offsets; + std::vector distances; + std::vector topk_per_nq_prefix_sum; + std::vector element_indices; + + SearchGroupBy(&op_context, + iterators, + search_info, + group_by_values, + *growing, + seg_offsets, + distances, + topk_per_nq_prefix_sum, + &element_indices); + + ASSERT_EQ(seg_offsets, (std::vector{0, 0, 1, 1, 2})); + ASSERT_EQ(element_indices, (std::vector{0, 1, 0, 1, 0})); + ASSERT_EQ(topk_per_nq_prefix_sum, (std::vector{0, 5})); + ASSERT_EQ(group_by_values.size(), seg_offsets.size()); + ASSERT_EQ(distances, (std::vector{0.1F, 0.2F, 0.4F, 0.5F, 0.6F})); +} + +TEST(GroupBY, SearchGroupByNodeKeepsElementIndices) { + int dim = 4; + auto schema = std::make_shared(); + auto vec_fid = schema->AddDebugVectorArrayField("structA[array_vec]", + DataType::VECTOR_FLOAT, + dim, + knowhere::metric::L2); + auto pk_fid = schema->AddDebugField("id", DataType::INT64); + schema->set_primary_field_id(pk_fid); + + size_t N = 3; + int array_len = 3; + auto raw_data = DataGen(schema, N, 42, 0, 1, array_len); + + auto segment = CreateGrowingSegment(schema, empty_index_meta); + segment->PreInsert(N); + segment->Insert(0, + N, + raw_data.row_ids_.data(), + raw_data.timestamps_.data(), + raw_data.raw_); + + auto* growing = dynamic_cast(segment.get()); + ASSERT_NE(growing, nullptr); + auto array_offsets = growing->GetArrayOffsets(vec_fid); + ASSERT_NE(array_offsets, nullptr); + + SearchInfo search_info; + search_info.topk_ = 10; + search_info.group_size_ = 2; + search_info.metric_type_ = knowhere::metric::L2; + search_info.group_by_field_ids_.push_back(pk_fid); + search_info.array_offsets_ = array_offsets; + + SearchResult search_result; + search_result.total_nq_ = 1; + search_result.unity_topK_ = search_info.topk_; + search_result.element_level_ = true; + search_result.vector_iterators_ = + std::vector>{ + std::make_shared( + std::vector>{ + {0, 0.1F}, + {1, 0.2F}, + {2, 0.3F}, + {3, 0.4F}, + {4, 0.5F}, + {6, 0.6F}, + })}; + + OpContext op_context; + auto query_context = std::make_shared( + "search_group_by_node_element_level", + growing, + N, + MAX_TIMESTAMP, + 0, + 0, + query::PlanOptions{false}, + std::make_shared( + std::unordered_map{})); + query_context->set_search_info(search_info); + query_context->set_array_offsets(array_offsets); + query_context->set_search_result(std::move(search_result)); + query_context->set_op_context(&op_context); + + auto group_by_plan = + std::make_shared("group_by"); + auto task = + milvus::exec::Task::Create("search_group_by_node_element_level", + milvus::plan::PlanFragment(group_by_plan), + 0, + query_context); + milvus::exec::DriverContext driver_context(task, 0, 0, 0, 0); + milvus::exec::PhySearchGroupByNode node(1, &driver_context, group_by_plan); + + auto input = std::make_shared(std::vector{}); + node.AddInput(input); + node.NoMoreInput(); + auto output = node.GetOutput(); + ASSERT_NE(output, nullptr); + + auto grouped = query_context->get_search_result(); + ASSERT_TRUE(grouped.element_level_); + ASSERT_EQ(grouped.seg_offsets_, (std::vector{0, 0, 1, 1, 2})); + ASSERT_EQ(grouped.element_indices_, (std::vector{0, 1, 0, 1, 0})); + ASSERT_EQ(grouped.topk_per_nq_prefix_sum_, (std::vector{0, 5})); + ASSERT_TRUE(grouped.composite_group_by_values_.has_value()); + ASSERT_EQ(grouped.composite_group_by_values_->size(), + grouped.seg_offsets_.size()); + ASSERT_TRUE(grouped.group_size_.has_value()); + ASSERT_EQ(grouped.group_size_.value(), 2); +} + TEST(GroupBY, GrowingRawData) { //0. set up growing segment int dim = 128; @@ -695,4 +896,4 @@ TEST(GroupBY, GrowingIndex) { ASSERT_EQ(group_size, map_pair.second); } } -} \ No newline at end of file +} diff --git a/internal/proxy/query_pipeline_test.go b/internal/proxy/query_pipeline_test.go index 448a585343..c05cd7c2b5 100644 --- a/internal/proxy/query_pipeline_test.go +++ b/internal/proxy/query_pipeline_test.go @@ -42,6 +42,25 @@ func testSchema() *schemapb.CollectionSchema { } } +func testSchemaWithStructArray() *schemapb.CollectionSchema { + schema := testSchema() + schema.StructArrayFields = []*schemapb.StructArrayFieldSchema{ + { + FieldID: 200, + Name: "items", + Fields: []*schemapb.FieldSchema{ + { + FieldID: 201, + Name: "items[tag]", + DataType: schemapb.DataType_Array, + ElementType: schemapb.DataType_VarChar, + }, + }, + }, + } + return schema +} + func makeTestIntIDs(ids []int64) *schemapb.IDs { return &schemapb.IDs{IdField: &schemapb.IDs_IntId{IntId: &schemapb.LongArray{Data: ids}}} } @@ -267,6 +286,67 @@ func TestNewQueryPipeline_GroupBy(t *testing.T) { assert.Equal(t, 2, len(result.GetFieldsData())) } +func TestNewQueryPipeline_GroupByCountStarWithStructSchema(t *testing.T) { + schema := testSchemaWithStructArray() + + countAggs, err := agg.NewAggregate("count", 0, "count(*)", schemapb.DataType_None) + require.NoError(t, err) + aggsBases := make([]agg.AggregateBase, len(countAggs)) + copy(aggsBases, countAggs) + outputMap, err := agg.NewAggregationFieldMap( + []string{"pk", "count(*)"}, + []string{"pk"}, + aggsBases, + ) + require.NoError(t, err) + + pipeline, err := NewQueryPipeline( + schema, 10, 0, reduce.IReduceNoOrder, + nil, + []int64{100}, + []*planpb.Aggregate{{Op: planpb.AggregateOp_count, FieldId: 0}}, + outputMap, + nil, + ) + require.NoError(t, err) + + r1 := &internalpb.RetrieveResults{ + FieldsData: []*schemapb.FieldData{ + makeTestInt64Field(0, "", []int64{1, 2}), + makeTestInt64Field(0, "", []int64{2, 1}), + }, + } + r2 := &internalpb.RetrieveResults{ + FieldsData: []*schemapb.FieldData{ + makeTestInt64Field(0, "", []int64{1, 3}), + makeTestInt64Field(0, "", []int64{3, 4}), + }, + } + + result, err := pipeline.Execute(context.Background(), []*internalpb.RetrieveResults{r1, r2}) + require.NoError(t, err) + require.NotNil(t, result) + + result.OutputFields = []string{"pk", "count(*)"} + reconstructStructFieldDataForQuery(result, schema) + + require.Len(t, result.GetFieldsData(), 2) + assert.Equal(t, "pk", result.GetFieldsData()[0].GetFieldName()) + assert.Equal(t, "count(*)", result.GetFieldsData()[1].GetFieldName()) + assert.Equal(t, int64(0), result.GetFieldsData()[1].GetFieldId()) + assert.Equal(t, []string{"pk", "count(*)"}, result.GetOutputFields()) + + pks := result.GetFieldsData()[0].GetScalars().GetLongData().GetData() + counts := result.GetFieldsData()[1].GetScalars().GetLongData().GetData() + require.Len(t, pks, len(counts)) + + countByPK := make(map[int64]int64, len(pks)) + for i, pk := range pks { + countByPK[pk] = counts[i] + } + assert.Equal(t, map[int64]int64{1: 5, 2: 1, 3: 4}, countByPK) +} + func TestNewQueryPipeline_GroupByOrderBy(t *testing.T) { schema := testSchema() diff --git a/internal/proxy/search_pipeline.go b/internal/proxy/search_pipeline.go index c0c603d609..a74f3b5918 100644 --- a/internal/proxy/search_pipeline.go +++ b/internal/proxy/search_pipeline.go @@ -628,6 +628,8 @@ func prepareElementLevelHybridResult(result *milvuspb.SearchResults) (*milvuspb. AllSearchCount: data.GetAllSearchCount(), PrimaryFieldName: data.GetPrimaryFieldName(), ElementIndices: &schemapb.LongArray{}, + GroupByFieldValues: append([]*schemapb.FieldData(nil), data.GetGroupByFieldValues()...), + GroupByFieldValue: data.GetGroupByFieldValue(), SearchIteratorV2Results: data.GetSearchIteratorV2Results(), } if len(data.GetDistances()) > 0 { @@ -665,6 +667,8 @@ func prepareElementLevelHybridResult(result *milvuspb.SearchResults) (*milvuspb. AllSearchCount: data.GetAllSearchCount(), PrimaryFieldName: data.GetPrimaryFieldName(), ElementIndices: data.GetElementIndices(), + GroupByFieldValues: append([]*schemapb.FieldData(nil), data.GetGroupByFieldValues()...), + GroupByFieldValue: data.GetGroupByFieldValue(), SearchIteratorV2Results: data.GetSearchIteratorV2Results(), } if len(data.GetDistances()) > 0 { @@ -753,6 +757,8 @@ func restoreElementLevelHybridRankResult(rankResult *milvuspb.SearchResults) (*m AllSearchCount: data.GetAllSearchCount(), PrimaryFieldName: data.GetPrimaryFieldName(), ElementIndices: &schemapb.LongArray{Data: elementIndices}, + GroupByFieldValues: append([]*schemapb.FieldData(nil), data.GetGroupByFieldValues()...), + GroupByFieldValue: data.GetGroupByFieldValue(), SearchIteratorV2Results: data.GetSearchIteratorV2Results(), } if len(data.GetDistances()) > 0 { diff --git a/internal/proxy/search_pipeline_test.go b/internal/proxy/search_pipeline_test.go index 94477251bc..92be55d90d 100644 --- a/internal/proxy/search_pipeline_test.go +++ b/internal/proxy/search_pipeline_test.go @@ -739,6 +739,14 @@ func (s *SearchPipelineSuite) TestElementLevelHybridPrepareAndRestoreKeys() { }, Scores: []float32{0.8, 0.9, 0.7}, ElementIndices: &schemapb.LongArray{Data: []int64{0, 2, 1}}, + GroupByFieldValues: []*schemapb.FieldData{{ + FieldId: 100, + FieldName: "pk", + Type: schemapb.DataType_Int64, + Field: &schemapb.FieldData_Scalars{Scalars: &schemapb.ScalarField{ + Data: &schemapb.ScalarField_LongData{LongData: &schemapb.LongArray{Data: []int64{10, 10, 20}}}, + }}, + }}, }, } @@ -746,6 +754,8 @@ func (s *SearchPipelineSuite) TestElementLevelHybridPrepareAndRestoreKeys() { s.Require().NoError(err) preparedIDs := prepared.GetResults().GetIds().GetStrId().GetData() s.Require().Len(preparedIDs, 3) + s.Require().Len(prepared.GetResults().GetGroupByFieldValues(), 1) + s.Equal([]int64{10, 10, 20}, prepared.GetResults().GetGroupByFieldValues()[0].GetScalars().GetLongData().GetData()) rankResult := &milvuspb.SearchResults{ Status: merr.Success(), @@ -759,6 +769,14 @@ func (s *SearchPipelineSuite) TestElementLevelHybridPrepareAndRestoreKeys() { }, }, Scores: []float32{0.99, 0.88}, + GroupByFieldValues: []*schemapb.FieldData{{ + FieldId: 100, + FieldName: "pk", + Type: schemapb.DataType_Int64, + Field: &schemapb.FieldData_Scalars{Scalars: &schemapb.ScalarField{ + Data: &schemapb.ScalarField_LongData{LongData: &schemapb.LongArray{Data: []int64{10, 20}}}, + }}, + }}, }, } @@ -767,6 +785,8 @@ func (s *SearchPipelineSuite) TestElementLevelHybridPrepareAndRestoreKeys() { s.Equal([]int64{10, 20}, restored.GetResults().GetIds().GetIntId().GetData()) s.Equal([]int64{2, 1}, restored.GetResults().GetElementIndices().GetData()) s.Equal([]float32{0.99, 0.88}, restored.GetResults().GetScores()) + s.Require().Len(restored.GetResults().GetGroupByFieldValues(), 1) + s.Equal([]int64{10, 20}, restored.GetResults().GetGroupByFieldValues()[0].GetScalars().GetLongData().GetData()) _, err = prepareElementLevelHybridResult(&milvuspb.SearchResults{ Status: merr.Success(), diff --git a/internal/proxy/task_search.go b/internal/proxy/task_search.go index 5403b9ad75..7b4602f29f 100644 --- a/internal/proxy/task_search.go +++ b/internal/proxy/task_search.go @@ -580,6 +580,10 @@ func (t *searchTask) initAdvancedSearchRequest(ctx context.Context) error { subSearchInfo.Collapse = collapseConfig } + // ArrayOfVector hybrid validation is kind-specific: embedding-list + // rejects range/iterator here, element-level rejects legacy iterator + // here, and group-by is validated after same-struct inference below. + // Hybrid search only supports plain top-K on ArrayOfVector fields. Both // element-level and embedding-list searches reject advanced controls here. if annsField != nil && annsField.GetDataType() == schemapb.DataType_ArrayOfVector { @@ -590,14 +594,10 @@ func (t *searchTask) initAdvancedSearchRequest(ctx context.Context) error { if isStructEmbListSubSearch { searchKind = "embedding-list" } - if gjson.Get(queryInfo.GetSearchParams(), radiusKey).Exists() { + if isStructEmbListSubSearch && gjson.Get(queryInfo.GetSearchParams(), radiusKey).Exists() { return merr.WrapErrParameterInvalid("", "", "range search is not supported for vector array ("+searchKind+") fields in hybrid search, fieldName:"+annsField.GetName()) } - if t.rankParams.GetGroupByFieldId() > 0 || len(t.rankParams.GetGroupByFieldIds()) > 0 { - return merr.WrapErrParameterInvalid("", "", - "group by search is not supported for vector array ("+searchKind+") fields in hybrid search, fieldName:"+annsField.GetName()) - } if subIsIterator { return merr.WrapErrParameterInvalid("", "", "search iterator is not supported for vector array ("+searchKind+") fields in hybrid search, fieldName:"+annsField.GetName()) @@ -696,6 +696,9 @@ func (t *searchTask) initAdvancedSearchRequest(ctx context.Context) error { } t.hybridElementLevel = inferElementLevelHybrid(t.hybridSubSearchInfos) + if err := t.validateHybridArrayOfVectorGroupBy(); err != nil { + return err + } for index, info := range t.hybridSubSearchInfos { if t.hybridElementLevel && info.ElementScopeProvided { return merr.WrapErrParameterInvalidMsg("%s is not allowed for same-struct element-level hybrid search", elementScopeKey) @@ -729,6 +732,50 @@ func (t *searchTask) initAdvancedSearchRequest(ctx context.Context) error { return nil } +func (t *searchTask) validateHybridArrayOfVectorGroupBy() error { + groupByFieldIDs := t.rankParams.GetGroupByFieldIds() + if len(groupByFieldIDs) == 0 { + return nil + } + + fieldName := func(index int) string { + if index >= 0 && index < len(t.queryInfos) { + if field := typeutil.GetField(t.schema.CollectionSchema, t.queryInfos[index].GetQueryFieldId()); field != nil { + return field.GetName() + } + } + return "" + } + + hasStructElementSubSearch := false + for index, info := range t.hybridSubSearchInfos { + switch info.Kind { + case hybridSubSearchStructEmbList: + return merr.WrapErrParameterInvalid("", "", + "group by search is not supported for vector array (embedding-list) fields in hybrid search, fieldName:"+fieldName(index)) + case hybridSubSearchStructElement: + hasStructElementSubSearch = true + if !t.hybridElementLevel { + return merr.WrapErrParameterInvalid("", "", + "group by search is only supported for same-struct element-level vector array fields in hybrid search, fieldName:"+fieldName(index)) + } + } + } + if !hasStructElementSubSearch { + return nil + } + + pkField, err := t.schema.GetPkField() + if err != nil { + return err + } + if len(groupByFieldIDs) != 1 || pkField == nil || groupByFieldIDs[0] != pkField.GetFieldID() { + return merr.WrapErrParameterInvalid("", "", + "only group by primary key is supported for same-struct element-level vector array fields in hybrid search") + } + return nil +} + func (t *searchTask) fillResult() { limit := t.GetTopk() - t.GetOffset() resultSizeInsufficient := false diff --git a/internal/proxy/task_search_test.go b/internal/proxy/task_search_test.go index 75443c0282..0fc3d5f4d1 100644 --- a/internal/proxy/task_search_test.go +++ b/internal/proxy/task_search_test.go @@ -6495,7 +6495,8 @@ func TestSearchTask_ArrayOfVectorHybridSearch(t *testing.T) { Name: "struct_array", Fields: []*schemapb.FieldSchema{ {FieldID: 104, Name: "emb_vec", DataType: schemapb.DataType_ArrayOfVector, ElementType: schemapb.DataType_FloatVector, TypeParams: []*commonpb.KeyValuePair{{Key: common.DimKey, Value: "4"}}}, - {FieldID: 105, Name: typeutil.ConcatStructFieldName("struct_array", "price"), DataType: schemapb.DataType_Array, ElementType: schemapb.DataType_Int64}, + {FieldID: 105, Name: "emb_text_vec", DataType: schemapb.DataType_ArrayOfVector, ElementType: schemapb.DataType_FloatVector, TypeParams: []*commonpb.KeyValuePair{{Key: common.DimKey, Value: "4"}}}, + {FieldID: 106, Name: typeutil.ConcatStructFieldName("struct_array", "price"), DataType: schemapb.DataType_Array, ElementType: schemapb.DataType_Int64}, }, }, }, @@ -6514,19 +6515,35 @@ func TestSearchTask_ArrayOfVectorHybridSearch(t *testing.T) { return bs } - buildHybridTaskWithMetric := func(annsField string, metricType string, phType commonpb.PlaceholderType, rangeRadius string, withIterator bool, groupByField string) *searchTask { - paramsJSON := `{"nprobe": 10}` - if rangeRadius != "" { - paramsJSON = `{"nprobe": 10, "radius": ` + rangeRadius + `}` - } - subParams := []*commonpb.KeyValuePair{ - {Key: common.MetricTypeKey, Value: metricType}, - {Key: ParamsKey, Value: paramsJSON}, - {Key: AnnsFieldKey, Value: annsField}, - {Key: TopKKey, Value: "10"}, - } - if withIterator { - subParams = append(subParams, &commonpb.KeyValuePair{Key: IteratorField, Value: "True"}) + type hybridSubSpec struct { + annsField string + metricType string + phType commonpb.PlaceholderType + rangeRadius string + withIterator bool + } + + buildHybridTaskWithSubSpecs := func(groupByField string, specs ...hybridSubSpec) *searchTask { + subReqs := make([]*milvuspb.SubSearchRequest, 0, len(specs)) + for _, spec := range specs { + paramsJSON := `{"nprobe": 10}` + if spec.rangeRadius != "" { + paramsJSON = `{"nprobe": 10, "radius": ` + spec.rangeRadius + `}` + } + subParams := []*commonpb.KeyValuePair{ + {Key: common.MetricTypeKey, Value: spec.metricType}, + {Key: ParamsKey, Value: paramsJSON}, + {Key: AnnsFieldKey, Value: spec.annsField}, + {Key: TopKKey, Value: "10"}, + } + if spec.withIterator { + subParams = append(subParams, &commonpb.KeyValuePair{Key: IteratorField, Value: "True"}) + } + subReqs = append(subReqs, &milvuspb.SubSearchRequest{ + Dsl: "", + PlaceholderGroup: makePlaceholderGroup(spec.phType), + SearchParams: subParams, + }) } outerParams := []*commonpb.KeyValuePair{ @@ -6547,19 +6564,23 @@ func TestSearchTask_ArrayOfVectorHybridSearch(t *testing.T) { request: &milvuspb.SearchRequest{ CollectionName: "test_collection", SearchParams: outerParams, - SubReqs: []*milvuspb.SubSearchRequest{ - { - Dsl: "", - PlaceholderGroup: makePlaceholderGroup(phType), - SearchParams: subParams, - }, - }, + SubReqs: subReqs, }, schema: schemaInfo, tr: timerecord.NewTimeRecorder("test"), } } + buildHybridTaskWithMetric := func(annsField string, metricType string, phType commonpb.PlaceholderType, rangeRadius string, withIterator bool, groupByField string) *searchTask { + return buildHybridTaskWithSubSpecs(groupByField, hybridSubSpec{ + annsField: annsField, + metricType: metricType, + phType: phType, + rangeRadius: rangeRadius, + withIterator: withIterator, + }) + } + buildHybridTask := func(annsField string, rangeRadius string, withIterator bool, groupByField string) *searchTask { return buildHybridTaskWithMetric(annsField, metric.MaxSimL2, commonpb.PlaceholderType_EmbListFloatVector, rangeRadius, withIterator, groupByField) } @@ -6568,6 +6589,13 @@ func TestSearchTask_ArrayOfVectorHybridSearch(t *testing.T) { return buildHybridTaskWithMetric(annsField, metric.L2, commonpb.PlaceholderType_FloatVector, rangeRadius, withIterator, groupByField) } + buildSameStructElementHybridTask := func(groupByField string) *searchTask { + return buildHybridTaskWithSubSpecs(groupByField, + hybridSubSpec{annsField: "emb_vec", metricType: metric.L2, phType: commonpb.PlaceholderType_FloatVector}, + hybridSubSpec{annsField: "emb_text_vec", metricType: metric.L2, phType: commonpb.PlaceholderType_FloatVector}, + ) + } + t.Run("hybrid with ArrayOfVector EmbList metric plain topK should succeed", func(t *testing.T) { qt := buildHybridTaskWithMetric("emb_vec", metric.MaxSimCosine, commonpb.PlaceholderType_EmbListFloatVector, "", false, "") err := qt.initAdvancedSearchRequest(ctx) @@ -6598,6 +6626,14 @@ func TestSearchTask_ArrayOfVectorHybridSearch(t *testing.T) { assert.Contains(t, err.Error(), "group by search is not supported for vector array (embedding-list) fields in hybrid search") }) + t.Run("hybrid with ArrayOfVector group by PK should fail for embedding-list", func(t *testing.T) { + qt := buildHybridTask("emb_vec", "", false, "pk") + err := qt.initAdvancedSearchRequest(ctx) + assert.Error(t, err) + assert.ErrorIs(t, err, merr.ErrParameterInvalid) + assert.Contains(t, err.Error(), "group by search is not supported for vector array (embedding-list) fields in hybrid search") + }) + t.Run("hybrid with element-level ArrayOfVector plain topK should succeed", func(t *testing.T) { qt := buildElementHybridTask("emb_vec", "", false, "") err := qt.initAdvancedSearchRequest(ctx) @@ -6624,12 +6660,10 @@ func TestSearchTask_ArrayOfVectorHybridSearch(t *testing.T) { assert.NoError(t, err) }) - t.Run("hybrid with element-level ArrayOfVector range search should fail", func(t *testing.T) { + t.Run("hybrid with element-level ArrayOfVector range search should succeed", func(t *testing.T) { qt := buildElementHybridTask("emb_vec", "0.2", false, "") err := qt.initAdvancedSearchRequest(ctx) - assert.Error(t, err) - assert.ErrorIs(t, err, merr.ErrParameterInvalid) - assert.Contains(t, err.Error(), "range search is not supported for vector array (element-level) fields in hybrid search") + assert.NoError(t, err) }) t.Run("hybrid with element-level ArrayOfVector iterator should fail", func(t *testing.T) { @@ -6640,12 +6674,47 @@ func TestSearchTask_ArrayOfVectorHybridSearch(t *testing.T) { assert.Contains(t, err.Error(), "search iterator is not supported for vector array (element-level) fields in hybrid search") }) - t.Run("hybrid with element-level ArrayOfVector group by should fail", func(t *testing.T) { - qt := buildElementHybridTask("emb_vec", "", false, "pk") + t.Run("hybrid with same-struct element-level ArrayOfVector group by PK should succeed", func(t *testing.T) { + qt := buildSameStructElementHybridTask("pk") + err := qt.initAdvancedSearchRequest(ctx) + assert.NoError(t, err) + assert.True(t, qt.hybridElementLevel) + assert.Equal(t, int64(100), qt.GroupByFieldId) + }) + + t.Run("hybrid with element-level ArrayOfVector group by non-PK should fail", func(t *testing.T) { + qt := buildElementHybridTask("emb_vec", "", false, "scalar_field") err := qt.initAdvancedSearchRequest(ctx) assert.Error(t, err) assert.ErrorIs(t, err, merr.ErrParameterInvalid) - assert.Contains(t, err.Error(), "group by search is not supported for vector array (element-level) fields in hybrid search") + assert.Contains(t, err.Error(), "only group by primary key is supported") + }) + + t.Run("hybrid with same-struct element-level ArrayOfVector composite group by should fail", func(t *testing.T) { + qt := buildSameStructElementHybridTask("") + qt.request.SearchParams = append(qt.request.SearchParams, &commonpb.KeyValuePair{Key: GroupByFieldsKey, Value: "pk,scalar_field"}) + err := qt.initAdvancedSearchRequest(ctx) + assert.Error(t, err) + assert.ErrorIs(t, err, merr.ErrParameterInvalid) + assert.Contains(t, err.Error(), "only group by primary key is supported") + }) + + t.Run("hybrid with mixed element-level ArrayOfVector group by should fail", func(t *testing.T) { + qt := buildHybridTaskWithSubSpecs("pk", + hybridSubSpec{annsField: "emb_vec", metricType: metric.L2, phType: commonpb.PlaceholderType_FloatVector}, + hybridSubSpec{annsField: "regular_vec", metricType: metric.L2, phType: commonpb.PlaceholderType_FloatVector}, + ) + err := qt.initAdvancedSearchRequest(ctx) + assert.Error(t, err) + assert.ErrorIs(t, err, merr.ErrParameterInvalid) + assert.Contains(t, err.Error(), "same-struct element-level") + }) + + t.Run("hybrid with single element-level ArrayOfVector group by PK should succeed", func(t *testing.T) { + qt := buildElementHybridTask("emb_vec", "", false, "pk") + err := qt.initAdvancedSearchRequest(ctx) + assert.NoError(t, err) + assert.True(t, qt.hybridElementLevel) }) t.Run("hybrid with normal vector advanced controls should succeed", func(t *testing.T) { diff --git a/internal/proxy/util.go b/internal/proxy/util.go index 27d8a8df58..fcf46f3221 100644 --- a/internal/proxy/util.go +++ b/internal/proxy/util.go @@ -3099,9 +3099,11 @@ func reconstructStructFieldData( if _, ok := regularFieldIDs[fieldID]; ok { newFieldsData = append(newFieldsData, field) reconstructedOutputFields = append(reconstructedOutputFields, field.GetFieldName()) - } else { - structFieldID := subFieldToStructMap[fieldID] + } else if structFieldID, ok := subFieldToStructMap[fieldID]; ok { groupedStructFields[structFieldID] = append(groupedStructFields[structFieldID], field) + } else { + newFieldsData = append(newFieldsData, field) + reconstructedOutputFields = append(reconstructedOutputFields, field.GetFieldName()) } } diff --git a/internal/proxy/util_test.go b/internal/proxy/util_test.go index c2e2116580..2478eb4675 100644 --- a/internal/proxy/util_test.go +++ b/internal/proxy/util_test.go @@ -4615,6 +4615,71 @@ func Test_reconstructStructFieldData(t *testing.T) { assert.Equal(t, originalOutputFields, resultOutputFields) }) + t.Run("group by with count(*) should preserve aggregate field", func(t *testing.T) { + fieldsData := []*schemapb.FieldData{ + { + FieldName: "id", + FieldId: 100, + Type: schemapb.DataType_Int64, + Field: &schemapb.FieldData_Scalars{ + Scalars: &schemapb.ScalarField{ + Data: &schemapb.ScalarField_LongData{ + LongData: &schemapb.LongArray{Data: []int64{1, 2}}, + }, + }, + }, + }, + { + FieldName: "count(*)", + FieldId: 0, + Type: schemapb.DataType_Int64, + Field: &schemapb.FieldData_Scalars{ + Scalars: &schemapb.ScalarField{ + Data: &schemapb.ScalarField_LongData{ + LongData: &schemapb.LongArray{Data: []int64{3, 4}}, + }, + }, + }, + }, + } + outputFields := []string{"id", "count(*)"} + + schema := &schemapb.CollectionSchema{ + Fields: []*schemapb.FieldSchema{ + { + FieldID: 100, + Name: "id", + IsPrimaryKey: true, + DataType: schemapb.DataType_Int64, + }, + }, + StructArrayFields: []*schemapb.StructArrayFieldSchema{ + { + FieldID: 102, + Name: "test_struct", + Fields: []*schemapb.FieldSchema{ + { + FieldID: 1021, + Name: "test_struct[sub_field]", + DataType: schemapb.DataType_Array, + ElementType: schemapb.DataType_Int32, + }, + }, + }, + }, + } + + resultFieldsData, resultOutputFields := reconstructStructFieldData(fieldsData, outputFields, schema) + + assert.Len(t, resultFieldsData, 2) + assert.Equal(t, "id", resultFieldsData[0].FieldName) + assert.Equal(t, int64(100), resultFieldsData[0].FieldId) + assert.Equal(t, "count(*)", resultFieldsData[1].FieldName) + assert.Equal(t, int64(0), resultFieldsData[1].FieldId) + assert.Equal(t, []string{"id", "count(*)"}, resultOutputFields) + assert.Equal(t, []int64{3, 4}, resultFieldsData[1].GetScalars().GetLongData().GetData()) + }) + t.Run("struct field query - should reconstruct struct field", func(t *testing.T) { fieldsData := []*schemapb.FieldData{ {