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 <tangchenjie1210@gmail.com>
This commit is contained in:
Spade A
2026-06-24 14:44:26 +08:00
committed by GitHub
parent 79a42f5d44
commit 6dc84c51bb
19 changed files with 1089 additions and 96 deletions
+3
View File
@@ -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";
+103 -20
View File
@@ -17,6 +17,7 @@
#include "ProjectNode.h"
#include <algorithm>
#include <optional>
#include <utility>
#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<int64_t> row_offsets;
std::vector<int32_t> 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<int64_t>& values) {
auto selected_count = values.size();
FixedVector<int64_t> output(selected_count);
milvus::fastmem::FastMemcpy(
output.data(), values.data(), selected_count * sizeof(int64_t));
auto field_data = std::make_shared<FieldData<int64_t>>(
DataType::INT64, false, std::move(output));
auto valid_map = TargetBitmap(selected_count, true);
return std::make_shared<ColumnVector>(
std::move(field_data), std::move(valid_map), 0);
}
ColumnVectorPtr
MakeElementIndexColumn(const std::vector<int32_t>& element_indices) {
auto selected_count = element_indices.size();
FixedVector<int64_t> output(selected_count);
for (size_t i = 0; i < selected_count; i++) {
output[i] = element_indices[i];
}
auto field_data = std::make_shared<FieldData<int64_t>>(
DataType::INT64, false, std::move(output));
auto valid_map = TargetBitmap(selected_count, true);
return std::make_shared<ColumnVector>(
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<int64_t>(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<int64_t> offsets(selected_count);
milvus::fastmem::FastMemcpy(offsets.data(),
selected_offsets.data(),
selected_count * sizeof(int64_t));
auto field_data = std::make_shared<FieldData<int64_t>>(
DataType::INT64, false, std::move(offsets));
auto valid_map = TargetBitmap(selected_count, true);
auto offset_col = std::make_shared<ColumnVector>(
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
}; // namespace milvus
@@ -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<FieldId> fields_to_project_;
QueryContext* query_context_;
OpContext* op_context_;
};
} // namespace exec
@@ -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
} // namespace milvus
@@ -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 <typename T, typename InnerRawType = T>
static std::function<GroupByValueType(int64_t)>
@@ -189,18 +200,21 @@ GroupIteratorResult(const std::shared_ptr<VectorIterator>& iterator,
std::vector<CompositeGroupKey>& composite_group_by_values,
std::vector<int64_t>& offsets,
std::vector<float>& distances,
const SearchInfo& search_info) {
const SearchInfo& search_info,
std::vector<int32_t>* 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<std::tuple<int64_t, float, CompositeGroupKey>> res;
std::vector<GroupedResult> res;
CompositeGroupKey scratch_key;
while (iterator->HasNext() && !groupMap.IsGroupResEnough()) {
auto offset_dis_pair = iterator->Next();
@@ -212,35 +226,41 @@ GroupIteratorResult(const std::shared_ptr<VectorIterator>& 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<int32_t>(raw_offset))
.first;
auto [doc_id, elem_idx] =
search_info.array_offsets_->ElementIDToRowID(
static_cast<int32_t>(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<int64_t>& seg_offsets,
std::vector<float>& distances,
std::vector<size_t>& topk_per_nq_prefix_sum) {
std::vector<size_t>& topk_per_nq_prefix_sum,
std::vector<int32_t>* 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());
}
}
@@ -432,7 +432,8 @@ SearchGroupBy(milvus::OpContext* op_ctx,
const segcore::SegmentInternalInterface& segment,
std::vector<int64_t>& seg_offsets,
std::vector<float>& distances,
std::vector<size_t>& topk_per_nq_prefix_sum);
std::vector<size_t>& topk_per_nq_prefix_sum,
std::vector<int32_t>* element_indices = nullptr);
} // namespace exec
} // namespace milvus
+16 -3
View File
@@ -390,7 +390,8 @@ std::tuple<plan::PlanNodePtr, std::vector<FieldId>, std::vector<FieldId>>
BuildOrderByProjectNode(const proto::plan::QueryPlanNode& query,
const planpb::PlanNode& plan_node_proto,
const SchemaPtr& schema,
const std::vector<plan::PlanNodePtr>& sources) {
const std::vector<plan::PlanNodePtr>& sources,
bool is_element_level) {
std::vector<FieldId> project_ids;
std::vector<std::string> project_names;
std::vector<milvus::DataType> 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<milvus::plan::PlanNodePtr>{plannode};
plan_node->deferred_field_ids_ = std::move(deferred);
+51 -8
View File
@@ -15,6 +15,7 @@
#include <algorithm>
#include <chrono>
#include <cstdint>
#include <limits>
#include <map>
#include <ratio>
#include <type_traits>
@@ -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<int32_t>::max(),
"invalid element index: {}",
element_index);
auto* elem_indices = results->add_element_indices();
elem_indices->add_indices(static_cast<int32_t>(element_index));
}
}
milvus::OpContext op_ctx;
// Two-project mode: bulk-fetch deferred fields using segment offsets.
+17 -8
View File
@@ -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();
@@ -10,7 +10,9 @@
// or implied. See the License for the specific language governing permissions and limitations under the License
#include <gtest/gtest.h>
#include <map>
#include <set>
#include <vector>
#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<int64_t>& 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<std::vector<int32_t>>& 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<milvus::plan::PlanNodePtr> 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>();
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>();
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<int64_t, int64_t> 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) {
@@ -12,6 +12,7 @@
#include <gtest/gtest.h>
#include <algorithm>
#include <set>
#include <vector>
#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<int64_t>& 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<std::vector<int32_t>>& 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<bool> {
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>();
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<int32_t> 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<int32_t>{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<int32_t> 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<int32_t>{0, 1}));
}
TEST_P(QueryOrderByTest, OrderBySingleInt64Asc) {
auto nullable = GetParam();
std::vector<FieldId> pipeline_ids;
+202 -1
View File
@@ -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<std::pair<int64_t, float>> results)
: results_(std::move(results)) {
}
bool
HasNext() override {
return pos_ < results_.size();
}
std::optional<std::pair<int64_t, float>>
Next() override {
if (!HasNext()) {
return std::nullopt;
}
return results_[pos_++];
}
private:
std::vector<std::pair<int64_t, float>> 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<Schema>();
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<SegmentGrowingImpl*>(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<std::shared_ptr<VectorIterator>> iterators{
std::make_shared<FixedVectorIterator>(
std::vector<std::pair<int64_t, float>>{
{0, 0.1F},
{1, 0.2F},
{2, 0.3F},
{3, 0.4F},
{4, 0.5F},
{6, 0.6F},
})};
OpContext op_context;
std::vector<CompositeGroupKey> group_by_values;
std::vector<int64_t> seg_offsets;
std::vector<float> distances;
std::vector<size_t> topk_per_nq_prefix_sum;
std::vector<int32_t> 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<int64_t>{0, 0, 1, 1, 2}));
ASSERT_EQ(element_indices, (std::vector<int32_t>{0, 1, 0, 1, 0}));
ASSERT_EQ(topk_per_nq_prefix_sum, (std::vector<size_t>{0, 5}));
ASSERT_EQ(group_by_values.size(), seg_offsets.size());
ASSERT_EQ(distances, (std::vector<float>{0.1F, 0.2F, 0.4F, 0.5F, 0.6F}));
}
TEST(GroupBY, SearchGroupByNodeKeepsElementIndices) {
int dim = 4;
auto schema = std::make_shared<Schema>();
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<SegmentGrowingImpl*>(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::shared_ptr<VectorIterator>>{
std::make_shared<FixedVectorIterator>(
std::vector<std::pair<int64_t, float>>{
{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<milvus::exec::QueryContext>(
"search_group_by_node_element_level",
growing,
N,
MAX_TIMESTAMP,
0,
0,
query::PlanOptions{false},
std::make_shared<milvus::exec::QueryConfig>(
std::unordered_map<std::string, std::string>{}));
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<milvus::plan::SearchGroupByNode>("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<RowVector>(std::vector<VectorPtr>{});
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<int64_t>{0, 0, 1, 1, 2}));
ASSERT_EQ(grouped.element_indices_, (std::vector<int32_t>{0, 1, 0, 1, 0}));
ASSERT_EQ(grouped.topk_per_nq_prefix_sum_, (std::vector<size_t>{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);
}
}
}
}
+80
View File
@@ -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()
+6
View File
@@ -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 {
+20
View File
@@ -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(),
+52 -5
View File
@@ -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
+97 -28
View File
@@ -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) {
+4 -2
View File
@@ -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())
}
}
+65
View File
@@ -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{
{