diff --git a/src/binding/c/c_api.cc b/src/binding/c/c_api.cc index 6476752..714b711 100644 --- a/src/binding/c/c_api.cc +++ b/src/binding/c/c_api.cc @@ -5908,6 +5908,29 @@ zvec_error_code_t zvec_sub_query_set_query_vector( return ZVEC_OK; } +zvec_error_code_t zvec_sub_query_set_sparse_vector( + zvec_sub_query_t *query, const uint32_t *indices, const float *values, + size_t count) { + if (!query || (!indices && count > 0) || (!values && count > 0)) { + SET_LAST_ERROR(ZVEC_ERROR_INVALID_ARGUMENT, + "Sub-vector query, indices or values pointer is null"); + return ZVEC_ERROR_INVALID_ARGUMENT; + } + + auto *ptr = reinterpret_cast(query); + auto &payload = std::get(ptr->target_.clause_); + if (count == 0) { + payload.sparse_indices_.clear(); + payload.sparse_values_.clear(); + return ZVEC_OK; + } + payload.sparse_indices_.assign( + reinterpret_cast(indices), count * sizeof(uint32_t)); + payload.sparse_values_.assign( + reinterpret_cast(values), count * sizeof(float)); + return ZVEC_OK; +} + zvec_error_code_t zvec_sub_query_set_sparse_indices( zvec_sub_query_t *query, const uint32_t *indices, size_t count) { if (!query || (!indices && count > 0)) { diff --git a/src/db/collection.cc b/src/db/collection.cc index 0e799fb..d88efb8 100644 --- a/src/db/collection.cc +++ b/src/db/collection.cc @@ -14,8 +14,6 @@ #include #include -#include -#include #include #include #include @@ -36,6 +34,7 @@ #include #include "db/common/constants.h" #include "db/common/file_helper.h" +#include "db/common/global_resource.h" #include "db/common/profiler.h" #include "db/common/typedef.h" #include "db/index/common/delete_store.h" @@ -1581,7 +1580,8 @@ Status CollectionImpl::DeleteByFilter(const std::string &filter) { query.output_fields_ = std::vector{}; query.include_doc_id_ = true; - auto ret = sql_engine_->execute(schema_, query, get_all_segments()); + auto ret = + sql_engine_->execute(schema_, std::move(query), get_all_segments()); if (!ret.has_value()) { return ret.error(); } @@ -1619,7 +1619,7 @@ Result CollectionImpl::Query(const SearchQuery &query) const { return DocPtrList(); } - return sql_engine_->execute(schema_, sanitized, segments); + return sql_engine_->execute(schema_, std::move(sanitized), segments); } Result CollectionImpl::Query(const MultiQuery &query) const { @@ -1628,13 +1628,14 @@ Result CollectionImpl::Query(const MultiQuery &query) const { CHECK_DESTROY_RETURN_STATUS_EXPECTED(destroyed_, false); if (query.queries.size() < 2) { - return tl::make_unexpected( - Status::InvalidArgument("Query requires at least 2 sub-queries")); + return tl::make_unexpected(Status::InvalidArgument( + "Invalid query: MultiQuery requires at least 2 sub-queries, got ", + query.queries.size())); } if (!query.reranker) { - return tl::make_unexpected( - Status::InvalidArgument("Reranker is required for multi-vector query")); + return tl::make_unexpected(Status::InvalidArgument( + "Invalid query: MultiQuery requires a reranker")); } auto segments = get_all_segments(); @@ -1642,10 +1643,15 @@ Result CollectionImpl::Query(const MultiQuery &query) const { return DocPtrList(); } + struct PendingQuery { + std::string field_name; + SearchQuery query; + }; + // Convert each SubQuery to a SearchQuery and validate. std::set seen_fields; - std::vector converted_queries; - converted_queries.reserve(query.queries.size()); + std::vector pending_queries; + pending_queries.reserve(query.queries.size()); for (const auto &sub : query.queries) { const auto &target = sub.target_; @@ -1670,27 +1676,39 @@ Result CollectionImpl::Query(const MultiQuery &query) const { auto s = sq.validate_and_sanitize(field_schema); CHECK_RETURN_STATUS_EXPECTED(s); - converted_queries.push_back(std::move(sq)); - } - - // Execute each sub-query concurrently and collect results per field. - std::vector>> futures; - futures.reserve(converted_queries.size()); - for (const auto &sq : converted_queries) { - futures.push_back(std::async(std::launch::async, [&]() { - auto engine = sqlengine::SQLEngine::create(std::make_shared()); - return engine->execute(schema_, sq, segments); - })); + pending_queries.push_back({target.field_name_, std::move(sq)}); } std::map query_results; - for (size_t i = 0; i < converted_queries.size(); ++i) { - auto result = futures[i].get(); - if (!result.has_value()) { - return tl::make_unexpected(result.error()); + + auto execute_query = [&](PendingQuery &pending) -> Result { + auto engine = sqlengine::SQLEngine::create(std::make_shared()); + return engine->execute(schema_, std::move(pending.query), segments); + }; + + std::vector> results(pending_queries.size()); + + // Single-segment queries have no segment-level fanout; multi-segment queries + // already use the query pool per sub-query. + if (segments.size() == 1) { + auto group = GlobalResource::Instance().query_thread_pool()->make_group(); + for (size_t i = 0; i < pending_queries.size(); ++i) { + group->execute( + [&, i]() { results[i] = execute_query(pending_queries[i]); }); } - query_results[converted_queries[i].target_.field_name_] = - std::move(result.value()); + group->wait_finish(); + } else { + for (size_t i = 0; i < pending_queries.size(); ++i) { + results[i] = execute_query(pending_queries[i]); + } + } + + for (size_t i = 0; i < pending_queries.size(); ++i) { + if (!results[i]) { + return tl::make_unexpected(results[i].error()); + } + query_results[pending_queries[i].field_name] = + std::move(results[i].value()); } // Merge and rerank results diff --git a/src/db/common/profiler.h b/src/db/common/profiler.h index 52b8838..57a7ef8 100644 --- a/src/db/common/profiler.h +++ b/src/db/common/profiler.h @@ -14,6 +14,8 @@ #pragma once #include +#include +#include #include #include #include @@ -198,4 +200,31 @@ class ScopedLatency { Profiler::Ptr profiler_; }; +//! RAII helper that closes a profiler stage on every exit path. +class ScopedProfilerStage { + public: + ScopedProfilerStage(Profiler::Ptr profiler, const std::string &name) + : profiler_(std::move(profiler)) { + active_ = profiler_ && profiler_->open_stage(name) == 0; + } + + ~ScopedProfilerStage() { + close(); + } + + void close() { + if (active_) { + profiler_->close_stage(); + active_ = false; + } + } + + ScopedProfilerStage(const ScopedProfilerStage &) = delete; + ScopedProfilerStage &operator=(const ScopedProfilerStage &) = delete; + + private: + Profiler::Ptr profiler_; + bool active_{false}; +}; + } // namespace zvec diff --git a/src/db/sqlengine/analyzer/query_analyzer.cc b/src/db/sqlengine/analyzer/query_analyzer.cc index c4af8f3..2e144dd 100644 --- a/src/db/sqlengine/analyzer/query_analyzer.cc +++ b/src/db/sqlengine/analyzer/query_analyzer.cc @@ -21,8 +21,6 @@ #include #include #include -#include "db/common/constants.h" -#include "db/common/error_code.h" #include "db/index/common/type_helper.h" #include "db/sqlengine/analyzer/query_node.h" #include "db/sqlengine/common/util.h" @@ -515,42 +513,35 @@ Status QueryAnalyzer::check_and_convert_vector( vector_field_name); } - std::string vector_term; uint32_t dimension = vector_meta->dimension(); - std::string vector_sparse_indices; - std::string vector_sparse_values; - QueryParams::Ptr query_params; const QueryNode::Ptr &vector_value_node = query_rel_node->right(); // for pb request if (vector_value_node->op() == QueryNodeOp::Q_VECTOR_MATRIX_VALUE) { // for format vector = [,,,] - const QueryVectorMatrixNode::Ptr &vector_node = + QueryVectorMatrixNode::Ptr vector_node = std::dynamic_pointer_cast(vector_value_node); - // we only have vector matrix, other info is not available - vector_term = vector_node->matrix(); - vector_sparse_indices = vector_node->sparse_indices(); - vector_sparse_values = vector_node->sparse_values(); - query_params = vector_node->query_params(); + // Consume the vector payload; this node is detached from the search + // condition after conversion. + auto vector_data = vector_node->take_node(); + auto core_data_type = + DataTypeCodeBook::to_data_type(vector_meta->data_type()); + if (core_data_type == core::IndexMeta::DataType::DT_UNDEFINED) { + return Status::InvalidArgument("invalid data type:", + (int)vector_meta->data_type()); + } + + *vector_cond = std::make_shared( + vector_meta, vector_data->take_matrix(), core_data_type, dimension, + vector_data->take_sparse_indices(), vector_data->take_sparse_values(), + vector_data->take_query_params()); + return Status::OK(); } else { return Status::InvalidArgument("invalid vector value node. op[", vector_value_node->op_name(), "], text[", vector_value_node->text(), "]"); } - - auto core_data_type = - DataTypeCodeBook::to_data_type(vector_meta->data_type()); - if (core_data_type == core::IndexMeta::DataType::DT_UNDEFINED) { - return Status::InvalidArgument("invalid data type:", - (int)vector_meta->data_type()); - } - - *vector_cond = std::make_shared( - vector_meta, vector_term, core_data_type, dimension, - std::move(vector_sparse_indices), std::move(vector_sparse_values), - std::move(query_params)); - return Status::OK(); } diff --git a/src/db/sqlengine/analyzer/query_info.h b/src/db/sqlengine/analyzer/query_info.h index ad9b381..c21ea48 100644 --- a/src/db/sqlengine/analyzer/query_info.h +++ b/src/db/sqlengine/analyzer/query_info.h @@ -47,13 +47,13 @@ class QueryInfo { using Ptr = std::shared_ptr; QueryVectorCondInfo(const FieldSchema *vector_schema, - const std::string &vector_term, + std::string vector_term, core::IndexMeta::DataType core_data_type, int dimension, std::string vector_sparse_indices, std::string vector_sparse_values, QueryParams::Ptr query_params) : vector_schema_(vector_schema), - vector_term_(vector_term), + vector_term_(std::move(vector_term)), data_type_(core_data_type), dimension_(dimension), vector_sparse_indices_(std::move(vector_sparse_indices)), diff --git a/src/db/sqlengine/analyzer/query_node.h b/src/db/sqlengine/analyzer/query_node.h index f54d1dc..6d2e352 100644 --- a/src/db/sqlengine/analyzer/query_node.h +++ b/src/db/sqlengine/analyzer/query_node.h @@ -166,7 +166,7 @@ class QueryNode : public Generic_Node { QueryNode::Ptr detach_from_invert_cond(QueryInfo *query_info_ptr); - virtual std::string text() const override; + std::string text() const override; virtual void set_text(std::string /*new_val*/) { /* for QueryConstantNode only */ @@ -218,8 +218,12 @@ class QueryVectorMatrixNode : public QueryNode { return node_->query_params(); } + std::shared_ptr take_node() { + return std::move(node_); + } + private: - std::shared_ptr node_{nullptr}; + std::shared_ptr node_{nullptr}; }; class QueryConstantNode : public QueryNode { @@ -261,7 +265,7 @@ class QueryFuncNode : public QueryNode { using Ptr = std::shared_ptr; QueryFuncNode(); - virtual ~QueryFuncNode() = default; + ~QueryFuncNode() override = default; void set_func_name_node(QueryNode::Ptr func_name_node); const QueryNode::Ptr &get_func_name_node() const; diff --git a/src/db/sqlengine/parser/node.h b/src/db/sqlengine/parser/node.h index a08dd4f..344b1ea 100644 --- a/src/db/sqlengine/parser/node.h +++ b/src/db/sqlengine/parser/node.h @@ -118,7 +118,7 @@ class Node : public Generic_Node { NodeType type(); - virtual std::string text() const override; + std::string text() const override; std::string to_string(); private: @@ -137,7 +137,7 @@ class RangeNode : public Node { RangeNode(); RangeNode(bool m_min_equal, bool m_max_equal); - virtual ~RangeNode() = default; + ~RangeNode() override = default; void set_min_equal(bool value); void set_max_equal(bool value); @@ -171,18 +171,34 @@ class VectorMatrixNode : public Node { return matrix_; } + std::string take_matrix() { + return std::move(matrix_); + } + const std::string &sparse_indices() const { return sparse_indices_; } + std::string take_sparse_indices() { + return std::move(sparse_indices_); + } + const std::string &sparse_values() const { return sparse_values_; } + std::string take_sparse_values() { + return std::move(sparse_values_); + } + const QueryParams::Ptr &query_params() const { return query_params_; } + QueryParams::Ptr take_query_params() { + return std::move(query_params_); + } + std::string text() const override { // do not distinguish between matrix and vector static std::string txt = "[...]"; @@ -231,7 +247,7 @@ class FuncNode : public Node { using Ptr = std::shared_ptr; FuncNode(); - virtual ~FuncNode() = default; + ~FuncNode() override = default; void set_func_name_node(Node::Ptr func_name_node); const Node::Ptr &get_func_name_node(); diff --git a/src/db/sqlengine/parser/sql_info_helper.cc b/src/db/sqlengine/parser/sql_info_helper.cc index 6ddd3a5..8b2c937 100644 --- a/src/db/sqlengine/parser/sql_info_helper.cc +++ b/src/db/sqlengine/parser/sql_info_helper.cc @@ -26,16 +26,17 @@ namespace zvec::sqlengine { using namespace zvec; -Node::Ptr handle_vector(const SearchQuery &request, std::string * /*err_msg*/) { - const auto *vc = request.target_.get_vector_clause(); +Node::Ptr handle_vector(SearchQuery *request) { + auto *vc = request->target_.get_vector_clause(); if (vc == nullptr) { return nullptr; } Node::Ptr rel_exp = std::make_shared(NodeOp::T_EQ); - rel_exp->set_left(std::make_shared(request.target_.field_name_)); + rel_exp->set_left(std::make_shared(request->target_.field_name_)); rel_exp->set_right(std::make_shared( - vc->query_vector_, vc->sparse_indices_, vc->sparse_values_, - request.target_.query_params_)); + std::move(vc->query_vector_), std::move(vc->sparse_indices_), + std::move(vc->sparse_values_), + std::move(request->target_.query_params_))); return rel_exp; } @@ -65,18 +66,18 @@ void handle_query_field(const SearchQuery *query, SelectInfo *selected_info) { } } -bool SQLInfoHelper::MessageToSQLInfo(const SearchQuery *query, - Node::Ptr filter_node, - std::shared_ptr group_by, - sqlengine::SQLInfo::Ptr *sql_info, - std::string *err_msg) { +Result SQLInfoHelper::BuildSQLInfoFromSearchQuery( + SearchQuery query, Node::Ptr filter_node, + std::shared_ptr group_by) { Node::Ptr index_params_node_ptr = nullptr; - if (const auto *vc = query->target_.get_vector_clause(); + if (const auto *vc = query.target_.get_vector_clause(); vc != nullptr && (!vc->query_vector_.empty() || !vc->sparse_indices_.empty())) { - index_params_node_ptr = handle_vector(*query, err_msg); + index_params_node_ptr = handle_vector(&query); if (index_params_node_ptr == nullptr) { - return false; + return tl::make_unexpected(Status::InvalidArgument( + "Failed to build vector condition for field: ", + query.target_.field_name_)); } } @@ -92,13 +93,13 @@ bool SQLInfoHelper::MessageToSQLInfo(const SearchQuery *query, } SelectInfo::Ptr select_info = std::make_shared(""); - handle_query_field(query, select_info.get()); + handle_query_field(&query, select_info.get()); select_info->set_search_cond(cond_expr); - uint32_t topk = query->topk_; + uint32_t topk = query.topk_; select_info->set_limit(topk); - select_info->set_include_vector(query->include_vector_); - select_info->set_include_doc_id(query->include_doc_id_); + select_info->set_include_vector(query.include_vector_); + select_info->set_include_doc_id(query.include_doc_id_); select_info->set_group_by(std::move(group_by)); // @@ -111,8 +112,7 @@ bool SQLInfoHelper::MessageToSQLInfo(const SearchQuery *query, // select_info->add_order_by_elem(std::move(orderby_elem_info)); // } - *sql_info = std::make_shared(SQLInfo::SQLType::SELECT, select_info); - return true; + return std::make_shared(SQLInfo::SQLType::SELECT, select_info); } } // namespace zvec::sqlengine diff --git a/src/db/sqlengine/parser/sql_info_helper.h b/src/db/sqlengine/parser/sql_info_helper.h index 4c73c45..09d45fa 100644 --- a/src/db/sqlengine/parser/sql_info_helper.h +++ b/src/db/sqlengine/parser/sql_info_helper.h @@ -15,6 +15,7 @@ #pragma once #include +#include #include "db/sqlengine/common/group_by.h" #include "db/sqlengine/parser/node.h" #include "db/sqlengine/parser/sql_info.h" @@ -23,11 +24,11 @@ namespace zvec::sqlengine { class SQLInfoHelper { public: - //! Perform QueryRequest to sql info conversion: - static bool MessageToSQLInfo(const SearchQuery *query, Node::Ptr filter_node, - std::shared_ptr group_by, - sqlengine::SQLInfo::Ptr *sql_info, - std::string *err_msg); + //! Build SQLInfo from SearchQuery. Takes query by value so callers may copy + //! or move it; vector payloads can be moved while building SQLInfo. + static Result BuildSQLInfoFromSearchQuery( + SearchQuery query, Node::Ptr filter_node, + std::shared_ptr group_by); }; } // namespace zvec::sqlengine diff --git a/src/db/sqlengine/sqlengine.h b/src/db/sqlengine/sqlengine.h index 011cff3..518970f 100644 --- a/src/db/sqlengine/sqlengine.h +++ b/src/db/sqlengine/sqlengine.h @@ -27,7 +27,7 @@ class SQLEngine { virtual ~SQLEngine(); virtual Result execute( - CollectionSchema::Ptr collection, const SearchQuery &query, + CollectionSchema::Ptr collection, SearchQuery query, const std::vector &segments) = 0; virtual Result execute_group_by( diff --git a/src/db/sqlengine/sqlengine_impl.cc b/src/db/sqlengine/sqlengine_impl.cc index 5fbbbb0..51e8b1b 100644 --- a/src/db/sqlengine/sqlengine_impl.cc +++ b/src/db/sqlengine/sqlengine_impl.cc @@ -56,13 +56,13 @@ SQLEngineImpl::SQLEngineImpl(zvec::Profiler::Ptr profiler) : profiler_(std::move(profiler)) {} Result SQLEngineImpl::execute( - CollectionSchema::Ptr collection, const SearchQuery &query, + CollectionSchema::Ptr collection, SearchQuery query, const std::vector &segments) { if (segments.empty()) { return DocPtrList{}; } - auto query_info = parse_request(collection, query, nullptr); + auto query_info = build_query_info(collection, std::move(query), nullptr); if (!query_info) { return tl::make_unexpected(query_info.error()); } @@ -100,7 +100,7 @@ Result SQLEngineImpl::execute_group_by( } SearchQuery query = from_group_by(group_by_query); - auto query_info = parse_request( + auto query_info = build_query_info( collection, query, std::make_shared(group_by_query.group_by_field_name_, group_by_query.group_topk_, @@ -231,24 +231,21 @@ Result SQLEngineImpl::parse_fts_query( Result SQLEngineImpl::parse_sql_info( const CollectionSchema &schema, const SQLInfo::Ptr &sql_info) { - profiler_->open_stage("analyze stage"); + ScopedProfilerStage stage_guard(profiler_, "analyze_sql_info"); QueryAnalyzer analyzer; auto query_info = analyzer.analyze(schema, sql_info); if (!query_info) { return tl::make_unexpected(Status::InvalidArgument( "Analyze SQL info failed: ", query_info.error().c_str())); } - profiler_->close_stage(); LOG_DEBUG("query_info: [%s]", query_info.value()->to_string().c_str()); return query_info.value(); } -Result SQLEngineImpl::parse_request( - CollectionSchema::Ptr collection, const SearchQuery &request, +Result SQLEngineImpl::build_query_info( + CollectionSchema::Ptr collection, SearchQuery request, std::shared_ptr group_by) { - profiler_->open_stage("message_to_sqlinfo"); - sqlengine::SQLInfo::Ptr sql_info; - std::string err_msg; + ScopedProfilerStage stage_guard(profiler_, "build_sql_info"); Node::Ptr filter_node; if (!request.filter_.empty()) { ZVecParser::Ptr parser = ZVecParser::create(); @@ -271,19 +268,7 @@ Result SQLEngineImpl::parse_request( } } - sqlengine::SQLInfoHelper::MessageToSQLInfo(&request, std::move(filter_node), - std::move(group_by), &sql_info, - &err_msg); - profiler_->close_stage(); - if (!err_msg.empty()) { - LOG_ERROR("QueryAgent, message to sql info failed, err_msg: %s", - err_msg.c_str()); - return tl::make_unexpected(Status::InvalidArgument( - "Convert message to SQL info failed: ", err_msg)); - } - - // If the request carries an FTS query, parse it and attach to SelectInfo - // so that query_analyzer can propagate it to QueryInfo. + FtsCondInfo::Ptr fts_cond_info; if (const auto *fts_clause = request.target_.get_fts_clause()) { auto fts_result = parse_fts_query(collection, request.target_.field_name_, *fts_clause, @@ -291,13 +276,30 @@ Result SQLEngineImpl::parse_request( if (!fts_result) { return tl::make_unexpected(fts_result.error()); } - auto select_info = - std::dynamic_pointer_cast(sql_info->base_info()); - select_info->set_fts_cond_info(std::move(fts_result.value())); + fts_cond_info = std::move(fts_result.value()); } - LOG_DEBUG("Sql info is %s", sql_info->to_string().c_str()); - return parse_sql_info(*collection, std::move(sql_info)); + auto sql_info = sqlengine::SQLInfoHelper::BuildSQLInfoFromSearchQuery( + std::move(request), std::move(filter_node), std::move(group_by)); + if (!sql_info) { + return tl::make_unexpected(sql_info.error()); + } + + // Attach FTS info to SelectInfo so query_analyzer can propagate it to + // QueryInfo. + if (fts_cond_info) { + auto select_info = + std::dynamic_pointer_cast(sql_info.value()->base_info()); + if (!select_info) { + return tl::make_unexpected(Status::InternalError( + "BuildSQLInfoFromSearchQuery did not produce SelectInfo")); + } + select_info->set_fts_cond_info(std::move(fts_cond_info)); + } + + LOG_DEBUG("Sql info is %s", sql_info.value()->to_string().c_str()); + stage_guard.close(); + return parse_sql_info(*collection, std::move(sql_info.value())); } Result> @@ -306,7 +308,7 @@ SQLEngineImpl::search_by_query_info( std::vector *query_infos) { global_init(); - profiler_->open_stage("plan stage"); + ScopedProfilerStage stage_guard(profiler_, "build_query_plan"); QueryPlanner planner(collection.get()); auto plan_info = planner.make_plan(segments, profiler_->trace_id(), query_infos); @@ -314,7 +316,6 @@ SQLEngineImpl::search_by_query_info( LOG_ERROR("plan query_info failed: [%s]", plan_info.error().c_str()); return tl::make_unexpected(plan_info.error()); } - profiler_->close_stage(); // LOG_DEBUG("plan_info: [%s]", plan_info->to_string().c_str()); return plan_info.value()->execute_to_reader(); } diff --git a/src/db/sqlengine/sqlengine_impl.h b/src/db/sqlengine/sqlengine_impl.h index 465c13e..bcc616b 100644 --- a/src/db/sqlengine/sqlengine_impl.h +++ b/src/db/sqlengine/sqlengine_impl.h @@ -34,10 +34,10 @@ class SQLEngineImpl : public SQLEngine { public: SQLEngineImpl(zvec::Profiler::Ptr profiler); - //! Parse pb request - Result parse_request(CollectionSchema::Ptr collection, - const SearchQuery &request, - std::shared_ptr group_by); + //! Build analyzed query info from a structured search query. + Result build_query_info(CollectionSchema::Ptr collection, + SearchQuery request, + std::shared_ptr group_by); //! Perform search with given query_info, segments and index filter Result> search_by_query_info( @@ -46,7 +46,7 @@ class SQLEngineImpl : public SQLEngine { std::vector *query_infos); Result execute( - CollectionSchema::Ptr collection, const SearchQuery &query, + CollectionSchema::Ptr collection, SearchQuery query, const std::vector &segments) override; Result execute_group_by( @@ -79,4 +79,4 @@ class SQLEngineImpl : public SQLEngine { std::string execution_time_info_{}; }; -} // namespace zvec::sqlengine \ No newline at end of file +} // namespace zvec::sqlengine diff --git a/src/include/zvec/c_api.h b/src/include/zvec/c_api.h index 291209f..64cec24 100644 --- a/src/include/zvec/c_api.h +++ b/src/include/zvec/c_api.h @@ -2167,6 +2167,18 @@ zvec_sub_query_get_field_name(const zvec_sub_query_t *query); ZVEC_EXPORT zvec_error_code_t ZVEC_CALL zvec_sub_query_set_query_vector( zvec_sub_query_t *query, const void *data, size_t size); +/** + * @brief Set sparse vector indices and values + * @param query Sub-query pointer + * @param indices Array of uint32_t indices + * @param values Array of float values + * @param count Number of sparse vector entries + * @return zvec_error_code_t Error code + */ +ZVEC_EXPORT zvec_error_code_t ZVEC_CALL zvec_sub_query_set_sparse_vector( + zvec_sub_query_t *query, const uint32_t *indices, const float *values, + size_t count); + /** * @brief Set sparse vector indices * @param query Sub-query pointer diff --git a/tests/c/c_api_test.c b/tests/c/c_api_test.c index ebbf480..662405f 100644 --- a/tests/c/c_api_test.c +++ b/tests/c/c_api_test.c @@ -4562,6 +4562,19 @@ void test_multi_vector_query_with_rrf_reranker(void) { zvec_free((char *)got_fields); } + zvec_sub_query_t *sparse_query = zvec_sub_query_create(); + TEST_ASSERT(sparse_query != NULL); + uint32_t sparse_indices[] = {1, 3}; + float sparse_values[] = {0.25f, 0.75f}; + err = zvec_sub_query_set_sparse_vector(sparse_query, sparse_indices, + sparse_values, 2); + TEST_ASSERT(err == ZVEC_OK); + err = zvec_sub_query_set_sparse_vector(sparse_query, NULL, sparse_values, 2); + TEST_ASSERT(err == ZVEC_ERROR_INVALID_ARGUMENT); + err = zvec_sub_query_set_sparse_vector(sparse_query, sparse_indices, NULL, 2); + TEST_ASSERT(err == ZVEC_ERROR_INVALID_ARGUMENT); + zvec_sub_query_destroy(sparse_query); + zvec_multi_query_destroy(mvq2); teardown_multi_query_fixture(&f); diff --git a/tests/db/sqlengine/optimizer_test.cc b/tests/db/sqlengine/optimizer_test.cc index f726060..bc0df72 100644 --- a/tests/db/sqlengine/optimizer_test.cc +++ b/tests/db/sqlengine/optimizer_test.cc @@ -100,7 +100,7 @@ TEST_F(OptimizerTest, Basic) { query.filter_ = "age > 200"; auto engine = std::make_shared(std::make_shared()); - auto ret = engine->parse_request(schema, query, nullptr); + auto ret = engine->build_query_info(schema, query, nullptr); ASSERT_TRUE(ret.has_value()); QueryInfo::Ptr query_info = ret.value(); @@ -123,7 +123,7 @@ TEST_F(OptimizerTest, Case1) { query.filter_ = "age > 12"; auto engine = std::make_shared(std::make_shared()); - auto ret = engine->parse_request(schema, query, nullptr); + auto ret = engine->build_query_info(schema, query, nullptr); ASSERT_TRUE(ret.has_value()); QueryInfo::Ptr query_info = ret.value(); @@ -146,7 +146,7 @@ TEST_F(OptimizerTest, Case2_1) { query.filter_ = "age > 100 and age > 101 or age > 102"; auto engine = std::make_shared(std::make_shared()); - auto ret = engine->parse_request(schema, query, nullptr); + auto ret = engine->build_query_info(schema, query, nullptr); ASSERT_TRUE(ret.has_value()); QueryInfo::Ptr query_info = ret.value(); @@ -170,7 +170,7 @@ TEST_F(OptimizerTest, Case2_2) { query.filter_ = "age > 100 or age > 90"; auto engine = std::make_shared(std::make_shared()); - auto ret = engine->parse_request(schema, query, nullptr); + auto ret = engine->build_query_info(schema, query, nullptr); ASSERT_TRUE(ret.has_value()); QueryInfo::Ptr query_info = ret.value(); @@ -194,7 +194,7 @@ TEST_F(OptimizerTest, Case3_1) { query.filter_ = "age > 100 and age > 101 and age > 10"; auto engine = std::make_shared(std::make_shared()); - auto ret = engine->parse_request(schema, query, nullptr); + auto ret = engine->build_query_info(schema, query, nullptr); ASSERT_TRUE(ret.has_value()); QueryInfo::Ptr query_info = ret.value(); @@ -218,7 +218,7 @@ TEST_F(OptimizerTest, Case3_2) { query.filter_ = "age > 10 and age > 11 and age > 100"; auto engine = std::make_shared(std::make_shared()); - auto ret = engine->parse_request(schema, query, nullptr); + auto ret = engine->build_query_info(schema, query, nullptr); ASSERT_TRUE(ret.has_value()); QueryInfo::Ptr query_info = ret.value(); @@ -241,7 +241,7 @@ TEST_F(OptimizerTest, Case3_3) { query.filter_ = "(age > 10 or age > 11) and age > 100"; auto engine = std::make_shared(std::make_shared()); - auto ret = engine->parse_request(schema, query, nullptr); + auto ret = engine->build_query_info(schema, query, nullptr); ASSERT_TRUE(ret.has_value()); QueryInfo::Ptr query_info = ret.value(); @@ -265,7 +265,7 @@ TEST_F(OptimizerTest, Case3_4) { query.filter_ = "age > 10 and (age > 101 and (age > 10 and age > 10))"; auto engine = std::make_shared(std::make_shared()); - auto ret = engine->parse_request(schema, query, nullptr); + auto ret = engine->build_query_info(schema, query, nullptr); ASSERT_TRUE(ret.has_value()); QueryInfo::Ptr query_info = ret.value(); @@ -289,7 +289,7 @@ TEST_F(OptimizerTest, Case4) { query.filter_ = "age in (10, 20)"; auto engine = std::make_shared(std::make_shared()); - auto ret = engine->parse_request(schema, query, nullptr); + auto ret = engine->build_query_info(schema, query, nullptr); ASSERT_TRUE(ret.has_value()); QueryInfo::Ptr query_info = ret.value(); @@ -304,7 +304,7 @@ TEST_F(OptimizerTest, Case4) { // in and optimizable, optimize optimizable query.filter_ = "age in (10, 20) and age > 100"; - ret = engine->parse_request(schema, query, nullptr); + ret = engine->build_query_info(schema, query, nullptr); ASSERT_TRUE(ret.has_value()); query_info = ret.value(); optimized = optimizer->optimize(segment.get(), query_info.get()); @@ -312,7 +312,7 @@ TEST_F(OptimizerTest, Case4) { // in or optimizable, not optimized query.filter_ = "age in (10, 20) or age > 100"; - ret = engine->parse_request(schema, query, nullptr); + ret = engine->build_query_info(schema, query, nullptr); ASSERT_TRUE(ret.has_value()); query_info = ret.value(); optimized = optimizer->optimize(segment.get(), query_info.get()); diff --git a/tests/db/sqlengine/query_info_test.cc b/tests/db/sqlengine/query_info_test.cc index 9beb4e7..b32fcd5 100644 --- a/tests/db/sqlengine/query_info_test.cc +++ b/tests/db/sqlengine/query_info_test.cc @@ -96,7 +96,7 @@ TEST_F(QueryInfoTest, BasicQueryRequest) { query.target_.query_params_->set_radius(0.8F); auto engine = std::make_shared(std::make_shared()); - auto ret = engine->parse_request(schema, query, nullptr); + auto ret = engine->build_query_info(schema, query, nullptr); ASSERT_TRUE(ret.has_value()) << ret.error().c_str(); QueryInfo::Ptr new_query_info = ret.value(); auto &query_fields = new_query_info->query_fields(); @@ -137,7 +137,7 @@ TEST_F(QueryInfoTest, QueryRequestWithFilter) { query.filter_ = "name<3 or name=4 or 1-dash_score_field='test'"; auto engine = std::make_shared(std::make_shared()); - auto ret = engine->parse_request(schema, query, nullptr); + auto ret = engine->build_query_info(schema, query, nullptr); ASSERT_TRUE(ret.has_value()); QueryInfo::Ptr new_query_info = ret.value(); auto &query_fields = new_query_info->query_fields(); @@ -221,7 +221,7 @@ TEST_F(QueryInfoTest, QueryRequestWithIncludeVector) { query.include_vector_ = true; auto engine = std::make_shared(std::make_shared()); - auto ret = engine->parse_request(schema, query, nullptr); + auto ret = engine->build_query_info(schema, query, nullptr); ASSERT_TRUE(ret.has_value()); QueryInfo::Ptr new_query_info = ret.value(); auto &query_fields = new_query_info->query_fields(); @@ -263,7 +263,7 @@ TEST_F(QueryInfoTest, OR_ANCESTOR) { query.filter_ = "name=1 and (name=2 or name=3)"; auto engine = std::make_shared(std::make_shared()); - auto ret = engine->parse_request(schema, query, nullptr); + auto ret = engine->build_query_info(schema, query, nullptr); ASSERT_TRUE(ret.has_value()); QueryInfo::Ptr new_query_info = ret.value(); } @@ -281,7 +281,7 @@ TEST_F(QueryInfoTest, QueryRequestWithInFilter) { "name=3 or name in (1, 2, 3) or category not in (\"a\", \"b\", \"c\")"; auto engine = std::make_shared(std::make_shared()); - auto ret = engine->parse_request(schema, query, nullptr); + auto ret = engine->build_query_info(schema, query, nullptr); ASSERT_TRUE(ret.has_value()); QueryInfo::Ptr new_query_info = ret.value(); @@ -370,23 +370,23 @@ TEST_F(QueryInfoTest, QueryRequestWithInFilterWrong) { query.target_.query_params_->set_radius(0.8F); auto engine = std::make_shared(std::make_shared()); - auto ret = engine->parse_request(schema, query, nullptr); + auto ret = engine->build_query_info(schema, query, nullptr); ASSERT_TRUE(ret.has_value()); query.filter_ = ("name in ()"); - ret = engine->parse_request(schema, query, nullptr); + ret = engine->build_query_info(schema, query, nullptr); ASSERT_FALSE(ret.has_value()); query.filter_ = ("name in (\"a\", 2, 3)"); - ret = engine->parse_request(schema, query, nullptr); + ret = engine->build_query_info(schema, query, nullptr); ASSERT_FALSE(ret.has_value()); query.filter_ = ("name in (1.1, 2, 3)"); - ret = engine->parse_request(schema, query, nullptr); + ret = engine->build_query_info(schema, query, nullptr); ASSERT_FALSE(ret.has_value()); query.filter_ = ("category in (1.1, \"b\")"); - ret = engine->parse_request(schema, query, nullptr); + ret = engine->build_query_info(schema, query, nullptr); ASSERT_FALSE(ret.has_value()); } @@ -410,7 +410,7 @@ TEST_F(QueryInfoTest, QueryRequestWithInFilterNum1024) { query.filter_ = filter_str; auto engine = std::make_shared(std::make_shared()); - auto ret = engine->parse_request(schema, query, nullptr); + auto ret = engine->build_query_info(schema, query, nullptr); ASSERT_TRUE(ret.has_value()); QueryInfo::Ptr new_query_info = ret.value(); @@ -451,7 +451,7 @@ TEST_F(QueryInfoTest, QueryRequestWithFilter_contain) { )"; auto engine = std::make_shared(std::make_shared()); - auto ret = engine->parse_request(schema, query, nullptr); + auto ret = engine->build_query_info(schema, query, nullptr); ASSERT_TRUE(ret.has_value()); QueryInfo::Ptr new_query_info = ret.value(); auto &query_fields = new_query_info->query_fields(); @@ -598,7 +598,7 @@ TEST_F(QueryInfoTest, SelectNonExistField) { query.include_vector_ = false; auto engine = std::make_shared(std::make_shared()); - auto ret = engine->parse_request(schema, query, nullptr); + auto ret = engine->build_query_info(schema, query, nullptr); ASSERT_FALSE(ret.has_value()); EXPECT_THAT(ret.error().message(), testing::HasSubstr("not defined in schema")); @@ -616,7 +616,7 @@ TEST_F(QueryInfoTest, ContainAllExceedLimit) { } query.filter_ += ")"; auto engine = std::make_shared(std::make_shared()); - auto ret = engine->parse_request(schema, query, nullptr); + auto ret = engine->build_query_info(schema, query, nullptr); ASSERT_FALSE(ret.has_value()); EXPECT_THAT(ret.error().message(), testing::HasSubstr( @@ -635,7 +635,7 @@ TEST_F(QueryInfoTest, ContainAnyExceedLimit) { } query.filter_ += ")"; auto engine = std::make_shared(std::make_shared()); - auto ret = engine->parse_request(schema, query, nullptr); + auto ret = engine->build_query_info(schema, query, nullptr); ASSERT_FALSE(ret.has_value()); EXPECT_THAT(ret.error().message(), testing::HasSubstr( @@ -647,7 +647,7 @@ TEST_F(QueryInfoTest, ArrayLengthNonExistField) { query.topk_ = 200; query.filter_ = "array_length(not_exist_field) > 1"; auto engine = std::make_shared(std::make_shared()); - auto ret = engine->parse_request(schema, query, nullptr); + auto ret = engine->build_query_info(schema, query, nullptr); ASSERT_FALSE(ret.has_value()); EXPECT_THAT(ret.error().message(), testing::HasSubstr("array_length argument not found in schema")); @@ -658,7 +658,7 @@ TEST_F(QueryInfoTest, ArrayLengthOnNonArrayField) { query.topk_ = 200; query.filter_ = "array_length(name) > 1"; auto engine = std::make_shared(std::make_shared()); - auto ret = engine->parse_request(schema, query, nullptr); + auto ret = engine->build_query_info(schema, query, nullptr); ASSERT_FALSE(ret.has_value()); EXPECT_THAT(ret.error().message(), testing::HasSubstr("array_length only support array")); @@ -669,7 +669,7 @@ TEST_F(QueryInfoTest, ArrayLengthInvalidArgument) { query.topk_ = 200; query.filter_ = "array_length(name_array) > '1'"; auto engine = std::make_shared(std::make_shared()); - auto ret = engine->parse_request(schema, query, nullptr); + auto ret = engine->build_query_info(schema, query, nullptr); ASSERT_FALSE(ret.has_value()); EXPECT_THAT( ret.error().message(), @@ -681,7 +681,7 @@ TEST_F(QueryInfoTest, ArrayLengthInvalidOp) { query.topk_ = 200; query.filter_ = "array_length(name_array) like '%'"; auto engine = std::make_shared(std::make_shared()); - auto ret = engine->parse_request(schema, query, nullptr); + auto ret = engine->build_query_info(schema, query, nullptr); ASSERT_FALSE(ret.has_value()); EXPECT_THAT(ret.error().message(), testing::HasSubstr("syntax error")); } diff --git a/tests/db/sqlengine/simple_rewriter_test.cc b/tests/db/sqlengine/simple_rewriter_test.cc index ad23107..c5e32a8 100644 --- a/tests/db/sqlengine/simple_rewriter_test.cc +++ b/tests/db/sqlengine/simple_rewriter_test.cc @@ -195,7 +195,7 @@ class SimpleRewriterTest : public testing::Test { query.filter_ = filter; auto engine = std::make_shared(profiler_); - auto ret = engine->parse_request(schema, query, nullptr); + auto ret = engine->build_query_info(schema, query, nullptr); // ASSERT_TRUE(ret.has_value()); QueryInfo::Ptr new_query_info = ret.value();