feat: impl StructArray -- support group by for element-level search (#47252)

issue: https://github.com/milvus-io/milvus/issues/42148

design doc:
https://github.com/milvus-io/milvus-design-docs/blob/main/design_docs/20260306-struct.md

---------

Signed-off-by: SpadeA <tangchenjie1210@gmail.com>
This commit is contained in:
Spade A
2026-03-15 19:35:24 +08:00
committed by GitHub
parent f163e94ff1
commit 5bcdbd1573
6 changed files with 611 additions and 164 deletions
-11
View File
@@ -317,16 +317,6 @@ class QueryContext : public Context {
return plan_options_;
}
void
set_element_level_query(bool element_level) {
element_level_query_ = element_level;
}
bool
element_level_query() const {
return element_level_query_;
}
void
set_struct_name(const std::string& field_name) {
struct_name_ = field_name;
@@ -417,7 +407,6 @@ class QueryContext : public Context {
query::PlanOptions plan_options_;
bool element_level_query_{false};
std::string struct_name_;
std::shared_ptr<const IArrayOffsets> array_offsets_{nullptr};
int64_t active_element_count_{0}; // Total elements in active documents
@@ -117,7 +117,6 @@ PhyElementFilterBitsNode::GetOutput() {
doc_bitset, doc_bitset_valid, array_offsets.get());
// Step 4: Set query context
query_context_->set_element_level_query(true);
query_context_->set_struct_name(struct_name_);
// Mark that bitset has been converted to element-level
query_context_->set_bitset_is_element_level(true);
@@ -79,6 +79,13 @@ PhySearchGroupByNode::GetOutput() {
auto op_context = query_context_->get_op_context();
auto search_result = query_context_->get_search_result();
search_info_.array_offsets_ = query_context_->get_array_offsets();
if (search_result.element_level_) {
AssertInfo(search_info_.array_offsets_ != nullptr,
"Array offsets not available");
}
if (search_result.vector_iterators_.has_value()) {
AssertInfo(search_result.vector_iterators_.value().size() ==
search_result.total_nq_,
@@ -105,6 +112,11 @@ PhySearchGroupByNode::GetOutput() {
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();
@@ -18,6 +18,7 @@
#include <algorithm>
#include <cstdint>
#include <tuple>
#include <unordered_set>
#include <variant>
#include "common/QueryInfo.h"
@@ -52,14 +53,11 @@ SearchGroupBy(milvus::OpContext* op_ctx,
auto dataGetter =
GetDataGetter<int8_t>(op_ctx, segment, group_by_field_id);
GroupIteratorsByType<int8_t>(iterators,
search_info.topk_,
search_info.group_size_,
search_info.strict_group_size_,
search_info,
dataGetter,
group_by_values,
seg_offsets,
distances,
search_info.metric_type_,
topk_per_nq_prefix_sum);
break;
}
@@ -67,14 +65,11 @@ SearchGroupBy(milvus::OpContext* op_ctx,
auto dataGetter =
GetDataGetter<int16_t>(op_ctx, segment, group_by_field_id);
GroupIteratorsByType<int16_t>(iterators,
search_info.topk_,
search_info.group_size_,
search_info.strict_group_size_,
search_info,
dataGetter,
group_by_values,
seg_offsets,
distances,
search_info.metric_type_,
topk_per_nq_prefix_sum);
break;
}
@@ -82,14 +77,11 @@ SearchGroupBy(milvus::OpContext* op_ctx,
auto dataGetter =
GetDataGetter<int32_t>(op_ctx, segment, group_by_field_id);
GroupIteratorsByType<int32_t>(iterators,
search_info.topk_,
search_info.group_size_,
search_info.strict_group_size_,
search_info,
dataGetter,
group_by_values,
seg_offsets,
distances,
search_info.metric_type_,
topk_per_nq_prefix_sum);
break;
}
@@ -97,14 +89,11 @@ SearchGroupBy(milvus::OpContext* op_ctx,
auto dataGetter =
GetDataGetter<int64_t>(op_ctx, segment, group_by_field_id);
GroupIteratorsByType<int64_t>(iterators,
search_info.topk_,
search_info.group_size_,
search_info.strict_group_size_,
search_info,
dataGetter,
group_by_values,
seg_offsets,
distances,
search_info.metric_type_,
topk_per_nq_prefix_sum);
break;
}
@@ -112,14 +101,11 @@ SearchGroupBy(milvus::OpContext* op_ctx,
auto dataGetter =
GetDataGetter<int64_t>(op_ctx, segment, group_by_field_id);
GroupIteratorsByType<int64_t>(iterators,
search_info.topk_,
search_info.group_size_,
search_info.strict_group_size_,
search_info,
dataGetter,
group_by_values,
seg_offsets,
distances,
search_info.metric_type_,
topk_per_nq_prefix_sum);
break;
}
@@ -127,14 +113,11 @@ SearchGroupBy(milvus::OpContext* op_ctx,
auto dataGetter =
GetDataGetter<bool>(op_ctx, segment, group_by_field_id);
GroupIteratorsByType<bool>(iterators,
search_info.topk_,
search_info.group_size_,
search_info.strict_group_size_,
search_info,
dataGetter,
group_by_values,
seg_offsets,
distances,
search_info.metric_type_,
topk_per_nq_prefix_sum);
break;
}
@@ -142,14 +125,11 @@ SearchGroupBy(milvus::OpContext* op_ctx,
auto dataGetter =
GetDataGetter<std::string>(op_ctx, segment, group_by_field_id);
GroupIteratorsByType<std::string>(iterators,
search_info.topk_,
search_info.group_size_,
search_info.strict_group_size_,
search_info,
dataGetter,
group_by_values,
seg_offsets,
distances,
search_info.metric_type_,
topk_per_nq_prefix_sum);
break;
}
@@ -167,17 +147,13 @@ SearchGroupBy(milvus::OpContext* op_ctx,
search_info.json_path_,
search_info.json_type_,
search_info.strict_cast_);
GroupIteratorsByType<bool>(
iterators,
search_info.topk_,
search_info.group_size_,
search_info.strict_group_size_,
data_getter,
group_by_values,
seg_offsets,
distances,
search_info.metric_type_,
topk_per_nq_prefix_sum);
GroupIteratorsByType<bool>(iterators,
search_info,
data_getter,
group_by_values,
seg_offsets,
distances,
topk_per_nq_prefix_sum);
break;
}
case DataType::INT8: {
@@ -188,17 +164,13 @@ SearchGroupBy(milvus::OpContext* op_ctx,
search_info.json_path_,
search_info.json_type_,
search_info.strict_cast_);
GroupIteratorsByType<int8_t>(
iterators,
search_info.topk_,
search_info.group_size_,
search_info.strict_group_size_,
data_getter,
group_by_values,
seg_offsets,
distances,
search_info.metric_type_,
topk_per_nq_prefix_sum);
GroupIteratorsByType<int8_t>(iterators,
search_info,
data_getter,
group_by_values,
seg_offsets,
distances,
topk_per_nq_prefix_sum);
break;
}
case DataType::INT16: {
@@ -209,17 +181,13 @@ SearchGroupBy(milvus::OpContext* op_ctx,
search_info.json_path_,
search_info.json_type_,
search_info.strict_cast_);
GroupIteratorsByType<int16_t>(
iterators,
search_info.topk_,
search_info.group_size_,
search_info.strict_group_size_,
data_getter,
group_by_values,
seg_offsets,
distances,
search_info.metric_type_,
topk_per_nq_prefix_sum);
GroupIteratorsByType<int16_t>(iterators,
search_info,
data_getter,
group_by_values,
seg_offsets,
distances,
topk_per_nq_prefix_sum);
break;
}
case DataType::INT32: {
@@ -230,17 +198,13 @@ SearchGroupBy(milvus::OpContext* op_ctx,
search_info.json_path_,
search_info.json_type_,
search_info.strict_cast_);
GroupIteratorsByType<int32_t>(
iterators,
search_info.topk_,
search_info.group_size_,
search_info.strict_group_size_,
data_getter,
group_by_values,
seg_offsets,
distances,
search_info.metric_type_,
topk_per_nq_prefix_sum);
GroupIteratorsByType<int32_t>(iterators,
search_info,
data_getter,
group_by_values,
seg_offsets,
distances,
topk_per_nq_prefix_sum);
break;
}
case DataType::INT64: {
@@ -251,17 +215,13 @@ SearchGroupBy(milvus::OpContext* op_ctx,
search_info.json_path_,
search_info.json_type_,
search_info.strict_cast_);
GroupIteratorsByType<int64_t>(
iterators,
search_info.topk_,
search_info.group_size_,
search_info.strict_group_size_,
data_getter,
group_by_values,
seg_offsets,
distances,
search_info.metric_type_,
topk_per_nq_prefix_sum);
GroupIteratorsByType<int64_t>(iterators,
search_info,
data_getter,
group_by_values,
seg_offsets,
distances,
topk_per_nq_prefix_sum);
break;
}
case DataType::VARCHAR: {
@@ -275,14 +235,11 @@ SearchGroupBy(milvus::OpContext* op_ctx,
search_info.strict_cast_);
GroupIteratorsByType<std::string>(
iterators,
search_info.topk_,
search_info.group_size_,
search_info.strict_group_size_,
search_info,
data_getter,
group_by_values,
seg_offsets,
distances,
search_info.metric_type_,
topk_per_nq_prefix_sum);
break;
}
@@ -301,17 +258,13 @@ SearchGroupBy(milvus::OpContext* op_ctx,
search_info.json_path_,
search_info.json_type_,
search_info.strict_cast_);
GroupIteratorsByType<std::string>(
iterators,
search_info.topk_,
search_info.group_size_,
search_info.strict_group_size_,
data_getter,
group_by_values,
seg_offsets,
distances,
search_info.metric_type_,
topk_per_nq_prefix_sum);
GroupIteratorsByType<std::string>(iterators,
search_info,
data_getter,
group_by_values,
seg_offsets,
distances,
topk_per_nq_prefix_sum);
}
break;
}
@@ -328,26 +281,20 @@ template <typename T>
void
GroupIteratorsByType(
const std::vector<std::shared_ptr<VectorIterator>>& iterators,
int64_t topK,
int64_t group_size,
bool strict_group_size,
const SearchInfo& search_info,
const std::shared_ptr<DataGetter<T>>& data_getter,
std::vector<GroupByValueType>& group_by_values,
std::vector<int64_t>& seg_offsets,
std::vector<float>& distances,
const knowhere::MetricType& metrics_type,
std::vector<size_t>& topk_per_nq_prefix_sum) {
topk_per_nq_prefix_sum.push_back(0);
for (auto& iterator : iterators) {
GroupIteratorResult<T>(iterator,
topK,
group_size,
strict_group_size,
search_info,
data_getter,
group_by_values,
seg_offsets,
distances,
metrics_type);
distances);
topk_per_nq_prefix_sum.push_back(seg_offsets.size());
}
}
@@ -355,21 +302,27 @@ GroupIteratorsByType(
template <typename T>
void
GroupIteratorResult(const std::shared_ptr<VectorIterator>& iterator,
int64_t topK,
int64_t group_size,
bool strict_group_size,
const SearchInfo& search_info,
const std::shared_ptr<DataGetter<T>>& data_getter,
std::vector<GroupByValueType>& group_by_values,
std::vector<int64_t>& offsets,
std::vector<float>& distances,
const knowhere::MetricType& metrics_type) {
std::vector<float>& distances) {
//1.
GroupByMap<T> groupMap(topK, group_size, strict_group_size);
GroupByMap<T> groupMap(search_info.topk_,
search_info.group_size_,
search_info.strict_group_size_);
auto is_element_id = search_info.element_level();
auto array_offsets = search_info.array_offsets_;
//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, std::optional<T>>> res;
// For element-level search, multiple elements from the same row may be
// returned by the iterator. We must deduplicate by row_offset so that
// the same document does not consume multiple slots in a group.
std::unordered_set<int64_t> seen_rows;
while (iterator->HasNext() && !groupMap.IsGroupResEnough()) {
auto offset_dis_pair = iterator->Next();
AssertInfo(
@@ -378,16 +331,30 @@ GroupIteratorResult(const std::shared_ptr<VectorIterator>& iterator,
"tells hasNext, terminate groupBy operation");
auto offset = offset_dis_pair.value().first;
auto dis = offset_dis_pair.value().second;
std::optional<T> row_data = data_getter->Get(offset);
// When array_offsets is present, iterator returns element_id,
// but DataGetter expects row_id. Convert element_id to row_id.
int64_t row_offset = offset;
if (is_element_id) {
row_offset =
array_offsets->ElementIDToRowID(static_cast<int32_t>(offset))
.first;
// Skip if we already saw this row (closest element wins)
if (!seen_rows.insert(row_offset).second) {
continue;
}
}
std::optional<T> row_data = data_getter->Get(row_offset);
if (groupMap.Push(row_data)) {
res.emplace_back(offset, dis, row_data);
res.emplace_back(row_offset, dis, row_data);
}
}
//3. sorted 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), metrics_type);
std::get<1>(lhs), std::get<1>(rhs), search_info.metric_type_);
};
std::sort(res.begin(), res.end(), customComparator);
@@ -343,14 +343,11 @@ template <typename T>
void
GroupIteratorsByType(
const std::vector<std::shared_ptr<VectorIterator>>& iterators,
int64_t topK,
int64_t group_size,
bool strict_group_size,
const SearchInfo& search_info,
const std::shared_ptr<DataGetter<T>>& data_getter,
std::vector<GroupByValueType>& group_by_values,
std::vector<int64_t>& seg_offsets,
std::vector<float>& distances,
const knowhere::MetricType& metrics_type,
std::vector<size_t>& topk_per_nq_prefix_sum);
template <typename T>
@@ -413,13 +410,11 @@ struct GroupByMap {
template <typename T>
void
GroupIteratorResult(const std::shared_ptr<VectorIterator>& iterator,
int64_t topK,
int64_t group_size,
bool strict_group_size,
const SearchInfo& search_info,
const std::shared_ptr<DataGetter<T>>& data_getter,
std::vector<GroupByValueType>& group_by_values,
std::vector<int64_t>& offsets,
std::vector<float>& distances,
const knowhere::MetricType& metrics_type);
std::vector<float>& distances);
} // namespace exec
} // namespace milvus
+512 -27
View File
@@ -13,12 +13,12 @@
#include <gtest/gtest.h>
#include <stddef.h>
#include <cstdint>
#include <iostream>
#include <map>
#include <memory>
#include <optional>
#include <string>
#include <tuple>
#include <unordered_set>
#include <utility>
#include <vector>
@@ -321,10 +321,6 @@ TEST_P(ElementFilterSealed, RangeExpr) {
for (size_t i = 0; i < search_result->seg_offsets_.size(); i++) {
int64_t doc_id = search_result->seg_offsets_[i];
int32_t elem_idx = search_result->element_indices_[i];
float distance = search_result->distances_[i];
std::cout << "doc_id: " << doc_id << ", element_index: " << elem_idx
<< ", distance: " << distance << std::endl;
// Verify the doc_id satisfies the predicate (id % 2 == 0) only
// when predicate is enabled
@@ -544,16 +540,6 @@ TEST_P(ElementFilterSealed, UnaryExpr) {
ASSERT_LE(search_result->element_indices_.size(),
static_cast<size_t>(topK * num_queries));
std::cout << "Element-level search returned ("
<< static_cast<int>(elem_type) << "):" << std::endl;
for (size_t i = 0; i < search_result->seg_offsets_.size(); i++) {
std::cout << "doc_id: " << search_result->seg_offsets_[i]
<< ", element_index: "
<< search_result->element_indices_[i]
<< ", distance: " << search_result->distances_[i]
<< std::endl;
}
// Verify distances are sorted
for (size_t i = 1; i < search_result->distances_.size(); ++i) {
ASSERT_LE(search_result->distances_[i - 1],
@@ -1227,16 +1213,9 @@ TEST_P(ElementFilterGrowing, RangeExpr) {
static_cast<size_t>(topK * num_queries))
<< "Should not exceed topK results";
std::cout << "Growing segment element-level search ("
<< static_cast<int>(elem_type) << "):" << std::endl;
for (size_t i = 0; i < search_result->seg_offsets_.size(); i++) {
int64_t doc_id = search_result->seg_offsets_[i];
int32_t elem_idx = search_result->element_indices_[i];
float distance = search_result->distances_[i];
std::cout << " [" << i << "] doc_id=" << doc_id
<< ", element_index=" << elem_idx
<< ", distance=" << distance << std::endl;
// Verify the doc_id satisfies the predicate (id % 2 == 0) only
// when predicate is enabled
@@ -2131,11 +2110,6 @@ TEST_P(ElementFilterNestedIndex, ExecutionMode) {
last_distance = search_result->distances_[i];
}
}
std::cout << "Test passed: index_type="
<< NestedIndexTypeToString(index_type)
<< ", offset_mode=" << (offset_mode ? "true" : "false")
<< ", valid_results=" << valid_count << std::endl;
}
INSTANTIATE_TEST_SUITE_P(
@@ -2154,3 +2128,514 @@ INSTANTIATE_TEST_SUITE_P(
name += offset_mode ? "_OffsetMode" : "_FullMode";
return name;
});
// Test element-level filter combined with group by on sealed segment with index
TEST(ElementFilterGroupBy, SealedWithIndex) {
int dim = 4;
size_t N = 500;
int array_len = 3;
auto schema = std::make_shared<Schema>();
auto vec_fid = schema->AddDebugVectorArrayField("structA[array_vec]",
DataType::VECTOR_FLOAT,
dim,
knowhere::metric::L2);
auto int_array_fid = schema->AddDebugArrayField(
"structA[price_array]", DataType::INT32, false);
auto int64_fid = schema->AddDebugField("id", DataType::INT64);
schema->set_primary_field_id(int64_fid);
// Generate test data
auto raw_data = DataGen(schema, N, 42, 0, 1, array_len);
// Customize int_array data: doc i has elements [i*3+1, i*3+2, i*3+3]
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() == int_array_fid.get()) {
field_data->mutable_scalars()
->mutable_array_data()
->mutable_data()
->Clear();
for (int row = 0; row < N; row++) {
auto* array_data = field_data->mutable_scalars()
->mutable_array_data()
->mutable_data()
->Add();
for (int elem = 0; elem < array_len; elem++) {
int value = row * array_len + elem + 1;
array_data->mutable_int_data()->mutable_data()->Add(value);
}
}
break;
}
}
// Create sealed segment and load field data
auto segment = CreateSealedWithFieldDataLoaded(schema, raw_data);
// Build and load vector index
auto array_vec_values = raw_data.get_col<VectorFieldProto>(vec_fid);
std::vector<float> vector_data(dim * N * array_len);
for (int i = 0; i < N; i++) {
const auto& float_vec = array_vec_values[i].float_vector().data();
for (int j = 0; j < array_len * dim; j++) {
vector_data[i * array_len * dim + j] = float_vec[j];
}
}
auto indexing = GenVecIndexing(N * array_len,
dim,
vector_data.data(),
knowhere::IndexEnum::INDEX_HNSW);
LoadIndexInfo load_index_info;
load_index_info.field_id = vec_fid.get();
load_index_info.index_params = GenIndexParams(indexing.get());
load_index_info.cache_index =
CreateTestCacheIndex("test", std::move(indexing));
load_index_info.index_params["metric_type"] = knowhere::metric::L2;
load_index_info.field_type = DataType::VECTOR_ARRAY;
load_index_info.element_type = DataType::VECTOR_FLOAT;
segment->LoadIndex(load_index_info);
int topK = 5;
int group_size = 4;
// Execute element-level search with group by primary key
ScopedSchemaHandle handle(*schema);
std::string expr =
"id % 2 == 0 && element_filter(structA, 400 > $[price_array] > 100)";
std::string search_params = R"({"ef": 50})";
auto plan_bytes = handle.ParseGroupBySearch(expr,
"structA[array_vec]",
topK,
knowhere::metric::L2,
search_params,
int64_fid.get(),
group_size);
auto plan =
CreateSearchPlanByExpr(schema, plan_bytes.data(), plan_bytes.size());
ASSERT_NE(plan, nullptr);
auto num_queries = 1;
auto ph_group_raw = CreatePlaceholderGroup(num_queries, dim, 1024, true);
auto ph_group =
ParsePlaceholderGroup(plan.get(), ph_group_raw.SerializeAsString());
auto search_result = segment->Search(plan.get(), ph_group.get(), 1L << 63);
ASSERT_NE(search_result, nullptr);
ASSERT_TRUE(search_result->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";
auto& group_by_values = search_result->group_by_values_.value();
ASSERT_LE(search_result->seg_offsets_.size(),
static_cast<size_t>(topK * group_size))
<< "Should not exceed topK * group_size results";
std::unordered_map<int64_t, int> group_counts;
for (size_t i = 0; i < search_result->seg_offsets_.size(); i++) {
int64_t doc_id = search_result->seg_offsets_[i];
if (i < group_by_values.size() && group_by_values[i].has_value()) {
if (std::holds_alternative<int64_t>(group_by_values[i].value())) {
int64_t group_val =
std::get<int64_t>(group_by_values[i].value());
ASSERT_EQ(group_val, doc_id)
<< "Group by primary key: group value should equal doc_id";
group_counts[group_val]++;
ASSERT_LE(group_counts[group_val], group_size)
<< "Each group should have at most group_size results";
}
}
ASSERT_EQ(doc_id % 2, 0)
<< "Result doc_id " << doc_id << " should satisfy (id % 2 == 0)";
}
ASSERT_LE(group_counts.size(), static_cast<size_t>(topK))
<< "Should have at most topK distinct groups";
for (size_t i = 1; i < search_result->distances_.size(); ++i) {
ASSERT_LE(search_result->distances_[i - 1],
search_result->distances_[i])
<< "Distances should be sorted in ascending order";
}
}
// Test: element level + group by + growing segment
TEST(ElementFilterGroupBy, GrowingSegment) {
int dim = 4;
auto schema = std::make_shared<Schema>();
auto vec_fid = schema->AddDebugVectorArrayField(
"structA[array_vec]", DataType::VECTOR_FLOAT, dim, "L2");
auto int_array_fid = schema->AddDebugArrayField(
"structA[price_array]", DataType::INT32, false);
auto int64_fid = schema->AddDebugField("id", DataType::INT64);
schema->set_primary_field_id(int64_fid);
size_t N = 500;
int array_len = 3;
auto raw_data = DataGen(schema, N, 42, 0, 1, array_len);
// Customize int_array data: doc i has elements [i*3+1, i*3+2, i*3+3]
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() == int_array_fid.get()) {
field_data->mutable_scalars()
->mutable_array_data()
->mutable_data()
->Clear();
for (int row = 0; row < N; row++) {
auto* array_data = field_data->mutable_scalars()
->mutable_array_data()
->mutable_data()
->Add();
for (int elem = 0; elem < array_len; elem++) {
int value = row * array_len + elem + 1;
array_data->mutable_int_data()->mutable_data()->Add(value);
}
}
break;
}
}
// Create growing segment and insert data
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_);
int topK = 5;
int group_size = 4;
// Execute element-level search with group by primary key
ScopedSchemaHandle handle(*schema);
std::string expr =
"id % 2 == 0 && element_filter(structA, 400 > $[price_array] > 100)";
std::string search_params = R"({"nprobe": 10})";
auto plan_bytes = handle.ParseGroupBySearch(expr,
"structA[array_vec]",
topK,
knowhere::metric::L2,
search_params,
int64_fid.get(),
group_size);
auto plan =
CreateSearchPlanByExpr(schema, plan_bytes.data(), plan_bytes.size());
ASSERT_NE(plan, nullptr);
auto num_queries = 1;
auto ph_group_raw = CreatePlaceholderGroup(num_queries, dim, 1024, true);
auto ph_group =
ParsePlaceholderGroup(plan.get(), ph_group_raw.SerializeAsString());
auto search_result = segment->Search(plan.get(), ph_group.get(), 1L << 63);
ASSERT_NE(search_result, nullptr);
ASSERT_TRUE(search_result->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";
auto& group_by_values = search_result->group_by_values_.value();
ASSERT_FALSE(search_result->seg_offsets_.empty())
<< "Should have search results";
ASSERT_LE(search_result->seg_offsets_.size(),
static_cast<size_t>(topK * group_size))
<< "Should not exceed topK * group_size results";
for (size_t i = 0; i < search_result->seg_offsets_.size(); i++) {
int64_t doc_id = search_result->seg_offsets_[i];
// Verify row-level filter
ASSERT_EQ(doc_id % 2, 0)
<< "Result doc_id " << doc_id << " should satisfy (id % 2 == 0)";
// Verify element-level filter: 100 < element_value < 400
// element_value = doc_id * array_len + elem + 1, elem in [0, array_len)
// For doc to have any matching element: doc_id * 3 + 1 < 400 and doc_id * 3 + 3 > 100
// So doc_id should be roughly in range [33, 132]
ASSERT_GE(doc_id, 33)
<< "doc_id " << doc_id << " should be >= 33 (element filter)";
ASSERT_LE(doc_id, 133)
<< "doc_id " << doc_id << " should be <= 133 (element filter)";
if (i < group_by_values.size() && group_by_values[i].has_value()) {
if (std::holds_alternative<int64_t>(group_by_values[i].value())) {
int64_t group_val =
std::get<int64_t>(group_by_values[i].value());
ASSERT_EQ(group_val, doc_id)
<< "Group by primary key: group value should equal doc_id";
}
}
}
}
// Test: normal group by (without element level)
TEST(ElementFilterGroupBy, NormalGroupBy) {
int dim = 4;
auto schema = std::make_shared<Schema>();
auto vec_fid =
schema->AddDebugField("vec", DataType::VECTOR_FLOAT, dim, "L2");
auto int64_fid = schema->AddDebugField("id", DataType::INT64);
auto category_fid = schema->AddDebugField("category", DataType::INT32);
schema->set_primary_field_id(int64_fid);
size_t N = 500;
auto raw_data = DataGen(schema, N, 42);
auto segment = CreateSealedWithFieldDataLoaded(schema, raw_data);
// Build vector index
auto vec_values = raw_data.get_col<float>(vec_fid);
auto indexing = GenVecIndexing(
N, dim, vec_values.data(), knowhere::IndexEnum::INDEX_HNSW);
LoadIndexInfo load_index_info;
load_index_info.field_id = vec_fid.get();
load_index_info.index_params = GenIndexParams(indexing.get());
load_index_info.cache_index =
CreateTestCacheIndex("test", std::move(indexing));
load_index_info.index_params["metric_type"] = knowhere::metric::L2;
load_index_info.field_type = DataType::VECTOR_FLOAT;
segment->LoadIndex(load_index_info);
int topK = 5;
int group_size = 2;
// Normal search with group by (no element level)
ScopedSchemaHandle handle(*schema);
std::string expr = "id % 2 == 0";
std::string search_params = R"({"ef": 50})";
auto plan_bytes = handle.ParseGroupBySearch(expr,
"vec",
topK,
knowhere::metric::L2,
search_params,
category_fid.get(),
group_size);
auto plan =
CreateSearchPlanByExpr(schema, plan_bytes.data(), plan_bytes.size());
ASSERT_NE(plan, nullptr);
auto num_queries = 1;
auto ph_group_raw = CreatePlaceholderGroup(num_queries, dim, 1024);
auto ph_group =
ParsePlaceholderGroup(plan.get(), ph_group_raw.SerializeAsString());
auto search_result = segment->Search(plan.get(), ph_group.get(), 1L << 63);
ASSERT_NE(search_result, nullptr);
ASSERT_TRUE(search_result->group_by_values_.has_value())
<< "Group by values should be present";
// Verify element_level_ is false
ASSERT_FALSE(search_result->element_level_)
<< "Normal search should not be element-level";
auto& group_by_values = search_result->group_by_values_.value();
ASSERT_FALSE(search_result->seg_offsets_.empty())
<< "Should have search results";
ASSERT_LE(search_result->seg_offsets_.size(),
static_cast<size_t>(topK * group_size))
<< "Should not exceed topK * group_size results";
std::unordered_map<int32_t, int> group_counts;
for (size_t i = 0; i < search_result->seg_offsets_.size(); i++) {
int64_t doc_id = search_result->seg_offsets_[i];
ASSERT_EQ(doc_id % 2, 0)
<< "Result doc_id " << doc_id << " should satisfy (id % 2 == 0)";
if (i < group_by_values.size() && group_by_values[i].has_value()) {
if (std::holds_alternative<int32_t>(group_by_values[i].value())) {
int32_t group_val =
std::get<int32_t>(group_by_values[i].value());
group_counts[group_val]++;
ASSERT_LE(group_counts[group_val], group_size)
<< "Each group should have at most group_size results";
}
}
}
ASSERT_LE(group_counts.size(), static_cast<size_t>(topK))
<< "Should have at most topK distinct groups";
}
// Test: element-level + group by non-unique field to verify row deduplication.
// When multiple elements from the same document are close to the query vector,
// each document should only appear once in the results (not consume multiple
// group slots).
TEST(ElementFilterGroupBy, DeduplicateRowsInGroup) {
int dim = 4;
size_t N = 500;
int array_len = 3;
auto schema = std::make_shared<Schema>();
auto vec_fid = schema->AddDebugVectorArrayField("structA[array_vec]",
DataType::VECTOR_FLOAT,
dim,
knowhere::metric::L2);
auto int_array_fid = schema->AddDebugArrayField(
"structA[price_array]", DataType::INT32, false);
auto int64_fid = schema->AddDebugField("id", DataType::INT64);
// category field: non-unique, used for group by
// Assign category = id % 10, so ~50 docs per category
auto category_fid = schema->AddDebugField("category", DataType::INT32);
schema->set_primary_field_id(int64_fid);
auto raw_data = DataGen(schema, N, 42, 0, 1, array_len);
// Customize int_array data: doc i has elements [i*3+1, i*3+2, i*3+3]
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() == int_array_fid.get()) {
field_data->mutable_scalars()
->mutable_array_data()
->mutable_data()
->Clear();
for (int row = 0; row < N; row++) {
auto* array_data = field_data->mutable_scalars()
->mutable_array_data()
->mutable_data()
->Add();
for (int elem = 0; elem < array_len; elem++) {
int value = row * array_len + elem + 1;
array_data->mutable_int_data()->mutable_data()->Add(value);
}
}
}
// Set category = id % 10
if (field_data->field_id() == category_fid.get()) {
field_data->mutable_scalars()
->mutable_int_data()
->mutable_data()
->Clear();
for (int row = 0; row < N; row++) {
field_data->mutable_scalars()
->mutable_int_data()
->mutable_data()
->Add(row % 10);
}
}
}
auto segment = CreateSealedWithFieldDataLoaded(schema, raw_data);
// Build and load vector index
auto array_vec_values = raw_data.get_col<VectorFieldProto>(vec_fid);
std::vector<float> vector_data(dim * N * array_len);
for (int i = 0; i < N; i++) {
const auto& float_vec = array_vec_values[i].float_vector().data();
for (int j = 0; j < array_len * dim; j++) {
vector_data[i * array_len * dim + j] = float_vec[j];
}
}
auto indexing = GenVecIndexing(N * array_len,
dim,
vector_data.data(),
knowhere::IndexEnum::INDEX_HNSW);
LoadIndexInfo load_index_info;
load_index_info.field_id = vec_fid.get();
load_index_info.index_params = GenIndexParams(indexing.get());
load_index_info.cache_index =
CreateTestCacheIndex("test", std::move(indexing));
load_index_info.index_params["metric_type"] = knowhere::metric::L2;
load_index_info.field_type = DataType::VECTOR_ARRAY;
load_index_info.element_type = DataType::VECTOR_FLOAT;
segment->LoadIndex(load_index_info);
int topK = 5;
int group_size = 3;
ScopedSchemaHandle handle(*schema);
std::string expr = "element_filter(structA, 2000 > $[price_array] > 100)";
std::string search_params = R"({"ef": 50})";
auto plan_bytes = handle.ParseGroupBySearch(expr,
"structA[array_vec]",
topK,
knowhere::metric::L2,
search_params,
category_fid.get(),
group_size);
auto plan =
CreateSearchPlanByExpr(schema, plan_bytes.data(), plan_bytes.size());
ASSERT_NE(plan, nullptr);
auto num_queries = 1;
auto ph_group_raw = CreatePlaceholderGroup(num_queries, dim, 1024, true);
auto ph_group =
ParsePlaceholderGroup(plan.get(), ph_group_raw.SerializeAsString());
auto search_result = segment->Search(plan.get(), ph_group.get(), 1L << 63);
ASSERT_NE(search_result, nullptr);
ASSERT_TRUE(search_result->group_by_values_.has_value())
<< "Group by values should be present";
ASSERT_FALSE(search_result->element_level_)
<< "Group by should return row-level results";
auto& group_by_values = search_result->group_by_values_.value();
ASSERT_FALSE(search_result->seg_offsets_.empty())
<< "Should have search results";
ASSERT_LE(search_result->seg_offsets_.size(),
static_cast<size_t>(topK * group_size))
<< "Should not exceed topK * group_size results";
// Verify: no duplicate row_offsets in results.
// Without deduplication, the same document could appear multiple times
// because different elements from one document may all be close to the
// query vector.
std::unordered_set<int64_t> seen_docs;
std::unordered_map<int32_t, int> group_counts;
for (size_t i = 0; i < search_result->seg_offsets_.size(); i++) {
int64_t doc_id = search_result->seg_offsets_[i];
ASSERT_TRUE(seen_docs.insert(doc_id).second)
<< "Duplicate doc_id " << doc_id
<< " in results: same document should not appear more than once";
if (i < group_by_values.size() && group_by_values[i].has_value()) {
if (std::holds_alternative<int32_t>(group_by_values[i].value())) {
int32_t group_val =
std::get<int32_t>(group_by_values[i].value());
ASSERT_EQ(group_val, doc_id % 10)
<< "Group value should equal category (doc_id % 10)";
group_counts[group_val]++;
ASSERT_LE(group_counts[group_val], group_size)
<< "Each group should have at most group_size results";
}
}
}
ASSERT_LE(group_counts.size(), static_cast<size_t>(topK))
<< "Should have at most topK distinct groups";
}