mirror of
https://github.com/milvus-io/milvus.git
synced 2026-07-21 02:05:41 +00:00
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:
@@ -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";
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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);
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -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);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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()
|
||||
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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(),
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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) {
|
||||
|
||||
@@ -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())
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -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{
|
||||
{
|
||||
|
||||
Reference in New Issue
Block a user