diff --git a/src/algorithm/hgraph/hgraph.h b/src/algorithm/hgraph/hgraph.h index 00d005455a..8515cdfa4e 100644 --- a/src/algorithm/hgraph/hgraph.h +++ b/src/algorithm/hgraph/hgraph.h @@ -58,6 +58,7 @@ namespace vsag { class FlattenOptimizedBuildInterface; class HGraphOptimizedBuildSession; class IteratorFilterContext; +class ReasoningContext; /** * @brief HGraph: hierarchical navigable graph index. @@ -805,6 +806,56 @@ class HGraph : public InnerIndexInterface { bool used_precise_float_csr{false}; }; + [[nodiscard]] QueryContext + create_query_context(const SearchRequest& request, + const HGraphSearchParameters& params, + int64_t k, + bool use_custom_distance, + SearchStatistics* stats, + std::shared_ptr& reasoning_ctx) const; + + void + search_route_graphs(const SearchRequest& request, + const HGraphSearchParameters& params, + InnerIdType entry_point, + bool use_custom_distance, + const void* query, + const VisitedListPtr& visited_list, + QueryContext* ctx, + InnerSearchParam& search_param) const; + + static void + configure_bottom_graph_search(const SearchRequest& request, + const HGraphSearchParameters& params, + bool is_range, + int64_t k, + bool use_custom_distance, + const FilterPtr& filter, + const std::optional& threshold, + QueryContext* ctx, + InnerSearchParam& search_param); + + [[nodiscard]] DatasetPtr + pack_search_result(const SearchRequest& request, + int64_t k, + DistHeapPtr search_result, + const QueryContext& ctx, + const MCIHybridSearchResult& mci_result, + const SearchStatistics& stats, + const std::shared_ptr& reasoning_ctx) const; + + [[nodiscard]] HGraphSearchParameters + parse_and_validate_search_params(const SearchRequest& request, + bool is_range, + int64_t k, + bool use_custom_distance) const; + + [[nodiscard]] std::shared_ptr + initialize_reasoning_context(const SearchRequest& request, + int64_t k, + bool use_custom_distance, + QueryContext* ctx) const; + [[nodiscard]] MCIHybridSearchResult try_mci_search(const SearchRequest& request, const HGraphSearchParameters& params, diff --git a/src/algorithm/hgraph/hgraph_mci_test.cpp b/src/algorithm/hgraph/hgraph_mci_test.cpp index d748ab7105..b6d0cc31b8 100644 --- a/src/algorithm/hgraph/hgraph_mci_test.cpp +++ b/src/algorithm/hgraph/hgraph_mci_test.cpp @@ -17,6 +17,7 @@ #include #include #include +#include #include #include #include @@ -83,6 +84,18 @@ class CountingValidIdsFilter : public HalfRatioAllValidFilter { mutable std::atomic check_count_{0}; }; +class NaNRatioAllValidFilter : public HalfRatioAllValidFilter { +public: + explicit NaNRatioAllValidFilter(const std::vector& ids) + : HalfRatioAllValidFilter(ids) { + } + + float + ValidRatio() const override { + return std::numeric_limits::quiet_NaN(); + } +}; + class CallbackOnlyFilter : public vsag::Filter { public: bool @@ -281,6 +294,43 @@ TEST_CASE("HGraph companion MCI incrementally updates cliques after Add", "[ut][ REQUIRE(std::stoull(result.value()->GetStatistics({"mci_seed_count"})[0]) == expected_seed_count); + result = index.value()->KnnSearch( + query, + 3, + R"({"hgraph":{"ef_search":16,"use_mci":true,"mci_seed_ratio":0.5,)" + R"("hgraph_valid_ratio_threshold":1.0,"brute_force_threshold":0.3}})", + filter); + REQUIRE(result.has_value()); + REQUIRE(result.value()->GetStatistics({"mci_hybrid_route"})[0] == R"("mci")"); + + auto nan_ratio_filter = std::make_shared(ids); + result = + index.value()->KnnSearch(query, + 3, + R"({"hgraph":{"ef_search":16,"use_mci":true,"mci_seed_ratio":0.5,)" + R"("hgraph_valid_ratio_threshold":1.0}})", + nan_ratio_filter); + REQUIRE(result.has_value()); + REQUIRE(result.value()->GetStatistics({"mci_hybrid_route"})[0] == R"("mci")"); + + result = index.value()->KnnSearch( + query, + 3, + R"({"hgraph":{"ef_search":16,"use_mci":true,"mci_seed_ratio":0.5,)" + R"("hgraph_valid_ratio_threshold":1.0,"brute_force_threshold":0.5}})", + filter); + REQUIRE(result.has_value()); + REQUIRE(result.value()->GetStatistics({"mci_hybrid_route"})[0] == R"("brute_force")"); + + result = index.value()->RangeSearch( + query, + std::numeric_limits::max(), + R"({"hgraph":{"ef_search":16,"use_mci":true,"mci_seed_ratio":0.5,)" + R"("hgraph_valid_ratio_threshold":1.0}})", + filter); + REQUIRE(result.has_value()); + REQUIRE(result.value()->GetStatistics({"mci_hybrid_route"})[0] == R"("hgraph")"); + result = index.value()->KnnSearch(query, 3, diff --git a/src/algorithm/hgraph/hgraph_search.cpp b/src/algorithm/hgraph/hgraph_search.cpp index 6a9eeeb6cc..100a5a77d1 100644 --- a/src/algorithm/hgraph/hgraph_search.cpp +++ b/src/algorithm/hgraph/hgraph_search.cpp @@ -20,6 +20,7 @@ #include "attr/argparse.h" #include "dataset_impl.h" #include "hgraph.h" // IWYU pragma: keep +#include "impl/filter/black_list_filter.h" #include "impl/filter/iterator_filter.h" #include "impl/heap/standard_heap.h" #include "impl/reasoning/search_reasoning.h" @@ -28,6 +29,97 @@ namespace vsag { +enum class search_plan { + K_BOTTOM_GRAPH, + K_BRUTE_FORCE, + K_MCI, +}; + +struct search_plan_input { + bool is_range_search{false}; + float valid_ratio{1.0F}; + float brute_force_threshold{0.0F}; + bool use_mci{false}; + bool mci_enabled{false}; + bool mci_has_clique_index{false}; + bool has_attribute_executor{false}; + bool has_valid_id_source{false}; + bool has_bitset_source{false}; + float mci_valid_ratio_threshold{0.0F}; +}; + +static bool +is_mci_available(const search_plan_input& input) { + return input.use_mci and input.mci_enabled and input.mci_has_clique_index and + not input.has_attribute_executor and + (input.has_valid_id_source or input.has_bitset_source); +} + +static search_plan +select_search_plan(const search_plan_input& input, bool mci_available) { + if (input.brute_force_threshold > 0.0F and input.valid_ratio <= input.brute_force_threshold) { + return search_plan::K_BRUTE_FORCE; + } + // A NaN ratio preserves the existing MCI attempt behavior. + if (not input.is_range_search and mci_available and + not(input.valid_ratio >= input.mci_valid_ratio_threshold)) { + return search_plan::K_MCI; + } + return search_plan::K_BOTTOM_GRAPH; +} + +static bool +has_valid_id_source(const FilterPtr& filter) { + if (filter == nullptr) { + return false; + } + const int64_t* valid_ids = nullptr; + int64_t valid_count = 0; + filter->GetValidIds(&valid_ids, valid_count); + return valid_ids != nullptr and valid_count > 0; +} + +static bool +has_bitset_source(const FilterPtr& filter) { + if (filter == nullptr) { + return false; + } + const auto bitset_filter = std::dynamic_pointer_cast(filter); + return bitset_filter != nullptr and bitset_filter->IsBitsetFilter(); +} + +class VisitedListGuard { +public: + VisitedListGuard(std::shared_ptr pool, VisitedListPtr visited_list) + : pool_(std::move(pool)), visited_list_(std::move(visited_list)) { + } + + VisitedListGuard(const VisitedListGuard&) = delete; + VisitedListGuard& + operator=(const VisitedListGuard&) = delete; + + VisitedListPtr& + Get() { + return visited_list_; + } + + void + Release() { + if (visited_list_ != nullptr) { + pool_->ReturnOne(visited_list_); + visited_list_.reset(); + } + } + + ~VisitedListGuard() { + Release(); + } + +private: + std::shared_ptr pool_; + VisitedListPtr visited_list_; +}; + static DatasetPtr make_empty_dataset_with_stats(const SearchStatistics& stats) { auto dataset_result = DatasetImpl::MakeEmptyDataset(); @@ -147,7 +239,7 @@ HGraph::KnnSearch(const DatasetPtr& query, } if (iter_filter_ctx->IsFirstUsed()) { ScopedDistancePhase routing_phase(ctx, DistanceEvaluationPhase::ROUTING); - for (auto i = static_cast(this->route_graphs_.size() - 1); i >= 0; --i) { + for (auto i = static_cast(this->route_graphs_.size()) - 1; i >= 0; --i) { auto result = this->search_one_graph(query_data, this->route_graphs_[i], this->basic_flatten_codes_, @@ -415,116 +507,37 @@ HGraph::RangeSearch(const DatasetPtr& query, return this->SearchWithRequest(req); } -[[nodiscard]] DatasetPtr -HGraph::SearchWithRequest(const SearchRequest& request) const { - ValidateSearchThreshold(request.threshold_); - SearchStatistics stats; - QueryContext ctx{.alloc = this->allocator_, .stats = &stats}; +QueryContext +HGraph::create_query_context(const SearchRequest& request, + const HGraphSearchParameters& params, + int64_t k, + bool use_custom_distance, + SearchStatistics* stats, + std::shared_ptr& reasoning_ctx) const { + QueryContext ctx; + ctx.alloc = this->allocator_; + ctx.stats = stats; + ctx.rabitq_error_rate = params.rabitq_error_rate; if (request.search_allocator_ != nullptr) { ctx.alloc = request.search_allocator_; } - - const auto& query = request.query_; - bool is_range = (request.mode_ == SearchMode::RANGE_SEARCH); - auto k = request.topk_; - const bool use_custom_distance = request.distance_batch_func_ != nullptr; - - if (use_custom_distance) { - CHECK_ARGUMENT(request.distance_batch_size_ > 0, - "distance_batch_size must be greater than 0"); - CHECK_ARGUMENT(not is_range, "HGraph custom distance only supports KNN search"); - } - - if (is_range) { - if (not use_custom_distance) { - this->validate_range_args(query, request.radius_, request.limited_size_); - } - } else { - if (not use_custom_distance) { - this->validate_knn_args(query, k); - } else { - CHECK_ARGUMENT(k > 0, "topk must be greater than 0"); - } - } - - auto params = HGraphSearchParameters::FromJson(request.params_str_); - ctx.rabitq_error_rate = params.rabitq_error_rate; - - if (use_custom_distance) { - CHECK_ARGUMENT(params.parallel_search_thread_count == 1, - "HGraph custom query distance does not support parallel search"); - CHECK_ARGUMENT(params.brute_force_threshold <= 0.0F, - "HGraph custom query distance does not support brute_force_threshold"); - } - - CHECK_ARGUMENT( // NOLINT - params.ef_search >= 1, - fmt::format("ef_search({}) must be at least 1", params.ef_search)); - - std::shared_lock force_remove_rlock; - std::shared_lock shared_lock; - if (!this->immutable_.load(std::memory_order_acquire)) { - if (this->support_force_remove()) { - force_remove_rlock = std::shared_lock(this->force_remove_mutex_); - } - shared_lock = this->acquire_global_read_lock(); - } - const auto element_count = GetNumElements(); - if (element_count == 0) { - return make_empty_dataset_with_stats(); - } - k = std::min(k, element_count); - - // Setup reasoning context (KNN only) - std::shared_ptr reasoning_ctx; - if (not is_range and not request.expected_labels_.empty()) { - reasoning_ctx = std::make_shared(this->allocator_); - reasoning_ctx->SetSearchParams( - k, "HGraph", use_custom_distance ? false : use_reorder_, request.filter_ != nullptr); - - UnorderedMap label_to_inner_id(this->allocator_); - for (const auto& label : request.expected_labels_) { - auto [success, inner_id] = label_table_->TryGetIdByLabel(label, true); - if (success) { - label_to_inner_id[label] = inner_id; - } - } - - Vector expected_labels_vec( - request.expected_labels_.begin(), request.expected_labels_.end(), this->allocator_); - reasoning_ctx->InitializeExpectedTargets(expected_labels_vec, label_to_inner_id); - - FlattenInterfacePtr precise_flatten = nullptr; - ComputerInterfacePtr computer = nullptr; - if (not use_custom_distance) { - precise_flatten = this->basic_flatten_codes_; - if (use_reorder_) { - precise_flatten = this->high_precise_codes_; - } - if (create_new_raw_vector_) { - precise_flatten = this->raw_vector_; - } - computer = precise_flatten->FactoryComputer(get_data(query)); - } - for (const auto& pair : label_to_inner_id) { - float dist = 0.0F; - const auto inner_id = pair.second; - if (use_custom_distance) { - const auto label = this->label_table_->GetLabelById(inner_id); - request.distance_batch_func_(&label, 1, &dist); - CHECK_ARGUMENT(std::isfinite(dist), "distance callback must return finite scores"); - stats.AddDistance(SearchStatistics::DistancePhase::APPROXIMATE, - DistanceEvaluationBackend::UNKNOWN); - } else { - precise_flatten->Query(&dist, computer, &inner_id, 1, &ctx); - } - reasoning_ctx->SetTrueDistance(inner_id, dist); - } + reasoning_ctx = this->initialize_reasoning_context(request, k, use_custom_distance, &ctx); + if (reasoning_ctx != nullptr) { ctx.reasoning_ctx = reasoning_ctx.get(); } + return ctx; +} - InnerSearchParam search_param; - search_param.ep = this->entry_point_id_; +void +HGraph::search_route_graphs(const SearchRequest& request, + const HGraphSearchParameters& params, + InnerIdType entry_point, + bool use_custom_distance, + const void* query, + const VisitedListPtr& visited_list, + QueryContext* ctx, + InnerSearchParam& search_param) const { + search_param.ep = entry_point; search_param.topk = 1; search_param.ef = 1; search_param.is_inner_id_allowed = nullptr; @@ -534,82 +547,60 @@ HGraph::SearchWithRequest(const SearchRequest& request) const { search_param.distance_batch_size = request.distance_batch_size_; if (search_param.ep == INVALID_ENTRY_POINT) { - return make_empty_dataset_with_stats(); - } - - struct visited_list_guard { - std::shared_ptr pool; - VisitedListPtr visited_list; - - void - Release() { - if (visited_list != nullptr) { - pool->ReturnOne(visited_list); - visited_list.reset(); - } - } - - ~visited_list_guard() { - Release(); - } - }; - visited_list_guard vt_guard{this->pool_, this->pool_->TakeOne()}; - auto& vt = vt_guard.visited_list; - - const auto* raw_query = use_custom_distance ? nullptr : get_data(query); - ctx.distance_phase = DistanceEvaluationPhase::ROUTING; - for (auto i = static_cast(this->route_graphs_.size() - 1); i >= 0; --i) { - auto result = this->search_one_graph( - raw_query, this->route_graphs_[i], this->basic_flatten_codes_, search_param, vt, &ctx); + return; + } + ctx->distance_phase = DistanceEvaluationPhase::ROUTING; + for (auto i = static_cast(this->route_graphs_.size()) - 1; i >= 0; --i) { + auto result = this->search_one_graph(query, + this->route_graphs_[i], + this->basic_flatten_codes_, + search_param, + visited_list, + ctx); // An unrankable route seed can still bridge to finite bottom-layer results. if (not result->Empty()) { search_param.ep = result->Top().second; } } - ctx.distance_phase = DistanceEvaluationPhase::APPROXIMATE; - - FilterPtr ft = this->create_search_filter(request.filter_, params.use_extra_info_filter); + ctx->distance_phase = DistanceEvaluationPhase::APPROXIMATE; +} - if (request.enable_attribute_filter_ and this->attr_filter_index_ != nullptr) { - auto& schema = this->attr_filter_index_->field_type_map_; - auto expr = AstParse(request.attribute_filter_str_, &schema); - auto executor = Executor::MakeInstance(this->allocator_, expr, this->attr_filter_index_); - executor->Init(); - search_param.executors.emplace_back(executor); - } +void +HGraph::configure_bottom_graph_search(const SearchRequest& request, + const HGraphSearchParameters& params, + bool is_range, + int64_t k, + bool use_custom_distance, + const FilterPtr& filter, + const std::optional& threshold, + QueryContext* ctx, + InnerSearchParam& search_param) { + search_param.is_inner_id_allowed = filter; + search_param.enable_reorder = use_custom_distance ? false : params.enable_reorder; + search_param.consider_duplicate = true; + search_param.enable_rabitq_one_bit_search = + use_custom_distance ? false : params.rabitq_one_bit_search; + search_param.parallel_search_thread_count = params.parallel_search_thread_count; if (is_range) { search_param.ef = std::max(params.ef_search, request.limited_size_); - search_param.is_inner_id_allowed = ft; search_param.radius = request.radius_; search_param.search_mode = RANGE_SEARCH; - search_param.consider_duplicate = true; search_param.range_search_limit_size = static_cast(request.limited_size_); - search_param.parallel_search_thread_count = params.parallel_search_thread_count; - search_param.enable_reorder = use_custom_distance ? false : params.enable_reorder; - search_param.enable_rabitq_one_bit_search = - use_custom_distance ? false : params.rabitq_one_bit_search; } else { search_param.ef = std::max(params.ef_search, k); - search_param.is_inner_id_allowed = ft; - search_param.distance_threshold = request.threshold_; + search_param.distance_threshold = threshold; search_param.topk = static_cast(search_param.ef); if (params.topk_factor > 1.0F) { search_param.topk = std::min(search_param.topk, static_cast(static_cast(k) * params.topk_factor)); } - search_param.enable_reorder = use_custom_distance ? false : params.enable_reorder; - search_param.consider_duplicate = true; - search_param.enable_rabitq_one_bit_search = - use_custom_distance ? false : params.rabitq_one_bit_search; if (params.enable_time_record) { search_param.time_cost = std::make_shared(); search_param.time_cost->SetThreshold(params.timeout_ms); - stats.is_timeout.store(false, std::memory_order_relaxed); + ctx->stats->is_timeout.store(false, std::memory_order_relaxed); } - search_param.parallel_search_thread_count = params.parallel_search_thread_count; - if (static_cast(params.hops_limit) <= static_cast(params.ef_search)) { search_param.hops_limit = std::numeric_limits::max(); if (params.hops_limit != std::numeric_limits::max()) { @@ -625,83 +616,17 @@ HGraph::SearchWithRequest(const SearchRequest& request) const { search_param.skip_ratio = params.skip_ratio; search_param.skip_strategy_type = params.skip_strategy_type; +} - DistanceRecordVector rabitq_lower_bound_candidates(ctx.alloc); - auto* rabitq_lower_bound_candidates_ptr = - search_param.enable_rabitq_one_bit_search and use_reorder_ and - search_param.enable_reorder and reorder_by_base_ - ? &rabitq_lower_bound_candidates - : nullptr; - - DistHeapPtr search_result; - bool brute_force_used = false; - MCIHybridSearchResult mci_result(params, ft); - if (not use_custom_distance) { - if (params.brute_force_threshold > 0.0F and - mci_result.valid_ratio <= params.brute_force_threshold) { - if (is_range) { - search_result = this->brute_force_search( - raw_query, ft, request.limited_size_, request.radius_, &ctx); - } else { - search_result = this->brute_force_search( - raw_query, ft, k, 0.0F, &ctx, request.threshold_); - } - brute_force_used = true; - mci_result.route = "brute_force"; - } else { - mci_result = this->try_mci_search(request, params, ft, raw_query, search_param, &ctx); - if (mci_result.route == "mci") { - search_result = std::move(mci_result.result); - } else { - search_result = this->search_one_graph(raw_query, - this->bottom_graph_, - this->basic_flatten_codes_, - search_param, - vt, - &ctx, - rabitq_lower_bound_candidates_ptr); - } - } - } else { - search_result = this->search_one_graph(raw_query, - this->bottom_graph_, - this->basic_flatten_codes_, - search_param, - vt, - &ctx, - rabitq_lower_bound_candidates_ptr); - } - vt_guard.Release(); - - // Reorder - if (mci_result.route != "mci" and not brute_force_used and use_reorder_ and - search_param.enable_reorder) { - auto limit = is_range ? request.limited_size_ : k; - auto reorder_threshold = is_range ? std::nullopt : request.threshold_; - this->reorder(raw_query, - this->get_reorder_codes(), - search_result, - limit, - nullptr, - ctx, - rabitq_lower_bound_candidates_ptr, - reorder_threshold); - } else if (mci_result.route != "mci" and not brute_force_used and - search_param.enable_reorder and params.rabitq_one_bit_search) { - auto limit = is_range ? request.limited_size_ : k; - auto reorder_threshold = is_range ? std::nullopt : request.threshold_; - this->reorder(raw_query, - this->basic_flatten_codes_, - search_result, - limit, - nullptr, - ctx, - nullptr, - reorder_threshold); - } - - // Trim and pack results - if (is_range) { +DatasetPtr +HGraph::pack_search_result(const SearchRequest& request, + int64_t k, + DistHeapPtr search_result, + const QueryContext& ctx, + const MCIHybridSearchResult& mci_result, + const SearchStatistics& stats, + const std::shared_ptr& reasoning_ctx) const { + if (request.mode_ == SearchMode::RANGE_SEARCH) { while (not search_result->Empty() and search_result->Top().first > request.radius_ + THRESHOLD_ERROR) { search_result->Pop(); @@ -732,11 +657,11 @@ HGraph::SearchWithRequest(const SearchRequest& request) const { search_result->Push(record); } filter_search_result_by_threshold(search_result, request.threshold_, ctx.alloc); + while (search_result->Size() > static_cast(k)) { search_result->Pop(); } - // return an empty dataset directly if searcher returns nothing if (search_result->Empty()) { auto dataset_result = DatasetImpl::MakeEmptyDataset(); dataset_result->Statistics(mci_result.MakeStatistics(stats).Dump()); @@ -749,10 +674,9 @@ HGraph::SearchWithRequest(const SearchRequest& request) const { auto count = static_cast(search_result->Size()); Vector result_inner_ids(static_cast(count), this->allocator_); - auto [dataset_results, dists, ids] = create_fast_dataset(count, ctx.alloc); char* extra_infos = nullptr; - if (extra_info_size_ > 0 && this->extra_infos_ != nullptr) { + if (extra_info_size_ > 0 and this->extra_infos_ != nullptr) { extra_infos = static_cast(ctx.alloc->Allocate(extra_info_size_ * search_result->Size())); dataset_results->ExtraInfos(extra_infos) @@ -770,14 +694,282 @@ HGraph::SearchWithRequest(const SearchRequest& request) const { } dataset_results->Statistics(mci_result.MakeStatistics(stats).Dump()); - // Generate reasoning report if reasoning context was created if (reasoning_ctx) { reasoning_ctx->MarkResult(result_inner_ids); reasoning_ctx->DiagnoseExpectedTargets(); dataset_results->Reasoning(reasoning_ctx->GenerateReport()); } - return std::move(dataset_results); } +HGraphSearchParameters +HGraph::parse_and_validate_search_params(const SearchRequest& request, + bool is_range, + int64_t k, + bool use_custom_distance) const { + if (use_custom_distance) { + CHECK_ARGUMENT(request.distance_batch_size_ > 0, + "distance_batch_size must be greater than 0"); + CHECK_ARGUMENT(not is_range, "HGraph custom distance only supports KNN search"); + } + + if (is_range) { + if (not use_custom_distance) { + this->validate_range_args(request.query_, request.radius_, request.limited_size_); + } + } else if (not use_custom_distance) { + this->validate_knn_args(request.query_, k); + } else { + CHECK_ARGUMENT(k > 0, "topk must be greater than 0"); + } + + auto params = HGraphSearchParameters::FromJson(request.params_str_); + if (use_custom_distance) { + CHECK_ARGUMENT(params.parallel_search_thread_count == 1, + "HGraph custom query distance does not support parallel search"); + CHECK_ARGUMENT(params.brute_force_threshold <= 0.0F, + "HGraph custom query distance does not support brute_force_threshold"); + } + CHECK_ARGUMENT( // NOLINT + params.ef_search >= 1, + fmt::format("ef_search({}) must be at least 1", params.ef_search)); + return params; +} + +std::shared_ptr +HGraph::initialize_reasoning_context(const SearchRequest& request, + int64_t k, + bool use_custom_distance, + QueryContext* ctx) const { + if (request.mode_ == SearchMode::RANGE_SEARCH or request.expected_labels_.empty()) { + return nullptr; + } + + auto reasoning_ctx = std::make_shared(this->allocator_); + reasoning_ctx->SetSearchParams( + k, "HGraph", use_custom_distance ? false : use_reorder_, request.filter_ != nullptr); + + UnorderedMap label_to_inner_id(this->allocator_); + for (const auto& label : request.expected_labels_) { + auto [success, inner_id] = label_table_->TryGetIdByLabel(label, true); + if (success) { + label_to_inner_id[label] = inner_id; + } + } + + Vector expected_labels_vec( + request.expected_labels_.begin(), request.expected_labels_.end(), this->allocator_); + reasoning_ctx->InitializeExpectedTargets(expected_labels_vec, label_to_inner_id); + + FlattenInterfacePtr precise_flatten = nullptr; + ComputerInterfacePtr computer = nullptr; + if (not use_custom_distance) { + precise_flatten = this->basic_flatten_codes_; + if (use_reorder_) { + precise_flatten = this->high_precise_codes_; + } + if (create_new_raw_vector_) { + precise_flatten = this->raw_vector_; + } + computer = precise_flatten->FactoryComputer(get_data(request.query_)); + } + for (const auto& pair : label_to_inner_id) { + float dist = 0.0F; + const auto inner_id = pair.second; + if (use_custom_distance) { + const auto label = this->label_table_->GetLabelById(inner_id); + request.distance_batch_func_(&label, 1, &dist); + CHECK_ARGUMENT(std::isfinite(dist), "distance callback must return finite scores"); + if (ctx != nullptr and ctx->stats != nullptr) { + ctx->stats->AddDistance(SearchStatistics::DistancePhase::APPROXIMATE, + DistanceEvaluationBackend::UNKNOWN); + } + } else { + precise_flatten->Query(&dist, computer, &inner_id, 1, ctx); + } + reasoning_ctx->SetTrueDistance(inner_id, dist); + } + return reasoning_ctx; +} + +[[nodiscard]] DatasetPtr +HGraph::SearchWithRequest(const SearchRequest& request) const { + ValidateSearchThreshold(request.threshold_); + SearchStatistics stats; + + const auto& query = request.query_; + bool is_range = (request.mode_ == SearchMode::RANGE_SEARCH); + auto k = request.topk_; + const bool use_custom_distance = request.distance_batch_func_ != nullptr; + + /***** Step 1: Parse and validate request-specific search parameters. *****/ + auto params = this->parse_and_validate_search_params(request, is_range, k, use_custom_distance); + + /***** Step 2: Keep the index state stable while searching mutable indexes. *****/ + std::shared_lock force_remove_rlock; + std::shared_lock shared_lock; + if (!this->immutable_.load(std::memory_order_acquire)) { + if (this->support_force_remove()) { + force_remove_rlock = std::shared_lock(this->force_remove_mutex_); + } + shared_lock = this->acquire_global_read_lock(); + } + const auto element_count = GetNumElements(); + const auto entry_point = this->entry_point_id_; + if (element_count == 0 or entry_point == INVALID_ENTRY_POINT) { + return make_empty_dataset_with_stats(); + } + k = std::min(k, element_count); + + /***** Step 3: Set up query-local allocation, statistics, and optional reasoning. *****/ + std::shared_ptr reasoning_ctx; + auto ctx = + this->create_query_context(request, params, k, use_custom_distance, &stats, reasoning_ctx); + + /***** Step 4: Navigate upper route graphs to obtain the bottom-graph entry point. *****/ + VisitedListGuard vt_guard{this->pool_, this->pool_->TakeOne()}; + auto& vt = vt_guard.Get(); + const auto* raw_query = use_custom_distance ? nullptr : get_data(query); + InnerSearchParam search_param; + this->search_route_graphs( + request, params, entry_point, use_custom_distance, raw_query, vt, &ctx, search_param); + + /***** Step 5: Build filters and configure the bottom-graph search parameters. *****/ + FilterPtr ft = this->create_search_filter(request.filter_, params.use_extra_info_filter); + if (request.enable_attribute_filter_ and this->attr_filter_index_ != nullptr) { + auto& schema = this->attr_filter_index_->field_type_map_; + auto expr = AstParse(request.attribute_filter_str_, &schema); + auto executor = Executor::MakeInstance(this->allocator_, expr, this->attr_filter_index_); + executor->Init(); + search_param.executors.emplace_back(executor); + } + + this->configure_bottom_graph_search(request, + params, + is_range, + k, + use_custom_distance, + ft, + request.threshold_, + &ctx, + search_param); + + /***** Step 6: Select and execute brute-force, MCI, or bottom-graph search. *****/ + DistanceRecordVector rabitq_lower_bound_candidates(ctx.alloc); + auto* rabitq_lower_bound_candidates_ptr = + search_param.enable_rabitq_one_bit_search and use_reorder_ and + search_param.enable_reorder and reorder_by_base_ + ? &rabitq_lower_bound_candidates + : nullptr; + + bool brute_force_used = false; + MCIHybridSearchResult mci_result(params, ft); + DistHeapPtr search_result; + if (not use_custom_distance) { + search_plan_input plan_input; + plan_input.is_range_search = is_range; + plan_input.valid_ratio = mci_result.valid_ratio; + plan_input.brute_force_threshold = params.brute_force_threshold; + plan_input.use_mci = params.use_mci; + plan_input.mci_enabled = this->mci_parameters_.enabled; + plan_input.has_attribute_executor = not search_param.executors.empty(); + plan_input.mci_valid_ratio_threshold = params.mci_hgraph_valid_ratio_threshold; + const bool brute_force_route = + select_search_plan(plan_input, false) == search_plan::K_BRUTE_FORCE; + bool mci_available = false; + if (not brute_force_route) { + // MCI seeds use the original external-label filter; ft wraps it for inner-ID search. + const auto bitset_seed_source = has_bitset_source(request.filter_); + bool valid_id_seed_source = false; + if (params.use_mci and this->mci_parameters_.enabled and + search_param.executors.empty()) { + valid_id_seed_source = has_valid_id_source(request.filter_); + if (valid_id_seed_source or bitset_seed_source) { + plan_input.mci_has_clique_index = + this->mci_cliques_ != nullptr and + this->mci_cliques_->HasCliqueIndex(this->total_count_.load()); + } + } + plan_input.has_valid_id_source = valid_id_seed_source; + plan_input.has_bitset_source = bitset_seed_source; + mci_available = is_mci_available(plan_input); + } + + switch (select_search_plan(plan_input, mci_available)) { + case search_plan::K_BRUTE_FORCE: + if (is_range) { + search_result = this->brute_force_search( + raw_query, ft, request.limited_size_, request.radius_, &ctx); + } else { + search_result = this->brute_force_search( + raw_query, ft, k, 0.0F, &ctx, request.threshold_); + } + brute_force_used = true; + mci_result.route = "brute_force"; + break; + case search_plan::K_MCI: + mci_result = + this->try_mci_search(request, params, ft, raw_query, search_param, &ctx); + if (mci_result.route == "mci") { + search_result = std::move(mci_result.result); + break; + } + [[fallthrough]]; + case search_plan::K_BOTTOM_GRAPH: + if (mci_available) { + mci_result.route = "hgraph"; + } + search_result = this->search_one_graph(raw_query, + this->bottom_graph_, + this->basic_flatten_codes_, + search_param, + vt, + &ctx, + rabitq_lower_bound_candidates_ptr); + break; + } + } else { + search_result = this->search_one_graph(raw_query, + this->bottom_graph_, + this->basic_flatten_codes_, + search_param, + vt, + &ctx, + rabitq_lower_bound_candidates_ptr); + } + /***** Step 7: Return the pooled visited list before post-processing results. *****/ + vt_guard.Release(); + + /***** Step 8: Reorder candidates when the selected route permits it. *****/ + if (mci_result.route != "mci" and not brute_force_used and use_reorder_ and + search_param.enable_reorder) { + auto limit = is_range ? request.limited_size_ : k; + auto reorder_threshold = is_range ? std::nullopt : request.threshold_; + this->reorder(raw_query, + this->get_reorder_codes(), + search_result, + limit, + nullptr, + ctx, + rabitq_lower_bound_candidates_ptr, + reorder_threshold); + } else if (mci_result.route != "mci" and not brute_force_used and + search_param.enable_reorder and params.rabitq_one_bit_search) { + auto limit = is_range ? request.limited_size_ : k; + auto reorder_threshold = is_range ? std::nullopt : request.threshold_; + this->reorder(raw_query, + this->basic_flatten_codes_, + search_result, + limit, + nullptr, + ctx, + nullptr, + reorder_threshold); + } + + /***** Step 9: Trim, pack, and annotate the final dataset. *****/ + return this->pack_search_result( + request, k, search_result, ctx, mci_result, stats, reasoning_ctx); +} + } // namespace vsag