From 8dcb6cbd7fcf76f3b89553728b9bcc236d59a398 Mon Sep 17 00:00:00 2001 From: egolearner Date: Fri, 29 May 2026 16:36:09 +0800 Subject: [PATCH] refactor: drop VectorQuery, unify single-target query on SearchQuery (#428) --- examples/c++/db/main.cc | 8 +- src/binding/c/c_api.cc | 82 ++++----- .../python/model/param/python_param.cc | 123 +++++++------ src/binding/python/model/python_collection.cc | 2 +- src/db/collection.cc | 46 +++-- src/db/index/common/doc.cc | 161 ----------------- src/db/index/common/query.cc | 163 ++++++++++++++++++ src/db/index/common/type_helper.cc | 43 +++++ src/db/index/common/type_helper.h | 7 + src/db/sqlengine/parser/sql_info_helper.cc | 20 ++- src/db/sqlengine/parser/sql_info_helper.h | 2 +- src/db/sqlengine/sqlengine.h | 2 +- src/db/sqlengine/sqlengine_impl.cc | 26 ++- src/db/sqlengine/sqlengine_impl.h | 4 +- src/include/zvec/c_api.h | 5 +- src/include/zvec/db/collection.h | 2 +- src/include/zvec/db/query.h | 118 ++++++++----- tests/db/collection_test.cc | 152 ++++++++-------- .../crash_recovery/optimize_recovery_test.cc | 16 +- tests/db/index/common/doc_test.cc | 107 ++++++------ tests/db/sqlengine/contain_test.cc | 28 +-- tests/db/sqlengine/forward_recall_test.cc | 70 ++++---- tests/db/sqlengine/invert_recall_test.cc | 58 +++---- tests/db/sqlengine/like_test.cc | 32 ++-- tests/db/sqlengine/optimizer_test.cc | 36 ++-- tests/db/sqlengine/query_info_test.cc | 130 +++++++------- tests/db/sqlengine/simple_rewriter_test.cc | 2 +- tests/db/sqlengine/sqlengine_test.cc | 34 ++-- tests/db/sqlengine/vector_recall_test.cc | 69 ++++---- 29 files changed, 830 insertions(+), 718 deletions(-) create mode 100644 src/db/index/common/query.cc diff --git a/examples/c++/db/main.cc b/examples/c++/db/main.cc index 3cb5bb6..2fbd36d 100644 --- a/examples/c++/db/main.cc +++ b/examples/c++/db/main.cc @@ -231,13 +231,13 @@ int main() { // query { - VectorQuery query; + SearchQuery query; query.topk_ = 10; - query.field_name_ = "dense"; + query.target_.field_name_ = "dense"; query.include_vector_ = true; std::vector query_vector = std::vector(128, 0.1); - query.query_vector_.assign((char *)query_vector.data(), - query_vector.size() * sizeof(float)); + query.target_.set_vector(std::string((char *)query_vector.data(), + query_vector.size() * sizeof(float))); auto res = coll->Query(query); if (!res.has_value()) { std::cout << res.error().message() << std::endl; diff --git a/src/binding/c/c_api.cc b/src/binding/c/c_api.cc index 170bdd9..fba5ed4 100644 --- a/src/binding/c/c_api.cc +++ b/src/binding/c/c_api.cc @@ -4840,12 +4840,13 @@ bool zvec_query_params_flat_get_is_using_refiner( } // ============================================================================= -// VectorQuery implementation - owns zvec::VectorQuery via raw pointer +// Query implementation - owns zvec::SearchQuery via raw pointer +// (external C symbol naming kept for ABI compatibility) // ============================================================================= zvec_vector_query_t *zvec_vector_query_create(void) { - ZVEC_TRY_RETURN_NULL("Failed to create VectorQuery", - auto *query = new zvec::VectorQuery(); + ZVEC_TRY_RETURN_NULL("Failed to create query object", + auto *query = new zvec::SearchQuery(); query->topk_ = 10; query->include_doc_id_ = true; query->include_vector_ = false; return reinterpret_cast(query);) @@ -4854,7 +4855,7 @@ zvec_vector_query_t *zvec_vector_query_create(void) { void zvec_vector_query_destroy(zvec_vector_query_t *query) { if (query) { - delete reinterpret_cast(query); + delete reinterpret_cast(query); } } @@ -4863,14 +4864,14 @@ zvec_error_code_t zvec_vector_query_set_topk(zvec_vector_query_t *query, int top SET_LAST_ERROR(ZVEC_ERROR_INVALID_ARGUMENT, "Vector query pointer is null"); return ZVEC_ERROR_INVALID_ARGUMENT; } - auto *ptr = reinterpret_cast(query); + auto *ptr = reinterpret_cast(query); ptr->topk_ = topk; return ZVEC_OK; } int zvec_vector_query_get_topk(const zvec_vector_query_t *query) { if (!query) return 10; - auto *ptr = reinterpret_cast(query); + auto *ptr = reinterpret_cast(query); return ptr->topk_; } @@ -4880,15 +4881,16 @@ zvec_error_code_t zvec_vector_query_set_field_name(zvec_vector_query_t *query, SET_LAST_ERROR(ZVEC_ERROR_INVALID_ARGUMENT, "Vector query pointer is null"); return ZVEC_ERROR_INVALID_ARGUMENT; } - auto *ptr = reinterpret_cast(query); - ptr->field_name_ = field_name ? field_name : ""; + auto *ptr = reinterpret_cast(query); + ptr->target_.field_name_ = field_name ? field_name : ""; return ZVEC_OK; } const char *zvec_vector_query_get_field_name(const zvec_vector_query_t *query) { if (!query) return nullptr; - auto *ptr = reinterpret_cast(query); - return ptr->field_name_.empty() ? nullptr : ptr->field_name_.c_str(); + auto *ptr = reinterpret_cast(query); + return ptr->target_.field_name_.empty() ? nullptr + : ptr->target_.field_name_.c_str(); } zvec_error_code_t zvec_vector_query_set_query_vector(zvec_vector_query_t *query, @@ -4899,8 +4901,8 @@ zvec_error_code_t zvec_vector_query_set_query_vector(zvec_vector_query_t *query, "Vector query pointer or data is null/empty"); return ZVEC_ERROR_INVALID_ARGUMENT; } - auto *ptr = reinterpret_cast(query); - ptr->query_vector_.assign(static_cast(data), size); + auto *ptr = reinterpret_cast(query); + ptr->target_.set_vector(std::string(static_cast(data), size)); return ZVEC_OK; } @@ -4910,14 +4912,14 @@ zvec_error_code_t zvec_vector_query_set_filter(zvec_vector_query_t *query, SET_LAST_ERROR(ZVEC_ERROR_INVALID_ARGUMENT, "Vector query pointer is null"); return ZVEC_ERROR_INVALID_ARGUMENT; } - auto *ptr = reinterpret_cast(query); + auto *ptr = reinterpret_cast(query); ptr->filter_ = filter ? filter : ""; return ZVEC_OK; } const char *zvec_vector_query_get_filter(const zvec_vector_query_t *query) { if (!query) return nullptr; - auto *ptr = reinterpret_cast(query); + auto *ptr = reinterpret_cast(query); return ptr->filter_.empty() ? nullptr : ptr->filter_.c_str(); } @@ -4927,14 +4929,14 @@ zvec_error_code_t zvec_vector_query_set_include_vector(zvec_vector_query_t *quer SET_LAST_ERROR(ZVEC_ERROR_INVALID_ARGUMENT, "Vector query pointer is null"); return ZVEC_ERROR_INVALID_ARGUMENT; } - auto *ptr = reinterpret_cast(query); + auto *ptr = reinterpret_cast(query); ptr->include_vector_ = include; return ZVEC_OK; } bool zvec_vector_query_get_include_vector(const zvec_vector_query_t *query) { if (!query) return false; - auto *ptr = reinterpret_cast(query); + auto *ptr = reinterpret_cast(query); return ptr->include_vector_; } @@ -4944,14 +4946,14 @@ zvec_error_code_t zvec_vector_query_set_include_doc_id(zvec_vector_query_t *quer SET_LAST_ERROR(ZVEC_ERROR_INVALID_ARGUMENT, "Vector query pointer is null"); return ZVEC_ERROR_INVALID_ARGUMENT; } - auto *ptr = reinterpret_cast(query); + auto *ptr = reinterpret_cast(query); ptr->include_doc_id_ = include; return ZVEC_OK; } bool zvec_vector_query_get_include_doc_id(const zvec_vector_query_t *query) { if (!query) return false; - auto *ptr = reinterpret_cast(query); + auto *ptr = reinterpret_cast(query); return ptr->include_doc_id_; } @@ -4962,7 +4964,7 @@ zvec_error_code_t zvec_vector_query_set_output_fields(zvec_vector_query_t *query SET_LAST_ERROR(ZVEC_ERROR_INVALID_ARGUMENT, "Vector query pointer is null"); return ZVEC_ERROR_INVALID_ARGUMENT; } - auto *ptr = reinterpret_cast(query); + auto *ptr = reinterpret_cast(query); if (!fields || count == 0) { ptr->output_fields_ = std::nullopt; } else { @@ -4984,7 +4986,7 @@ zvec_error_code_t zvec_vector_query_get_output_fields(const zvec_vector_query_t "Query, fields, or count pointer is null"); return ZVEC_ERROR_INVALID_ARGUMENT; } - auto *ptr = reinterpret_cast(query); + auto *ptr = reinterpret_cast(query); if (!ptr->output_fields_.has_value()) { *fields = nullptr; @@ -5007,7 +5009,7 @@ zvec_error_code_t zvec_vector_query_get_output_fields(const zvec_vector_query_t // ============================================================================= // Type-safe query params attachment functions (transfer ownership to -// VectorQuery) +// query object) // ============================================================================= zvec_error_code_t zvec_vector_query_set_query_params(zvec_vector_query_t *query, @@ -5021,12 +5023,11 @@ zvec_error_code_t zvec_vector_query_set_query_params(zvec_vector_query_t *query, return ZVEC_ERROR_INVALID_ARGUMENT; } - auto *query_ptr = reinterpret_cast(query); + auto *query_ptr = reinterpret_cast(query); - // Cast to QueryParams* and transfer ownership via shared_ptr - // The params pointer is now owned by VectorQuery's shared_ptr + // Cast to QueryParams* and transfer ownership via shared_ptr. auto *params_ptr = reinterpret_cast(params); - query_ptr->query_params_.reset(params_ptr); + query_ptr->target_.query_params_.reset(params_ptr); return ZVEC_OK; } @@ -5040,11 +5041,11 @@ zvec_error_code_t zvec_vector_query_set_hnsw_params( return ZVEC_ERROR_INVALID_ARGUMENT; } - auto *query_ptr = reinterpret_cast(query); + auto *query_ptr = reinterpret_cast(query); auto *params_ptr = reinterpret_cast(hnsw_params); // Transfer ownership via shared_ptr (polymorphic conversion) - query_ptr->query_params_.reset(params_ptr); + query_ptr->target_.query_params_.reset(params_ptr); return ZVEC_OK; } @@ -5057,10 +5058,10 @@ zvec_error_code_t zvec_vector_query_set_ivf_params(zvec_vector_query_t *query, return ZVEC_ERROR_INVALID_ARGUMENT; } - auto *query_ptr = reinterpret_cast(query); + auto *query_ptr = reinterpret_cast(query); auto *params_ptr = reinterpret_cast(ivf_params); - query_ptr->query_params_.reset(params_ptr); + query_ptr->target_.query_params_.reset(params_ptr); return ZVEC_OK; } @@ -5073,10 +5074,10 @@ zvec_error_code_t zvec_vector_query_set_flat_params( return ZVEC_ERROR_INVALID_ARGUMENT; } - auto *query_ptr = reinterpret_cast(query); + auto *query_ptr = reinterpret_cast(query); auto *params_ptr = reinterpret_cast(flat_params); - query_ptr->query_params_.reset(params_ptr); + query_ptr->target_.query_params_.reset(params_ptr); return ZVEC_OK; } @@ -5109,7 +5110,7 @@ zvec_error_code_t zvec_group_by_vector_query_set_field_name( return ZVEC_ERROR_INVALID_ARGUMENT; } auto *ptr = reinterpret_cast(query); - ptr->field_name_ = field_name ? field_name : ""; + ptr->target_.field_name_ = field_name ? field_name : ""; return ZVEC_OK; } @@ -5117,7 +5118,8 @@ const char *zvec_group_by_vector_query_get_field_name( const zvec_group_by_vector_query_t *query) { if (!query) return nullptr; auto *ptr = reinterpret_cast(query); - return ptr->field_name_.empty() ? nullptr : ptr->field_name_.c_str(); + return ptr->target_.field_name_.empty() ? nullptr + : ptr->target_.field_name_.c_str(); } zvec_error_code_t zvec_group_by_vector_query_set_group_by_field_name( @@ -5186,7 +5188,7 @@ zvec_error_code_t zvec_group_by_vector_query_set_query_vector( return ZVEC_ERROR_INVALID_ARGUMENT; } auto *ptr = reinterpret_cast(query); - ptr->query_vector_.assign(static_cast(data), size); + ptr->target_.set_vector(std::string(static_cast(data), size)); return ZVEC_OK; } @@ -5288,7 +5290,7 @@ zvec_error_code_t zvec_group_by_vector_query_set_query_params( auto *query_ptr = reinterpret_cast(query); auto *params_ptr = reinterpret_cast(params); - query_ptr->query_params_.reset(params_ptr); + query_ptr->target_.query_params_.reset(params_ptr); return ZVEC_OK; } @@ -5305,7 +5307,7 @@ zvec_error_code_t zvec_group_by_vector_query_set_hnsw_params( auto *query_ptr = reinterpret_cast(query); auto *params_ptr = reinterpret_cast(hnsw_params); - query_ptr->query_params_.reset(params_ptr); + query_ptr->target_.query_params_.reset(params_ptr); return ZVEC_OK; } @@ -5321,7 +5323,7 @@ zvec_error_code_t zvec_group_by_vector_query_set_ivf_params( auto *query_ptr = reinterpret_cast(query); auto *params_ptr = reinterpret_cast(ivf_params); - query_ptr->query_params_.reset(params_ptr); + query_ptr->target_.query_params_.reset(params_ptr); return ZVEC_OK; } @@ -5337,7 +5339,7 @@ zvec_error_code_t zvec_group_by_vector_query_set_flat_params( auto *query_ptr = reinterpret_cast(query); auto *params_ptr = reinterpret_cast(flat_params); - query_ptr->query_params_.reset(params_ptr); + query_ptr->target_.query_params_.reset(params_ptr); return ZVEC_OK; } @@ -6338,9 +6340,9 @@ zvec_error_code_t zvec_collection_query(const zvec_collection_t *collection, reinterpret_cast *>( collection); - // Cast zvec_vector_query_t* to zvec::VectorQuery* directly + // zvec_vector_query_t wraps zvec::SearchQuery internally. auto *internal_query = - reinterpret_cast(query); + reinterpret_cast(query); auto result = (*coll_ptr)->Query(*internal_query); zvec_error_code_t error_code = handle_expected_result(result); diff --git a/src/binding/python/model/param/python_param.cc b/src/binding/python/model/param/python_param.cc index 86376db..6a6e48f 100644 --- a/src/binding/python/model/param/python_param.cc +++ b/src/binding/python/model/param/python_param.cc @@ -1379,31 +1379,39 @@ void ZVecPyParams::bind_vector_query(py::module_ &m) { .def_readwrite("num_candidates", &SubQuery::num_candidates_) .def_static( "from_vector_query", - [](const VectorQuery &vq) { + [](const SearchQuery &sq) { SubQuery sub; - sub.num_candidates_ = vq.topk_; - sub.target_.field_name_ = vq.field_name_; - auto &clause = std::get(sub.target_.clause_); - clause.query_vector_ = vq.query_vector_; - clause.sparse_indices_ = vq.query_sparse_indices_; - clause.sparse_values_ = vq.query_sparse_values_; - sub.target_.query_params_ = vq.query_params_; + sub.num_candidates_ = sq.topk_; + sub.target_ = sq.target_; return sub; }, - py::arg("vector_query"), "Create a SubQuery from a VectorQuery"); + py::arg("vector_query"), + "Create a SubQuery from a single-target search query."); - py::class_(m, "_VectorQuery") + // _VectorQuery is the historical Python class name; it now wraps the + // single-target SearchQuery so external Python code keeps working unchanged. + py::class_(m, "_VectorQuery") .def(py::init<>()) // properties - .def_readwrite("topk", &VectorQuery::topk_) - .def_readwrite("field_name", &VectorQuery::field_name_) - .def_readwrite("filter", &VectorQuery::filter_) - .def_readwrite("include_vector", &VectorQuery::include_vector_) - .def_readwrite("query_params", &VectorQuery::query_params_) - .def_readwrite("output_fields", &VectorQuery::output_fields_) + .def_readwrite("topk", &SearchQuery::topk_) + .def_property( + "field_name", + [](const SearchQuery &s) { return s.target_.field_name_; }, + [](SearchQuery &s, std::string v) { + s.target_.field_name_ = std::move(v); + }) + .def_readwrite("filter", &SearchQuery::filter_) + .def_readwrite("include_vector", &SearchQuery::include_vector_) + .def_property( + "query_params", + [](const SearchQuery &s) { return s.target_.query_params_; }, + [](SearchQuery &s, QueryParams::Ptr p) { + s.target_.query_params_ = std::move(p); + }) + .def_readwrite("output_fields", &SearchQuery::output_fields_) // vector .def("set_vector", - [](VectorQuery &self, const FieldSchema &field_schema, + [](SearchQuery &self, const FieldSchema &field_schema, const py::object &obj) { const DataType data_type = field_schema.data_type(); @@ -1422,23 +1430,23 @@ void ZVecPyParams::bind_vector_query(py::module_ &m) { const auto buf = arr.request(); switch (data_type) { case DataType::VECTOR_FP32: { - self.query_vector_ = serialize_vector( - static_cast(buf.ptr), buf.size); + self.target_.set_vector(serialize_vector( + static_cast(buf.ptr), buf.size)); return; } case DataType::VECTOR_FP64: { - self.query_vector_ = serialize_vector( - static_cast(buf.ptr), buf.size); + self.target_.set_vector(serialize_vector( + static_cast(buf.ptr), buf.size)); return; } case DataType::VECTOR_INT8: { - self.query_vector_ = serialize_vector( - static_cast(buf.ptr), buf.size); + self.target_.set_vector(serialize_vector( + static_cast(buf.ptr), buf.size)); return; } case DataType::VECTOR_FP16: { - self.query_vector_ = serialize_vector( - static_cast(buf.ptr), buf.size); + self.target_.set_vector(serialize_vector( + static_cast(buf.ptr), buf.size)); return; } default: @@ -1466,8 +1474,8 @@ void ZVecPyParams::bind_vector_query(py::module_ &m) { "FLOAT"); return ailego::Float16(f); }); - self.query_sparse_indices_ = std::move(indices); - self.query_sparse_values_ = std::move(values); + self.target_.set_sparse_vector(std::move(indices), + std::move(values)); break; } case DataType::SPARSE_VECTOR_FP32: { @@ -1477,8 +1485,8 @@ void ZVecPyParams::bind_vector_query(py::module_ &m) { h, "Sparse value[" + std::to_string(idx) + "]", "FLOAT"); }); - self.query_sparse_indices_ = std::move(indices); - self.query_sparse_values_ = std::move(values); + self.target_.set_sparse_vector(std::move(indices), + std::move(values)); break; } default: @@ -1494,16 +1502,17 @@ void ZVecPyParams::bind_vector_query(py::module_ &m) { }) .def( "get_vector", - [](const VectorQuery &self, + [](const SearchQuery &self, const FieldSchema &field_schema) -> py::object { DataType data_type = field_schema.data_type(); + const VectorClause *vc = self.target_.get_vector_clause(); if (FieldSchema::is_dense_vector_field(data_type)) { - if (self.query_vector_.empty()) { + if (vc == nullptr || vc->query_vector_.empty()) { throw std::runtime_error("No dense vector has been set"); } - size_t byte_size = self.query_vector_.size(); - const void *data = self.query_vector_.data(); + size_t byte_size = vc->query_vector_.size(); + const void *data = vc->query_vector_.data(); switch (data_type) { case DataType::VECTOR_FP32: { @@ -1549,29 +1558,29 @@ void ZVecPyParams::bind_vector_query(py::module_ &m) { } } if (FieldSchema::is_sparse_vector_field(data_type)) { - if (self.query_sparse_indices_.empty()) { + if (vc == nullptr || vc->sparse_indices_.empty()) { return py::dict(); } // Deserialize indices: stored as uint32_t[] - size_t indices_byte_size = self.query_sparse_indices_.size(); + size_t indices_byte_size = vc->sparse_indices_.size(); if (indices_byte_size % sizeof(uint32_t) != 0) { throw std::runtime_error( "Sparse indices buffer size not aligned to uint32_t"); } size_t n = indices_byte_size / sizeof(uint32_t); const uint32_t *indices = reinterpret_cast( - self.query_sparse_indices_.data()); + vc->sparse_indices_.data()); // Deserialize values switch (data_type) { case DataType::SPARSE_VECTOR_FP32: { - if (self.query_sparse_values_.size() != n * sizeof(float)) { + if (vc->sparse_values_.size() != n * sizeof(float)) { throw std::runtime_error( "Sparse FP32 values buffer size mismatch"); } const float *values = reinterpret_cast( - self.query_sparse_values_.data()); + vc->sparse_values_.data()); py::dict result; for (size_t i = 0; i < n; ++i) { result[py::int_(indices[i])] = py::float_(values[i]); @@ -1579,13 +1588,12 @@ void ZVecPyParams::bind_vector_query(py::module_ &m) { return result; } case DataType::SPARSE_VECTOR_FP16: { - if (self.query_sparse_values_.size() != - n * sizeof(uint16_t)) { + if (vc->sparse_values_.size() != n * sizeof(uint16_t)) { throw std::runtime_error( "Sparse FP16 values buffer size mismatch"); } const uint16_t *raw_bits = reinterpret_cast( - self.query_sparse_values_.data()); + vc->sparse_values_.data()); py::dict result; for (size_t i = 0; i < n; ++i) { float f = ailego::FloatHelper::ToFP32(raw_bits[i]); @@ -1604,31 +1612,36 @@ void ZVecPyParams::bind_vector_query(py::module_ &m) { }, py::arg("field_schema")) .def(py::pickle( - [](const VectorQuery &self) { - return py::make_tuple( - self.topk_, self.field_name_, self.query_vector_, - self.query_sparse_indices_, self.query_sparse_values_, - self.filter_, self.include_vector_, self.output_fields_, - self.query_params_ ? py::cast(self.query_params_) : py::none()); + [](const SearchQuery &self) { + const VectorClause *vc = self.target_.get_vector_clause(); + return py::make_tuple(self.topk_, self.target_.field_name_, + vc ? vc->query_vector_ : std::string(), + vc ? vc->sparse_indices_ : std::string(), + vc ? vc->sparse_values_ : std::string(), + self.filter_, self.include_vector_, + self.output_fields_, + self.target_.query_params_ + ? py::cast(self.target_.query_params_) + : py::none()); }, [](py::tuple t) { if (t.size() != 9) - throw std::runtime_error("Invalid pickle data for VectorQuery"); + throw std::runtime_error("Invalid pickle data for _VectorQuery"); - VectorQuery obj{}; + SearchQuery obj{}; obj.topk_ = t[0].cast(); - obj.field_name_ = t[1].cast(); - obj.query_vector_ = t[2].cast(); - obj.query_sparse_indices_ = t[3].cast(); - obj.query_sparse_values_ = t[4].cast(); + obj.target_.field_name_ = t[1].cast(); + obj.target_.clause_ = + VectorClause{t[2].cast(), t[3].cast(), + t[4].cast()}; obj.filter_ = t[5].cast(); obj.include_vector_ = t[6].cast(); obj.output_fields_ = t[7].cast>(); if (!t[8].is_none()) { - obj.query_params_ = t[8].cast(); + obj.target_.query_params_ = t[8].cast(); } return obj; })); } -} // namespace zvec \ No newline at end of file +} // namespace zvec diff --git a/src/binding/python/model/python_collection.cc b/src/binding/python/model/python_collection.cc index d902e02..b1311f1 100644 --- a/src/binding/python/model/python_collection.cc +++ b/src/binding/python/model/python_collection.cc @@ -252,7 +252,7 @@ void ZVecPyCollection::bind_dml_methods( void ZVecPyCollection::bind_dql_methods( py::class_ &col) { col.def("Query", - [](const Collection &self, const VectorQuery &query) { + [](const Collection &self, const SearchQuery &query) { Result result; { py::gil_scoped_release release; diff --git a/src/db/collection.cc b/src/db/collection.cc index 05bde16..8f76ea3 100644 --- a/src/db/collection.cc +++ b/src/db/collection.cc @@ -118,7 +118,7 @@ class CollectionImpl : public Collection { Status DeleteByFilter(const std::string &filter) override; - Result Query(const VectorQuery &query) const override; + Result Query(const SearchQuery &query) const override; Result Query(const MultiQuery &query) const override; @@ -1560,7 +1560,7 @@ Status CollectionImpl::DeleteByFilter(const std::string &filter) { CHECK_DESTROY_RETURN_STATUS(destroyed_, false); - VectorQuery query; + SearchQuery query; query.filter_ = filter; query.topk_ = INT32_MAX; query.output_fields_ = std::vector{}; @@ -1584,14 +1584,14 @@ Status CollectionImpl::DeleteByFilter(const std::string &filter) { return Status::OK(); } -Result CollectionImpl::Query(const VectorQuery &query) const { +Result CollectionImpl::Query(const SearchQuery &query) const { std::shared_lock lock(schema_handle_mtx_); CHECK_DESTROY_RETURN_STATUS_EXPECTED(destroyed_, false); - VectorQuery sanitized = query; + SearchQuery sanitized = query; auto s = sanitized.validate_and_sanitize( - schema_->get_vector_field(sanitized.field_name_)); + schema_->get_vector_field(sanitized.target_.field_name_)); CHECK_RETURN_STATUS_EXPECTED(s); auto segments = get_all_segments(); @@ -1622,9 +1622,9 @@ Result CollectionImpl::Query(const MultiQuery &query) const { return DocPtrList(); } - // Convert SubVectorQuery to VectorQuery and validate + // Convert each SubQuery to a SearchQuery and validate. std::set seen_fields; - std::vector converted_queries; + std::vector converted_queries; converted_queries.reserve(query.queries.size()); for (const auto &sub : query.queries) { @@ -1640,31 +1640,26 @@ Result CollectionImpl::Query(const MultiQuery &query) const { "Vector field not found: ", target.field_name_)); } - VectorQuery vq; - vq.topk_ = sub.num_candidates_; - vq.field_name_ = target.field_name_; - const auto &vec_clause = std::get(target.clause_); - vq.query_vector_ = vec_clause.query_vector_; - vq.query_sparse_indices_ = vec_clause.sparse_indices_; - vq.query_sparse_values_ = vec_clause.sparse_values_; - vq.query_params_ = target.query_params_; - vq.filter_ = query.filter; - vq.include_vector_ = query.include_vector; - vq.include_doc_id_ = query.include_doc_id_; - vq.output_fields_ = query.output_fields; + SearchQuery sq; + sq.target_ = target; + sq.topk_ = sub.num_candidates_; + sq.filter_ = query.filter; + sq.include_vector_ = query.include_vector; + sq.include_doc_id_ = query.include_doc_id_; + sq.output_fields_ = query.output_fields; - auto s = vq.validate_and_sanitize(field_schema); + auto s = sq.validate_and_sanitize(field_schema); CHECK_RETURN_STATUS_EXPECTED(s); - converted_queries.push_back(std::move(vq)); + converted_queries.push_back(std::move(sq)); } - // Execute each VectorQuery concurrently and collect results per field + // Execute each sub-query concurrently and collect results per field. std::vector>> futures; futures.reserve(converted_queries.size()); - for (const auto &vq : converted_queries) { + 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_, vq, segments); + return engine->execute(schema_, sq, segments); })); } @@ -1674,7 +1669,8 @@ Result CollectionImpl::Query(const MultiQuery &query) const { if (!result.has_value()) { return tl::make_unexpected(result.error()); } - query_results[converted_queries[i].field_name_] = std::move(result.value()); + query_results[converted_queries[i].target_.field_name_] = + std::move(result.value()); } // Merge and rerank results diff --git a/src/db/index/common/doc.cc b/src/db/index/common/doc.cc index fe8be10..8948ca4 100644 --- a/src/db/index/common/doc.cc +++ b/src/db/index/common/doc.cc @@ -187,45 +187,6 @@ template overloaded(Ts...) -> overloaded; -bool sort_and_find_duplicates(uint32_t *indices, char *values, size_t n, - size_t value_byte_size) { - if (n <= 1) { - return false; - } - bool already_sorted = true; - for (size_t i = 1; i < n; ++i) { - if (indices[i] == indices[i - 1]) { - return true; - } - if (indices[i] < indices[i - 1]) { - already_sorted = false; - break; - } - } - if (already_sorted) { - return false; - } - std::vector perm(n); - std::iota(perm.begin(), perm.end(), size_t{0}); - std::sort(perm.begin(), perm.end(), - [&](size_t a, size_t b) { return indices[a] < indices[b]; }); - std::vector sorted_indices(n); - std::vector sorted_values(n * value_byte_size); - for (size_t i = 0; i < n; ++i) { - sorted_indices[i] = indices[perm[i]]; - std::memcpy(sorted_values.data() + i * value_byte_size, - values + perm[i] * value_byte_size, value_byte_size); - } - std::memcpy(indices, sorted_indices.data(), n * sizeof(uint32_t)); - std::memcpy(values, sorted_values.data(), n * value_byte_size); - for (size_t i = 1; i < n; ++i) { - if (indices[i] == indices[i - 1]) { - return true; - } - } - return false; -} - } // namespace @@ -1269,126 +1230,4 @@ bool Doc::operator==(const Doc &other) const { return true; } -Status VectorQuery::validate_and_sanitize(const FieldSchema *schema) { - if ((uint32_t)topk_ > kMaxQueryTopk) { - return Status::InvalidArgument("Invalid query: topk[", topk_, - "] exceeds the maximum allowed value of ", - kMaxQueryTopk); - } - if (output_fields_.has_value() && - output_fields_->size() > kMaxOutputFieldSize) { - return Status::InvalidArgument( - "Invalid query: too many output fields, the maximum allowed is ", - kMaxOutputFieldSize); - } - - if (schema == nullptr) { - if (query_vector_.empty() && query_sparse_indices_.empty()) { - // Scalar-only filter query - return Status::OK(); - } else { - // If a query vector was provided, the field must exist as a vector field - // since we are performing a vector similarity search. - return Status::InvalidArgument( - "Invalid query: query vector is provided, but query field[", - field_name_, - "] does not exist or is not a vector field in the collection"); - } - } - - // Vector query - if (schema->is_dense_vector()) { - // Validate dimension - auto dim = schema->dimension(); - switch (schema->data_type()) { - case DataType::VECTOR_FP16: - if (dim * sizeof(float16_t) != query_vector_.size()) { - return Status::InvalidArgument( - "Invalid query: dimension mismatch, expected ", dim, " but got ", - query_vector_.size() / sizeof(float16_t), " (FP16)"); - } - break; - case DataType::VECTOR_FP32: - if (dim * sizeof(float) != query_vector_.size()) { - return Status::InvalidArgument( - "Invalid query: dimension mismatch, expected ", dim, " but got ", - query_vector_.size() / sizeof(float), " (FP32)"); - } - break; - case DataType::VECTOR_FP64: - if (dim * sizeof(double) != query_vector_.size()) { - return Status::InvalidArgument( - "Invalid query: dimension mismatch, expected ", dim, " but got ", - query_vector_.size() / sizeof(double), " (FP64)"); - } - break; - case DataType::VECTOR_INT8: - if (dim * sizeof(int8_t) != query_vector_.size()) { - return Status::InvalidArgument( - "Invalid query: dimension mismatch, expected ", dim, " but got ", - query_vector_.size() / sizeof(int8_t), " (INT8)"); - } - break; - case DataType::VECTOR_INT16: - case DataType::VECTOR_INT4: - case DataType::VECTOR_BINARY32: - case DataType::VECTOR_BINARY64: - return Status::NotSupported( - "Invalid query: dense vector type of field[", field_name_, - "] is not supported"); - default: - return Status::InvalidArgument("Invalid query: field[", field_name_, - "] is not a dense vector field"); - } - } else if (schema->is_sparse_vector()) { - size_t value_byte_size = 0; - switch (schema->data_type()) { - case DataType::SPARSE_VECTOR_FP32: - value_byte_size = sizeof(float); - break; - case DataType::SPARSE_VECTOR_FP16: - value_byte_size = sizeof(float16_t); - break; - default: - return Status::InvalidArgument( - "Invalid query: sparse vector type of field[", field_name_, - "] is not supported"); - } - if (query_sparse_indices_.size() % sizeof(uint32_t) != 0 || - query_sparse_values_.size() % value_byte_size != 0 || - query_sparse_indices_.size() / sizeof(uint32_t) != - query_sparse_values_.size() / value_byte_size) { - return Status::InvalidArgument( - "Invalid query: sparse vector query for field[", field_name_, - "] has mismatched indices and values sizes"); - } - size_t n_indices = query_sparse_indices_.size() / sizeof(uint32_t); - if (n_indices > kSparseMaxDimSize) { - return Status::InvalidArgument( - "Invalid query: too many sparse indices, the maximum allowed is ", - kSparseMaxDimSize); - } - if (sort_and_find_duplicates( - reinterpret_cast(query_sparse_indices_.data()), - query_sparse_values_.data(), n_indices, value_byte_size)) { - return Status::InvalidArgument( - "Invalid query: sparse vector query for field[", field_name_, - "] contains duplicate indices"); - } - } else { - return Status::InvalidArgument("Invalid query: field[", field_name_, - "] is not a vector field"); - } - // Validate query_params type - if (query_params_ && query_params_->type() != schema->index_type()) { - return Status::InvalidArgument( - "Invalid query: query params type does not match the index type of " - "vector field[", - field_name_, "], expected ", - IndexTypeCodeBook::AsString(schema->index_type()), " but got ", - IndexTypeCodeBook::AsString(query_params_->type())); - } - return Status::OK(); -} - } // namespace zvec diff --git a/src/db/index/common/query.cc b/src/db/index/common/query.cc new file mode 100644 index 0000000..273bb4a --- /dev/null +++ b/src/db/index/common/query.cc @@ -0,0 +1,163 @@ +// Copyright 2025-present the zvec project +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +#include +#include +#include +#include "db/common/constants.h" +#include "db/index/common/type_helper.h" + +namespace zvec { + +Status SearchQuery::validate_and_sanitize(const FieldSchema *schema) { + if ((uint32_t)topk_ > kMaxQueryTopk) { + return Status::InvalidArgument("Invalid query: topk[", topk_, + "] exceeds the maximum allowed value of ", + kMaxQueryTopk); + } + if (output_fields_.has_value() && + output_fields_->size() > kMaxOutputFieldSize) { + return Status::InvalidArgument( + "Invalid query: too many output fields, the maximum allowed is ", + kMaxOutputFieldSize); + } + + auto *vc = target_.get_vector_clause(); + auto &field_name = target_.field_name_; + // A "scalar-only filter" query has no vector payload — either the clause + // is not a VectorClause (e.g., FtsClause) or its fields are all empty. + bool no_vector_payload = (vc == nullptr) || (vc->query_vector_.empty() && + vc->sparse_indices_.empty()); + + if (schema == nullptr) { + if (no_vector_payload) { + // Scalar-only filter query + return Status::OK(); + } else { + // If a query vector was provided, the field must exist as a vector field + // since we are performing a vector similarity search. + return Status::InvalidArgument( + "Invalid query: query vector is provided, but query field[", + field_name, + "] does not exist or is not a vector field in the collection"); + } + } + + // Schema is non-null from here on: a vector payload is required. + if (no_vector_payload) { + return Status::InvalidArgument( + "Invalid query: missing query clause for field[", field_name, "]"); + } + + auto &query_vector = vc->query_vector_; + auto &query_sparse_indices = vc->sparse_indices_; + auto &query_sparse_values = vc->sparse_values_; + auto &query_params = target_.query_params_; + + // Vector query + if (schema->is_dense_vector()) { + // Validate dimension + auto dim = schema->dimension(); + switch (schema->data_type()) { + case DataType::VECTOR_FP16: + if (dim * sizeof(float16_t) != query_vector.size()) { + return Status::InvalidArgument( + "Invalid query: dimension mismatch, expected ", dim, " but got ", + query_vector.size() / sizeof(float16_t), " (FP16)"); + } + break; + case DataType::VECTOR_FP32: + if (dim * sizeof(float) != query_vector.size()) { + return Status::InvalidArgument( + "Invalid query: dimension mismatch, expected ", dim, " but got ", + query_vector.size() / sizeof(float), " (FP32)"); + } + break; + case DataType::VECTOR_FP64: + if (dim * sizeof(double) != query_vector.size()) { + return Status::InvalidArgument( + "Invalid query: dimension mismatch, expected ", dim, " but got ", + query_vector.size() / sizeof(double), " (FP64)"); + } + break; + case DataType::VECTOR_INT8: + if (dim * sizeof(int8_t) != query_vector.size()) { + return Status::InvalidArgument( + "Invalid query: dimension mismatch, expected ", dim, " but got ", + query_vector.size() / sizeof(int8_t), " (INT8)"); + } + break; + case DataType::VECTOR_INT16: + case DataType::VECTOR_INT4: + case DataType::VECTOR_BINARY32: + case DataType::VECTOR_BINARY64: + return Status::NotSupported( + "Invalid query: dense vector type of field[", field_name, + "] is not supported"); + default: + return Status::InvalidArgument("Invalid query: field[", field_name, + "] is not a dense vector field"); + } + } else if (schema->is_sparse_vector()) { + size_t value_byte_size = 0; + switch (schema->data_type()) { + case DataType::SPARSE_VECTOR_FP32: + value_byte_size = sizeof(float); + break; + case DataType::SPARSE_VECTOR_FP16: + value_byte_size = sizeof(float16_t); + break; + default: + return Status::InvalidArgument( + "Invalid query: sparse vector type of field[", field_name, + "] is not supported"); + } + if (query_sparse_indices.size() % sizeof(uint32_t) != 0 || + query_sparse_values.size() % value_byte_size != 0 || + query_sparse_indices.size() / sizeof(uint32_t) != + query_sparse_values.size() / value_byte_size) { + return Status::InvalidArgument( + "Invalid query: sparse vector query for field[", field_name, + "] has mismatched indices and values sizes"); + } + size_t n_indices = query_sparse_indices.size() / sizeof(uint32_t); + if (n_indices > kSparseMaxDimSize) { + return Status::InvalidArgument( + "Invalid query: too many sparse indices, the maximum allowed is ", + kSparseMaxDimSize); + } + if (sort_and_find_duplicates( + reinterpret_cast(query_sparse_indices.data()), + query_sparse_values.data(), n_indices, value_byte_size)) { + return Status::InvalidArgument( + "Invalid query: sparse vector query for field[", field_name, + "] contains duplicate indices"); + } + } else { + return Status::InvalidArgument("Invalid query: field[", field_name, + "] is not a vector field"); + } + // Validate query_params type + if (query_params && query_params->type() != schema->index_type()) { + return Status::InvalidArgument( + "Invalid query: query params type does not match the index type of " + "vector field[", + field_name, "], expected ", + IndexTypeCodeBook::AsString(schema->index_type()), " but got ", + IndexTypeCodeBook::AsString(query_params->type())); + } + return Status::OK(); +} + +} // namespace zvec diff --git a/src/db/index/common/type_helper.cc b/src/db/index/common/type_helper.cc index 7f622bb..45f2c24 100644 --- a/src/db/index/common/type_helper.cc +++ b/src/db/index/common/type_helper.cc @@ -13,10 +13,53 @@ // limitations under the License. #include "type_helper.h" +#include +#include +#include +#include #include namespace zvec { +bool sort_and_find_duplicates(uint32_t *indices, char *values, size_t n, + size_t value_byte_size) { + if (n <= 1) { + return false; + } + bool already_sorted = true; + for (size_t i = 1; i < n; ++i) { + if (indices[i] == indices[i - 1]) { + return true; + } + if (indices[i] < indices[i - 1]) { + already_sorted = false; + break; + } + } + if (already_sorted) { + return false; + } + std::vector perm(n); + std::iota(perm.begin(), perm.end(), size_t{0}); + std::sort(perm.begin(), perm.end(), + [&](size_t a, size_t b) { return indices[a] < indices[b]; }); + std::vector sorted_indices(n); + std::vector sorted_values(n * value_byte_size); + for (size_t i = 0; i < n; ++i) { + sorted_indices[i] = indices[perm[i]]; + std::memcpy(sorted_values.data() + i * value_byte_size, + values + perm[i] * value_byte_size, value_byte_size); + } + std::memcpy(indices, sorted_indices.data(), n * sizeof(uint32_t)); + std::memcpy(values, sorted_values.data(), n * value_byte_size); + for (size_t i = 1; i < n; ++i) { + if (indices[i] == indices[i - 1]) { + return true; + } + } + return false; +} + core::IndexMeta::DataType DataTypeCodeBook::to_data_type(DataType type) { switch (type) { case DataType::VECTOR_FP32: diff --git a/src/db/index/common/type_helper.h b/src/db/index/common/type_helper.h index 02b7c0b..bda5075 100644 --- a/src/db/index/common/type_helper.h +++ b/src/db/index/common/type_helper.h @@ -14,12 +14,19 @@ #pragma once +#include +#include #include #include #include "proto/zvec.pb.h" namespace zvec { +//! Sort sparse (indices, values) pairs in place by index ascending and report +//! whether any duplicate index exists. value_byte_size is the per-value stride. +bool sort_and_find_duplicates(uint32_t *indices, char *values, size_t n, + size_t value_byte_size); + //! Index Type Codebook struct IndexTypeCodeBook { //! convert protobuf IndexType to C++ IndexType diff --git a/src/db/sqlengine/parser/sql_info_helper.cc b/src/db/sqlengine/parser/sql_info_helper.cc index ae7a2cc..6ddd3a5 100644 --- a/src/db/sqlengine/parser/sql_info_helper.cc +++ b/src/db/sqlengine/parser/sql_info_helper.cc @@ -26,16 +26,20 @@ namespace zvec::sqlengine { using namespace zvec; -Node::Ptr handle_vector(const VectorQuery &request, std::string * /*err_msg*/) { +Node::Ptr handle_vector(const SearchQuery &request, std::string * /*err_msg*/) { + const 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.field_name_)); + rel_exp->set_left(std::make_shared(request.target_.field_name_)); rel_exp->set_right(std::make_shared( - request.query_vector_, request.query_sparse_indices_, - request.query_sparse_values_, request.query_params_)); + vc->query_vector_, vc->sparse_indices_, vc->sparse_values_, + request.target_.query_params_)); return rel_exp; } -void handle_query_field(const VectorQuery *query, SelectInfo *selected_info) { +void handle_query_field(const SearchQuery *query, SelectInfo *selected_info) { if (!query->output_fields_.has_value()) { SelectedElemInfo::Ptr selected_elem_info = std::make_shared(); @@ -61,13 +65,15 @@ void handle_query_field(const VectorQuery *query, SelectInfo *selected_info) { } } -bool SQLInfoHelper::MessageToSQLInfo(const VectorQuery *query, +bool SQLInfoHelper::MessageToSQLInfo(const SearchQuery *query, Node::Ptr filter_node, std::shared_ptr group_by, sqlengine::SQLInfo::Ptr *sql_info, std::string *err_msg) { Node::Ptr index_params_node_ptr = nullptr; - if (!query->query_vector_.empty() || !query->query_sparse_indices_.empty()) { + 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); if (index_params_node_ptr == nullptr) { return false; diff --git a/src/db/sqlengine/parser/sql_info_helper.h b/src/db/sqlengine/parser/sql_info_helper.h index 760dbc4..4c73c45 100644 --- a/src/db/sqlengine/parser/sql_info_helper.h +++ b/src/db/sqlengine/parser/sql_info_helper.h @@ -24,7 +24,7 @@ namespace zvec::sqlengine { class SQLInfoHelper { public: //! Perform QueryRequest to sql info conversion: - static bool MessageToSQLInfo(const VectorQuery *query, Node::Ptr filter_node, + static bool MessageToSQLInfo(const SearchQuery *query, Node::Ptr filter_node, std::shared_ptr group_by, sqlengine::SQLInfo::Ptr *sql_info, std::string *err_msg); diff --git a/src/db/sqlengine/sqlengine.h b/src/db/sqlengine/sqlengine.h index 47143b6..011cff3 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 VectorQuery &query, + CollectionSchema::Ptr collection, const 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 1f5bd51..356736f 100644 --- a/src/db/sqlengine/sqlengine_impl.cc +++ b/src/db/sqlengine/sqlengine_impl.cc @@ -50,7 +50,7 @@ SQLEngineImpl::SQLEngineImpl(zvec::Profiler::Ptr profiler) : profiler_(std::move(profiler)) {} Result SQLEngineImpl::execute( - CollectionSchema::Ptr collection, const VectorQuery &query, + CollectionSchema::Ptr collection, const SearchQuery &query, const std::vector &segments) { if (segments.empty()) { return DocPtrList{}; @@ -76,18 +76,14 @@ Result SQLEngineImpl::execute( return fill_result(select_item_meta_ptrs, reader.value().get()); } -VectorQuery from_group_by(const GroupByVectorQuery &gq) { - VectorQuery vq; - vq.field_name_ = gq.field_name_; - vq.query_vector_ = gq.query_vector_; - vq.query_sparse_indices_ = gq.query_sparse_indices_; - vq.query_sparse_values_ = gq.query_sparse_values_; - vq.filter_ = gq.filter_; - vq.include_vector_ = gq.include_vector_; - vq.query_params_ = gq.query_params_; - vq.output_fields_ = gq.output_fields_; - vq.topk_ = 0; - return vq; +SearchQuery from_group_by(const GroupByVectorQuery &gq) { + SearchQuery sq; + sq.target_ = gq.target_; + sq.filter_ = gq.filter_; + sq.include_vector_ = gq.include_vector_; + sq.output_fields_ = gq.output_fields_; + sq.topk_ = 0; + return sq; } Result SQLEngineImpl::execute_group_by( @@ -97,7 +93,7 @@ Result SQLEngineImpl::execute_group_by( return GroupResults{}; } - VectorQuery query = from_group_by(group_by_query); + SearchQuery query = from_group_by(group_by_query); auto query_info = parse_request( collection, query, std::make_shared(group_by_query.group_by_field_name_, @@ -135,7 +131,7 @@ Result SQLEngineImpl::parse_sql_info( } Result SQLEngineImpl::parse_request( - CollectionSchema::Ptr collection, const VectorQuery &request, + CollectionSchema::Ptr collection, const SearchQuery &request, std::shared_ptr group_by) { profiler_->open_stage("message_to_sqlinfo"); sqlengine::SQLInfo::Ptr sql_info; diff --git a/src/db/sqlengine/sqlengine_impl.h b/src/db/sqlengine/sqlengine_impl.h index 5e25834..658d121 100644 --- a/src/db/sqlengine/sqlengine_impl.h +++ b/src/db/sqlengine/sqlengine_impl.h @@ -34,7 +34,7 @@ class SQLEngineImpl : public SQLEngine { //! Parse pb request Result parse_request(CollectionSchema::Ptr collection, - const VectorQuery &request, + const SearchQuery &request, std::shared_ptr group_by); //! Perform search with given query_info, segments and index filter @@ -44,7 +44,7 @@ class SQLEngineImpl : public SQLEngine { std::vector *query_infos); Result execute( - CollectionSchema::Ptr collection, const VectorQuery &query, + CollectionSchema::Ptr collection, const SearchQuery &query, const std::vector &segments) override; Result execute_group_by( diff --git a/src/include/zvec/c_api.h b/src/include/zvec/c_api.h index 187f870..28e8589 100644 --- a/src/include/zvec/c_api.h +++ b/src/include/zvec/c_api.h @@ -1018,9 +1018,10 @@ typedef struct zvec_flat_query_params_t zvec_flat_query_params_t; /** * @brief Vector query structure (opaque pointer) - * Aligned with zvec::VectorQuery + * Backed by zvec::SearchQuery internally; the C symbol name is kept for + * backward compatibility. * Use zvec_vector_query_create() to create and zvec_vector_query_destroy() to - * destroy + * destroy. */ typedef struct zvec_vector_query_t zvec_vector_query_t; diff --git a/src/include/zvec/db/collection.h b/src/include/zvec/db/collection.h index 3df2840..83a289c 100644 --- a/src/include/zvec/db/collection.h +++ b/src/include/zvec/db/collection.h @@ -97,7 +97,7 @@ class Collection { virtual Status DeleteByFilter(const std::string &filter) = 0; - virtual Result Query(const VectorQuery &query) const = 0; + virtual Result Query(const SearchQuery &query) const = 0; virtual Result Query(const MultiQuery &query) const = 0; diff --git a/src/include/zvec/db/query.h b/src/include/zvec/db/query.h index 98e2c73..cf1eeab 100644 --- a/src/include/zvec/db/query.h +++ b/src/include/zvec/db/query.h @@ -13,6 +13,7 @@ // limitations under the License. #pragma once +#include #include #include #include @@ -24,42 +25,6 @@ namespace zvec { -struct VectorQuery { - int topk_; - std::string field_name_; - std::string query_vector_; // fp16, void * - std::string query_sparse_indices_; - std::string query_sparse_values_; - std::string filter_; - bool include_vector_{false}; - bool include_doc_id_{false}; - // select * by default, select no field if output_fields_ is empty, select - // specific fields if output_fields_ is not empty - std::optional> output_fields_; - QueryParams::Ptr query_params_; - - Status validate_and_sanitize(const FieldSchema *schema); -}; - -struct GroupByVectorQuery { - std::string field_name_; - std::string query_vector_; - std::string query_sparse_indices_; - std::string query_sparse_values_; - std::string filter_; - bool include_vector_; - // select * by default, select no field if output_fields_ is empty, select - // specific fields if output_fields_ is not empty - std::optional> output_fields_; - std::string group_by_field_name_; - uint32_t group_count_ = 2; - uint32_t group_topk_ = 3; - QueryParams::Ptr query_params_; -}; - -//! Multi query structure for combining multiple sub-queries -//! (vector, full-text, etc.) with optional re-ranking of results. - struct VectorClause { std::string query_vector_; std::string sparse_indices_; @@ -75,8 +40,79 @@ struct QueryTarget { std::string field_name_; std::variant clause_; QueryParams::Ptr query_params_; + + // Mutators ensure clause_ holds a VectorClause. + void set_vector(std::string vector); + void set_sparse_vector(std::string indices, std::string values); + + // nullptr when clause_ holds a non-VectorClause alternative. + VectorClause *get_vector_clause() { + return std::get_if(&clause_); + } + const VectorClause *get_vector_clause() const { + return std::get_if(&clause_); + } + + private: + // Resets clause_ to an empty VectorClause unless it already holds one. + VectorClause &ensure_vector_clause() { + if (!std::holds_alternative(clause_)) { + clause_ = VectorClause{}; + } + return std::get(clause_); + } }; +inline void QueryTarget::set_vector(std::string vector) { + ensure_vector_clause().query_vector_ = std::move(vector); +} + +inline void QueryTarget::set_sparse_vector(std::string indices, + std::string values) { + auto &vc = ensure_vector_clause(); + vc.sparse_indices_ = std::move(indices); + vc.sparse_values_ = std::move(values); +} + +struct SearchQuery { + QueryTarget target_; + int topk_{0}; + std::string filter_; + bool include_vector_{false}; + bool include_doc_id_{false}; + // Field selection: + // nullopt -> select all fields (select *) + // empty -> select no field + // non-empty -> select only the listed fields + std::optional> output_fields_; + + // FtsClause currently bypasses validation (FTS not yet implemented). + Status validate_and_sanitize(const FieldSchema *schema); +}; + +struct GroupByVectorQuery { + QueryTarget target_; + std::string filter_; + bool include_vector_{false}; + // Field selection: + // nullopt -> select all fields (select *) + // empty -> select no field + // non-empty -> select only the listed fields + std::optional> output_fields_; + std::string group_by_field_name_; + uint32_t group_count_{2}; + uint32_t group_topk_{3}; +}; + +struct GroupResult { + std::string group_by_value_; + std::vector docs_; +}; + +using GroupResults = std::vector; + +//! Multi query structure for combining multiple sub-queries +//! (vector, full-text, etc.) with optional re-ranking of results. struct SubQuery { QueryTarget target_; int num_candidates_{10}; @@ -88,15 +124,13 @@ struct MultiQuery { std::string filter; bool include_vector{false}; bool include_doc_id_{false}; + // Field selection: + // nullopt -> select all fields (select *) + // empty -> select no field + // non-empty -> select only the listed fields std::optional> output_fields; std::shared_ptr reranker{nullptr}; }; -struct GroupResult { - std::string group_by_value_; - std::vector docs_; -}; - -using GroupResults = std::vector; } // namespace zvec diff --git a/tests/db/collection_test.cc b/tests/db/collection_test.cc index 9317401..2fcf3de 100644 --- a/tests/db/collection_test.cc +++ b/tests/db/collection_test.cc @@ -94,7 +94,7 @@ TEST_F(CollectionTest, Feature_CreateAndOpen_General) { ASSERT_FALSE(col->Delete({}).has_value()); ASSERT_FALSE(col->DeleteByFilter("").ok()); ASSERT_FALSE(col->Fetch({}).has_value()); - ASSERT_FALSE(col->Query(VectorQuery{}).has_value()); + ASSERT_FALSE(col->Query(SearchQuery{}).has_value()); ASSERT_FALSE(col->Query(MultiQuery{}).has_value()); ASSERT_FALSE(col->GroupByQuery({}).has_value()); ASSERT_FALSE(col->CreateIndex("", nullptr).ok()); @@ -1905,9 +1905,9 @@ TEST_F(CollectionTest, Feature_CreateIndex_Vector) { << ", code: " << GetDefaultMessage(s.code()) << std::endl; ASSERT_TRUE(s.ok()); - VectorQuery query; + SearchQuery query; query.topk_ = doc_count; - query.field_name_ = field_name; + query.target_.field_name_ = field_name; query.include_vector_ = true; auto field_scheama = schema->get_vector_field(field_name); ASSERT_NE(field_scheama, nullptr); @@ -1927,37 +1927,36 @@ TEST_F(CollectionTest, Feature_CreateIndex_Vector) { vector_fp16 = std::vector(field_scheama->dimension(), ailego::Float16(1.0f)); vector_fp16[0] = 0; - query.query_vector_.assign( - (char *)vector_fp16.data(), - vector_fp16.size() * sizeof(ailego::Float16)); + query.target_.set_vector( + std::string((char *)vector_fp16.data(), + vector_fp16.size() * sizeof(ailego::Float16))); } else if (field_scheama->data_type() == DataType::VECTOR_FP32) { vector = std::vector(field_scheama->dimension(), 1); vector[0] = 0; - query.query_vector_.assign((char *)vector.data(), - vector.size() * sizeof(float)); + query.target_.set_vector( + std::string((char *)vector.data(), vector.size() * sizeof(float))); } else { vector_int8 = std::vector(field_scheama->dimension(), 1); vector_int8[0] = 0; - query.query_vector_.assign((char *)vector_int8.data(), - vector_int8.size() * sizeof(int8_t)); + query.target_.set_vector(std::string( + (char *)vector_int8.data(), vector_int8.size() * sizeof(int8_t))); } } else { if (field_scheama->data_type() == DataType::SPARSE_VECTOR_FP32) { sparse_vector = {{1}, {1}}; - query.query_sparse_indices_.assign( - (char *)sparse_vector.first.data(), - sparse_vector.first.size() * sizeof(uint32_t)); - query.query_sparse_values_.assign( - (char *)sparse_vector.second.data(), - sparse_vector.second.size() * sizeof(float)); + query.target_.set_sparse_vector( + std::string((char *)sparse_vector.first.data(), + sparse_vector.first.size() * sizeof(uint32_t)), + std::string((char *)sparse_vector.second.data(), + sparse_vector.second.size() * sizeof(float))); } else { sparse_vector_fp16 = {{1}, {ailego::Float16(1.0f)}}; - query.query_sparse_indices_.assign( - (char *)sparse_vector_fp16.first.data(), - sparse_vector_fp16.first.size() * sizeof(uint32_t)); - query.query_sparse_values_.assign( - (char *)sparse_vector_fp16.second.data(), - sparse_vector_fp16.second.size() * sizeof(ailego::Float16)); + query.target_.set_sparse_vector( + std::string((char *)sparse_vector_fp16.first.data(), + sparse_vector_fp16.first.size() * sizeof(uint32_t)), + std::string( + (char *)sparse_vector_fp16.second.data(), + sparse_vector_fp16.second.size() * sizeof(ailego::Float16))); } } auto query_result = collection->Query(query); @@ -2816,15 +2815,16 @@ TEST_F(CollectionTest, Feature_Optimize_MetricType) { auto query_doc = TestHelper::CreateDoc(i, *schema); // std::cout << query_doc.to_detail_string() << std::endl; - VectorQuery query; + SearchQuery query; query.topk_ = 10; query.include_vector_ = true; - query.field_name_ = "dense_fp32"; + query.target_.field_name_ = "dense_fp32"; auto vector = query_doc.get>("dense_fp32"); ASSERT_TRUE(vector.has_value()); - query.query_vector_.assign((char *)vector.value().data(), - vector.value().size() * sizeof(float)); + query.target_.set_vector( + std::string((char *)vector.value().data(), + vector.value().size() * sizeof(float))); auto result = collection->Query(query); @@ -3329,9 +3329,9 @@ TEST_F(CollectionTest, Feature_Query_Validate) { auto query_doc = TestHelper::CreateDoc(1, *schema); { - VectorQuery query; + SearchQuery query; query.topk_ = 1024; - query.field_name_ = field_name; + query.target_.field_name_ = field_name; auto field_scheama = schema->get_vector_field(field_name); ASSERT_NE(field_scheama, nullptr); @@ -3340,18 +3340,18 @@ TEST_F(CollectionTest, Feature_Query_Validate) { if (field_scheama->is_dense_vector()) { auto vector = query_doc.get>(field_name); ASSERT_TRUE(vector.has_value()); - query.query_vector_.assign((char *)vector.value().data(), - vector.value().size() * sizeof(float)); + query.target_.set_vector( + std::string((char *)vector.value().data(), + vector.value().size() * sizeof(float))); } else { auto sparse_vector = query_doc.get, std::vector>>( field_name); - query.query_sparse_indices_.assign( - (char *)sparse_vector.value().first.data(), - sparse_vector.value().first.size() * sizeof(uint32_t)); - query.query_sparse_values_.assign( - (char *)sparse_vector.value().second.data(), - sparse_vector.value().second.size() * sizeof(float)); + query.target_.set_sparse_vector( + std::string((char *)sparse_vector.value().first.data(), + sparse_vector.value().first.size() * sizeof(uint32_t)), + std::string((char *)sparse_vector.value().second.data(), + sparse_vector.value().second.size() * sizeof(float))); } query.include_vector_ = true; @@ -3361,9 +3361,9 @@ TEST_F(CollectionTest, Feature_Query_Validate) { } { - VectorQuery query; + SearchQuery query; query.topk_ = 100001; - query.field_name_ = field_name; + query.target_.field_name_ = field_name; auto field_scheama = schema->get_vector_field(field_name); ASSERT_NE(field_scheama, nullptr); @@ -3372,18 +3372,18 @@ TEST_F(CollectionTest, Feature_Query_Validate) { if (field_scheama->is_dense_vector()) { auto vector = query_doc.get>(field_name); ASSERT_TRUE(vector.has_value()); - query.query_vector_.assign((char *)vector.value().data(), - vector.value().size() * sizeof(float)); + query.target_.set_vector( + std::string((char *)vector.value().data(), + vector.value().size() * sizeof(float))); } else { auto sparse_vector = query_doc.get, std::vector>>( field_name); - query.query_sparse_indices_.assign( - (char *)sparse_vector.value().first.data(), - sparse_vector.value().first.size() * sizeof(uint32_t)); - query.query_sparse_values_.assign( - (char *)sparse_vector.value().second.data(), - sparse_vector.value().second.size() * sizeof(float)); + query.target_.set_sparse_vector( + std::string((char *)sparse_vector.value().first.data(), + sparse_vector.value().first.size() * sizeof(uint32_t)), + std::string((char *)sparse_vector.value().second.data(), + sparse_vector.value().second.size() * sizeof(float))); } query.include_vector_ = true; @@ -3393,9 +3393,9 @@ TEST_F(CollectionTest, Feature_Query_Validate) { } { - VectorQuery query; + SearchQuery query; query.topk_ = 1024; - query.field_name_ = field_name; + query.target_.field_name_ = field_name; query.output_fields_ = std::make_optional>( std::vector(1025)); @@ -3406,18 +3406,18 @@ TEST_F(CollectionTest, Feature_Query_Validate) { if (field_scheama->is_dense_vector()) { auto vector = query_doc.get>(field_name); ASSERT_TRUE(vector.has_value()); - query.query_vector_.assign((char *)vector.value().data(), - vector.value().size() * sizeof(float)); + query.target_.set_vector( + std::string((char *)vector.value().data(), + vector.value().size() * sizeof(float))); } else { auto sparse_vector = query_doc.get, std::vector>>( field_name); - query.query_sparse_indices_.assign( - (char *)sparse_vector.value().first.data(), - sparse_vector.value().first.size() * sizeof(uint32_t)); - query.query_sparse_values_.assign( - (char *)sparse_vector.value().second.data(), - sparse_vector.value().second.size() * sizeof(float)); + query.target_.set_sparse_vector( + std::string((char *)sparse_vector.value().first.data(), + sparse_vector.value().first.size() * sizeof(uint32_t)), + std::string((char *)sparse_vector.value().second.data(), + sparse_vector.value().second.size() * sizeof(float))); } query.include_vector_ = true; @@ -3448,9 +3448,9 @@ TEST_F(CollectionTest, Feature_Query_General) { auto query_doc = TestHelper::CreateDoc(i, *schema); // std::cout << query_doc.to_detail_string() << std::endl; - VectorQuery query; + SearchQuery query; query.topk_ = 10; - query.field_name_ = field_name; + query.target_.field_name_ = field_name; auto field_scheama = schema->get_vector_field(field_name); ASSERT_NE(field_scheama, nullptr); @@ -3459,18 +3459,18 @@ TEST_F(CollectionTest, Feature_Query_General) { if (field_scheama->is_dense_vector()) { auto vector = query_doc.get>(field_name); ASSERT_TRUE(vector.has_value()); - query.query_vector_.assign((char *)vector.value().data(), - vector.value().size() * sizeof(float)); + query.target_.set_vector( + std::string((char *)vector.value().data(), + vector.value().size() * sizeof(float))); } else { auto sparse_vector = query_doc.get, std::vector>>( field_name); - query.query_sparse_indices_.assign( - (char *)sparse_vector.value().first.data(), - sparse_vector.value().first.size() * sizeof(uint32_t)); - query.query_sparse_values_.assign( - (char *)sparse_vector.value().second.data(), - sparse_vector.value().second.size() * sizeof(float)); + query.target_.set_sparse_vector( + std::string((char *)sparse_vector.value().first.data(), + sparse_vector.value().first.size() * sizeof(uint32_t)), + std::string((char *)sparse_vector.value().second.data(), + sparse_vector.value().second.size() * sizeof(float))); } query.include_vector_ = true; @@ -3519,7 +3519,7 @@ TEST_F(CollectionTest, Feature_Query_Empty) { auto query_doc = TestHelper::CreateDoc(i, *schema); // std::cout << query_doc.to_detail_string() << std::endl; - VectorQuery query; + SearchQuery query; query.topk_ = topk; query.include_vector_ = true; @@ -3562,7 +3562,7 @@ TEST_F(CollectionTest, Feature_Query_WithoutVector_CreateScalarIndex) { std::cout << stats.to_string_formatted() << std::endl; // validate query result - VectorQuery query; + SearchQuery query; query.topk_ = topk; query.include_vector_ = true; query.filter_ = filter; @@ -3648,7 +3648,7 @@ TEST_F(CollectionTest, Feature_Query_WithoutVector_WithScalarIndex) { std::cout << stats.to_string_formatted() << std::endl; // validate query result - VectorQuery query; + SearchQuery query; query.topk_ = topk; query.include_vector_ = true; query.filter_ = filter; @@ -4159,7 +4159,7 @@ TEST_F(CollectionTest, Feature_AddColumn_General) { // validate query result for (int i = 1; i < 2; i++) { - VectorQuery query; + SearchQuery query; query.topk_ = 10; query.include_vector_ = true; @@ -4345,7 +4345,7 @@ TEST_F(CollectionTest, Feature_AlterColumn_General) { // validate query result for (int i = 1; i < 2; i++) { - VectorQuery query; + SearchQuery query; query.topk_ = 10; query.include_vector_ = true; @@ -4439,7 +4439,7 @@ TEST_F(CollectionTest, Feature_AlterColumn_CornerCase) { // validate query result for (int i = 1; i < 2; i++) { - VectorQuery query; + SearchQuery query; query.topk_ = 10; query.include_vector_ = true; @@ -5407,13 +5407,13 @@ TEST_F(CollectionTest, Feature_Query_NullableFilter_WithoutIndex) { ASSERT_EQ(stats.doc_count, total); auto query_doc = TestHelper::CreateDoc(1, *schema); - VectorQuery query; + SearchQuery query; query.topk_ = total; - query.field_name_ = "dense_fp32"; + query.target_.field_name_ = "dense_fp32"; auto vec = query_doc.get>("dense_fp32"); ASSERT_TRUE(vec.has_value()); - query.query_vector_.assign((char *)vec.value().data(), - vec.value().size() * sizeof(float)); + query.target_.set_vector(std::string((char *)vec.value().data(), + vec.value().size() * sizeof(float))); query.filter_ = "int32 > 0"; query.output_fields_ = std::vector{"int32"}; diff --git a/tests/db/crash_recovery/optimize_recovery_test.cc b/tests/db/crash_recovery/optimize_recovery_test.cc index a2a7f66..0e88316 100644 --- a/tests/db/crash_recovery/optimize_recovery_test.cc +++ b/tests/db/crash_recovery/optimize_recovery_test.cc @@ -192,12 +192,12 @@ TEST_F(OptimizeRecoveryTest, CrashDuringOptimize) { } } - VectorQuery query; + SearchQuery query; query.topk_ = 10; std::vector feature(128, 0.0); - query.query_vector_.assign((const char *)feature.data(), - feature.size() * sizeof(float)); - query.field_name_ = "dense_fp32_field"; + query.target_.set_vector(std::string((const char *)feature.data(), + feature.size() * sizeof(float))); + query.target_.field_name_ = "dense_fp32_field"; auto query_result = collection->Query(query); ASSERT_TRUE(query_result); auto doc_list = query_result.value(); @@ -247,12 +247,12 @@ TEST_F(OptimizeRecoveryTest, CrashDuringOptimize) { } } - VectorQuery query; + SearchQuery query; query.topk_ = 10; std::vector feature(128, 0.0); - query.query_vector_.assign((const char *)feature.data(), - feature.size() * sizeof(float)); - query.field_name_ = "dense_fp32_field"; + query.target_.set_vector(std::string((const char *)feature.data(), + feature.size() * sizeof(float))); + query.target_.field_name_ = "dense_fp32_field"; auto query_result = collection->Query(query); ASSERT_TRUE(query_result); auto doc_list = query_result.value(); diff --git a/tests/db/index/common/doc_test.cc b/tests/db/index/common/doc_test.cc index 5431411..8c879c8 100644 --- a/tests/db/index/common/doc_test.cc +++ b/tests/db/index/common/doc_test.cc @@ -823,8 +823,7 @@ TEST_F(DocDetailedTest, ValidateAndSanitization) { auto schema = test::TestHelper::CreateNormalSchema(false); std::vector invalid_names = { // Too long (>64) - std::string(65, 'a'), - std::string(64, 'a') + "_", + std::string(65, 'a'), std::string(64, 'a') + "_", // Illegal characters "a b", // space @@ -1219,26 +1218,26 @@ TEST_F(DocDetailedTest, EqualityOperatorCoverage) { } -TEST(VectorQuery, ValidateAndSanitize) { +TEST(SearchQuery, ValidateAndSanitize) { // scalar-only query (no query vector): field schema is null { - VectorQuery query; + SearchQuery query; query.topk_ = 10; - query.field_name_ = "field_name"; + query.target_.field_name_ = "field_name"; auto s = query.validate_and_sanitize(nullptr); EXPECT_TRUE(s.ok()); } // vector query requires a non-null field schema { - VectorQuery query; + SearchQuery query; query.topk_ = 10; - query.field_name_ = "field_name"; + query.target_.field_name_ = "field_name"; std::vector query_vector = {1.0f, 2.0f, 3.0f, 4.0f}; std::string query_vector_str = std::string(reinterpret_cast(query_vector.data()), query_vector.size() * sizeof(float)); - query.query_vector_ = query_vector_str; + query.target_.set_vector(query_vector_str); auto s = query.validate_and_sanitize(nullptr); EXPECT_FALSE(s.ok()); EXPECT_EQ(s.code(), StatusCode::INVALID_ARGUMENT); @@ -1246,8 +1245,8 @@ TEST(VectorQuery, ValidateAndSanitize) { // output_fields count exceeds the allowed maximum { - VectorQuery query; - query.field_name_ = "field_name"; + SearchQuery query; + query.target_.field_name_ = "field_name"; query.topk_ = 10; query.output_fields_ = std::vector(1025); FieldSchema schema = FieldSchema("field_name", DataType::INT32); @@ -1258,20 +1257,20 @@ TEST(VectorQuery, ValidateAndSanitize) { // dense vector query dimension must match the field schema { - VectorQuery query; - query.field_name_ = "field_name"; + SearchQuery query; + query.target_.field_name_ = "field_name"; query.topk_ = 100; std::vector query_vector = {1.0f, 2.0f, 3.0f, 4.0f}; std::string query_vector_str = std::string(reinterpret_cast(query_vector.data()), query_vector.size() * sizeof(float)); - query.query_vector_ = query_vector_str; + query.target_.set_vector(query_vector_str); FieldSchema schema = FieldSchema("field_name", DataType::VECTOR_FP32, 4, true); auto s = query.validate_and_sanitize(&schema); EXPECT_TRUE(s.ok()); - query.query_vector_ = query_vector_str.substr(0, 3); + query.target_.set_vector(query_vector_str.substr(0, 3)); s = query.validate_and_sanitize(&schema); EXPECT_FALSE(s.ok()); EXPECT_EQ(s.code(), StatusCode::INVALID_ARGUMENT); @@ -1279,17 +1278,16 @@ TEST(VectorQuery, ValidateAndSanitize) { // sparse query indices count must not exceed the allowed maximum { - VectorQuery query; - query.field_name_ = "field_name"; + SearchQuery query; + query.target_.field_name_ = "field_name"; query.topk_ = 100; std::vector query_indices(16385); std::vector query_values(16385); - query.query_sparse_indices_ = + query.target_.set_sparse_vector( std::string(reinterpret_cast(query_indices.data()), - query_indices.size() * sizeof(uint32_t)); - query.query_sparse_values_ = + query_indices.size() * sizeof(uint32_t)), std::string(reinterpret_cast(query_values.data()), - query_values.size() * sizeof(float)); + query_values.size() * sizeof(float))); FieldSchema schema = FieldSchema("field_name", DataType::SPARSE_VECTOR_FP32); auto s = query.validate_and_sanitize(&schema); @@ -1299,10 +1297,9 @@ TEST(VectorQuery, ValidateAndSanitize) { // one valid index and matching value: accepted uint32_t one_index = 0; float one_value = 0.0f; - query.query_sparse_indices_ = - std::string(reinterpret_cast(&one_index), sizeof(uint32_t)); - query.query_sparse_values_ = - std::string(reinterpret_cast(&one_value), sizeof(float)); + query.target_.set_sparse_vector( + std::string(reinterpret_cast(&one_index), sizeof(uint32_t)), + std::string(reinterpret_cast(&one_value), sizeof(float))); s = query.validate_and_sanitize(&schema); EXPECT_TRUE(s.ok()); } @@ -1331,32 +1328,36 @@ TEST(VectorQuery, ValidateAndSanitize) { // unsorted indices are sorted in place { - VectorQuery query; - query.field_name_ = "field_name"; + SearchQuery query; + query.target_.field_name_ = "field_name"; query.topk_ = 100; - query.query_sparse_indices_ = pack_idx({42u, 7u, 128u, 3u, 99u}); - query.query_sparse_values_ = pack_val({0.1f, 0.2f, 0.3f, 0.4f, 0.5f}); + query.target_.set_sparse_vector(pack_idx({42u, 7u, 128u, 3u, 99u}), + pack_val({0.1f, 0.2f, 0.3f, 0.4f, 0.5f})); auto s = query.validate_and_sanitize(&schema); EXPECT_TRUE(s.ok()) << s.message(); - EXPECT_EQ(decode_idx(query.query_sparse_indices_), - (std::vector{3u, 7u, 42u, 99u, 128u})); - EXPECT_EQ(decode_val(query.query_sparse_values_), - (std::vector{0.4f, 0.2f, 0.1f, 0.5f, 0.3f})); + EXPECT_EQ( + decode_idx( + std::get(query.target_.clause_).sparse_indices_), + (std::vector{3u, 7u, 42u, 99u, 128u})); + EXPECT_EQ( + decode_val( + std::get(query.target_.clause_).sparse_values_), + (std::vector{0.4f, 0.2f, 0.1f, 0.5f, 0.3f})); } // duplicates are rejected { - VectorQuery query; - query.field_name_ = "field_name"; + SearchQuery query; + query.target_.field_name_ = "field_name"; query.topk_ = 100; - query.query_sparse_indices_ = pack_idx({3u, 7u, 42u, 42u, 99u}); - query.query_sparse_values_ = pack_val({0.1f, 0.2f, 0.3f, 0.4f, 0.5f}); + query.target_.set_sparse_vector(pack_idx({3u, 7u, 42u, 42u, 99u}), + pack_val({0.1f, 0.2f, 0.3f, 0.4f, 0.5f})); auto s = query.validate_and_sanitize(&schema); EXPECT_FALSE(s.ok()); EXPECT_EQ(s.code(), StatusCode::INVALID_ARGUMENT); - query.query_sparse_indices_ = pack_idx({42u, 3u, 7u, 42u, 99u}); - query.query_sparse_values_ = pack_val({0.1f, 0.2f, 0.3f, 0.4f, 0.5f}); + query.target_.set_sparse_vector(pack_idx({42u, 3u, 7u, 42u, 99u}), + pack_val({0.1f, 0.2f, 0.3f, 0.4f, 0.5f})); s = query.validate_and_sanitize(&schema); EXPECT_FALSE(s.ok()); EXPECT_EQ(s.code(), StatusCode::INVALID_ARGUMENT); @@ -1364,19 +1365,21 @@ TEST(VectorQuery, ValidateAndSanitize) { // mismatched counts are rejected { - VectorQuery query; - query.field_name_ = "field_name"; + SearchQuery query; + query.target_.field_name_ = "field_name"; query.topk_ = 100; const auto idx_before = pack_idx({3u, 2u, 1u, 4u}); const auto val_before = pack_val({0.1f, 0.2f, 0.3f, 0.4f}); - query.query_sparse_indices_ = idx_before; - query.query_sparse_values_ = val_before; + query.target_.set_sparse_vector(idx_before, val_before); auto s = query.validate_and_sanitize(&schema); EXPECT_TRUE(s.ok()) << s.message(); - EXPECT_EQ(query.query_sparse_indices_, pack_idx({1u, 2u, 3u, 4u})); - EXPECT_EQ(query.query_sparse_values_, pack_val({0.3f, 0.2f, 0.1f, 0.4f})); + EXPECT_EQ(std::get(query.target_.clause_).sparse_indices_, + pack_idx({1u, 2u, 3u, 4u})); + EXPECT_EQ(std::get(query.target_.clause_).sparse_values_, + pack_val({0.3f, 0.2f, 0.1f, 0.4f})); - query.query_sparse_values_ = pack_val({0.1f, 0.2f, 0.3f}); + std::get(query.target_.clause_).sparse_values_ = + pack_val({0.1f, 0.2f, 0.3f}); s = query.validate_and_sanitize(&schema); EXPECT_FALSE(s.ok()); EXPECT_EQ(s.code(), StatusCode::INVALID_ARGUMENT); @@ -1385,27 +1388,27 @@ TEST(VectorQuery, ValidateAndSanitize) { // query_params type must match the field's index type { - VectorQuery query; - query.field_name_ = "embedding"; + SearchQuery query; + query.target_.field_name_ = "embedding"; query.topk_ = 10; std::vector query_vector(128, 1.0f); - query.query_vector_ = + query.target_.set_vector( std::string(reinterpret_cast(query_vector.data()), - query_vector.size() * sizeof(float)); + query_vector.size() * sizeof(float))); FieldSchema schema = FieldSchema("embedding", DataType::VECTOR_FP32, 128, false, std::make_shared(MetricType::L2)); - query.query_params_ = std::make_shared(150); + query.target_.query_params_ = std::make_shared(150); auto s = query.validate_and_sanitize(&schema); EXPECT_TRUE(s.ok()); - query.query_params_ = std::make_shared(50); + query.target_.query_params_ = std::make_shared(50); s = query.validate_and_sanitize(&schema); EXPECT_FALSE(s.ok()); EXPECT_EQ(s.code(), StatusCode::INVALID_ARGUMENT); - query.query_params_ = nullptr; + query.target_.query_params_ = nullptr; s = query.validate_and_sanitize(&schema); EXPECT_TRUE(s.ok()); } diff --git a/tests/db/sqlengine/contain_test.cc b/tests/db/sqlengine/contain_test.cc index 635fb66..e6562e7 100644 --- a/tests/db/sqlengine/contain_test.cc +++ b/tests/db/sqlengine/contain_test.cc @@ -124,7 +124,7 @@ class ContainTest : public testing::Test { TEST_F(ContainTest, ContainAllInt32) { - VectorQuery query; + SearchQuery query; query.output_fields_ = std::vector{}; query.topk_ = 200; query.filter_ = "i32_array contain_all ("; @@ -153,7 +153,7 @@ TEST_F(ContainTest, ContainAllInt32) { } TEST_F(ContainTest, ContainAllInt64) { - VectorQuery query; + SearchQuery query; query.output_fields_ = std::vector{}; query.topk_ = 200; query.filter_ = "i64_array contain_all ("; @@ -182,7 +182,7 @@ TEST_F(ContainTest, ContainAllInt64) { } TEST_F(ContainTest, ContainAllUint32) { - VectorQuery query; + SearchQuery query; query.output_fields_ = std::vector{}; query.topk_ = 200; query.filter_ = "u32_array contain_all ("; @@ -211,7 +211,7 @@ TEST_F(ContainTest, ContainAllUint32) { } TEST_F(ContainTest, ContainAllUint64) { - VectorQuery query; + SearchQuery query; query.output_fields_ = std::vector{}; query.topk_ = 200; query.filter_ = "u64_array contain_all ("; @@ -240,7 +240,7 @@ TEST_F(ContainTest, ContainAllUint64) { } TEST_F(ContainTest, ContainAllFp32) { - VectorQuery query; + SearchQuery query; query.output_fields_ = std::vector{}; query.topk_ = 200; query.filter_ = "fp32_array contain_all ("; @@ -269,7 +269,7 @@ TEST_F(ContainTest, ContainAllFp32) { } TEST_F(ContainTest, ContainAllFp64) { - VectorQuery query; + SearchQuery query; query.output_fields_ = std::vector{}; query.topk_ = 200; query.filter_ = "fp64_array contain_all ("; @@ -298,7 +298,7 @@ TEST_F(ContainTest, ContainAllFp64) { } TEST_F(ContainTest, ContainAllString) { - VectorQuery query; + SearchQuery query; query.output_fields_ = std::vector{}; query.topk_ = 200; query.filter_ = "str_array contain_all ("; @@ -327,7 +327,7 @@ TEST_F(ContainTest, ContainAllString) { } TEST_F(ContainTest, ContainAnyInt32) { - VectorQuery query; + SearchQuery query; query.output_fields_ = std::vector{}; query.topk_ = 200; query.filter_ = "i32_array contain_any (98,99,100)"; @@ -349,7 +349,7 @@ TEST_F(ContainTest, ContainAnyInt32) { } TEST_F(ContainTest, ContainAnyInt64) { - VectorQuery query; + SearchQuery query; query.output_fields_ = std::vector{}; query.topk_ = 200; query.filter_ = "i64_array contain_any (98,99,100)"; @@ -371,7 +371,7 @@ TEST_F(ContainTest, ContainAnyInt64) { } TEST_F(ContainTest, ContainAnyUint32) { - VectorQuery query; + SearchQuery query; query.output_fields_ = std::vector{}; query.topk_ = 200; query.filter_ = "u32_array contain_any (98,99,100)"; @@ -393,7 +393,7 @@ TEST_F(ContainTest, ContainAnyUint32) { } TEST_F(ContainTest, ContainAnyUint64) { - VectorQuery query; + SearchQuery query; query.output_fields_ = std::vector{}; query.topk_ = 200; query.filter_ = "u64_array contain_any (98,99,100)"; @@ -415,7 +415,7 @@ TEST_F(ContainTest, ContainAnyUint64) { } TEST_F(ContainTest, ContainAnyFp32) { - VectorQuery query; + SearchQuery query; query.output_fields_ = std::vector{}; query.topk_ = 200; query.filter_ = "fp32_array contain_any (98,99,100)"; @@ -437,7 +437,7 @@ TEST_F(ContainTest, ContainAnyFp32) { } TEST_F(ContainTest, ContainAnyFp64) { - VectorQuery query; + SearchQuery query; query.output_fields_ = std::vector{}; query.topk_ = 200; query.filter_ = "fp64_array contain_any (98,99,100)"; @@ -459,7 +459,7 @@ TEST_F(ContainTest, ContainAnyFp64) { } TEST_F(ContainTest, ContainAnyString) { - VectorQuery query; + SearchQuery query; query.output_fields_ = std::vector{}; query.topk_ = 200; query.filter_ = "str_array contain_any ('name98','name99','name100')"; diff --git a/tests/db/sqlengine/forward_recall_test.cc b/tests/db/sqlengine/forward_recall_test.cc index 6eea146..6f05147 100644 --- a/tests/db/sqlengine/forward_recall_test.cc +++ b/tests/db/sqlengine/forward_recall_test.cc @@ -24,7 +24,7 @@ namespace zvec::sqlengine { class ForwardRecallTest : public RecallTest {}; TEST_F(ForwardRecallTest, Basic) { - VectorQuery query; + SearchQuery query; query.output_fields_ = {"id", "name", "age"}; query.topk_ = 200; @@ -48,7 +48,7 @@ TEST_F(ForwardRecallTest, Basic) { } TEST_F(ForwardRecallTest, BasicWithDocId) { - VectorQuery query; + SearchQuery query; query.output_fields_ = {"id", "name", "age"}; query.topk_ = 200; query.include_doc_id_ = true; @@ -74,7 +74,7 @@ TEST_F(ForwardRecallTest, BasicWithDocId) { } TEST_F(ForwardRecallTest, OutputNoFields) { - VectorQuery query; + SearchQuery query; query.output_fields_ = std::vector{}; query.topk_ = 200; @@ -94,7 +94,7 @@ TEST_F(ForwardRecallTest, OutputNoFields) { } TEST_F(ForwardRecallTest, DenseVector) { - VectorQuery query; + SearchQuery query; query.output_fields_ = {"id", "dense"}; query.topk_ = 200; query.include_vector_ = true; @@ -121,7 +121,7 @@ TEST_F(ForwardRecallTest, DenseVector) { } TEST_F(ForwardRecallTest, SparseVector) { - VectorQuery query; + SearchQuery query; query.output_fields_ = {"id", "sparse"}; query.topk_ = 200; query.include_vector_ = true; @@ -160,7 +160,7 @@ TEST_F(ForwardRecallTest, SparseVector) { } TEST_F(ForwardRecallTest, MultiSegment) { - VectorQuery query; + SearchQuery query; query.output_fields_ = std::vector(); query.topk_ = 200; query.include_vector_ = true; @@ -206,7 +206,7 @@ TEST_F(ForwardRecallTest, MultiSegment) { } TEST_F(ForwardRecallTest, Eq) { - VectorQuery query; + SearchQuery query; query.output_fields_ = {"id", "name", "age"}; query.topk_ = 200; query.filter_ = "age = 1"; @@ -229,7 +229,7 @@ TEST_F(ForwardRecallTest, Eq) { } TEST_F(ForwardRecallTest, Gt) { - VectorQuery query; + SearchQuery query; query.output_fields_ = {"id", "name", "age"}; query.topk_ = 200; query.filter_ = "id > 1000"; @@ -252,7 +252,7 @@ TEST_F(ForwardRecallTest, Gt) { } TEST_F(ForwardRecallTest, Ge) { - VectorQuery query; + SearchQuery query; query.output_fields_ = {"id", "name", "age"}; query.topk_ = 200; query.filter_ = "id >= 1000"; @@ -275,7 +275,7 @@ TEST_F(ForwardRecallTest, Ge) { } TEST_F(ForwardRecallTest, Lt) { - VectorQuery query; + SearchQuery query; query.output_fields_ = {"id", "name", "age"}; query.topk_ = 200; query.filter_ = "id < 100"; @@ -298,7 +298,7 @@ TEST_F(ForwardRecallTest, Lt) { } TEST_F(ForwardRecallTest, Le) { - VectorQuery query; + SearchQuery query; query.output_fields_ = {"id", "name", "age"}; query.topk_ = 200; query.filter_ = "id <= 100"; @@ -321,7 +321,7 @@ TEST_F(ForwardRecallTest, Le) { } TEST_F(ForwardRecallTest, And) { - VectorQuery query; + SearchQuery query; query.output_fields_ = {"id", "name", "age"}; query.topk_ = 200; query.filter_ = "id <= 100 and id > 50"; @@ -344,7 +344,7 @@ TEST_F(ForwardRecallTest, And) { } TEST_F(ForwardRecallTest, Or) { - VectorQuery query; + SearchQuery query; query.output_fields_ = {"id", "name", "age"}; query.topk_ = 200; query.filter_ = "id < 100 or id > 200"; @@ -368,7 +368,7 @@ TEST_F(ForwardRecallTest, Or) { } TEST_F(ForwardRecallTest, StrEq) { - VectorQuery query; + SearchQuery query; query.output_fields_ = {"id", "name", "age"}; query.topk_ = 200; query.filter_ = "name = 'user_1'"; @@ -391,7 +391,7 @@ TEST_F(ForwardRecallTest, StrEq) { } TEST_F(ForwardRecallTest, StrGe) { - VectorQuery query; + SearchQuery query; query.output_fields_ = {"id", "name", "age"}; query.topk_ = 200; query.filter_ = "name >= 'user_1'"; @@ -417,7 +417,7 @@ TEST_F(ForwardRecallTest, StrGe) { } TEST_F(ForwardRecallTest, StrIn) { - VectorQuery query; + SearchQuery query; query.output_fields_ = {"id", "name", "age"}; query.topk_ = 200; query.filter_ = "name IN ('user_1', 'user_2')"; @@ -446,7 +446,7 @@ TEST_F(ForwardRecallTest, StrIn) { } TEST_F(ForwardRecallTest, StrNotIn) { - VectorQuery query; + SearchQuery query; query.output_fields_ = {"id", "name", "age"}; query.topk_ = 200; query.filter_ = "name NOT IN ('user_1', 'user_2')"; @@ -475,7 +475,7 @@ TEST_F(ForwardRecallTest, StrNotIn) { } TEST_F(ForwardRecallTest, StrLike) { - VectorQuery query; + SearchQuery query; query.output_fields_ = {"id", "name", "age"}; query.topk_ = 200; query.filter_ = "name like 'user_9%'"; @@ -506,7 +506,7 @@ TEST_F(ForwardRecallTest, StrLike) { } TEST_F(ForwardRecallTest, IsNull) { - VectorQuery query; + SearchQuery query; query.output_fields_ = {"id", "name", "age"}; query.topk_ = 200; query.filter_ = "optional_age is null"; @@ -529,7 +529,7 @@ TEST_F(ForwardRecallTest, IsNull) { } TEST_F(ForwardRecallTest, IsNotNull) { - VectorQuery query; + SearchQuery query; query.output_fields_ = {"id", "name", "age"}; query.topk_ = 200; query.filter_ = "optional_age is not null"; @@ -555,7 +555,7 @@ TEST_F(ForwardRecallTest, IsNotNull) { } TEST_F(ForwardRecallTest, IsNullNoResult) { - VectorQuery query; + SearchQuery query; query.output_fields_ = {"id", "name", "age"}; query.topk_ = 200; query.filter_ = "age is null"; @@ -568,7 +568,7 @@ TEST_F(ForwardRecallTest, IsNullNoResult) { } TEST_F(ForwardRecallTest, ContainAll) { - VectorQuery query; + SearchQuery query; query.output_fields_ = {"id", "name", "age"}; query.topk_ = 200; query.filter_ = "category_set contain_all ("; @@ -603,7 +603,7 @@ TEST_F(ForwardRecallTest, ContainAll) { } TEST_F(ForwardRecallTest, NotContainAll) { - VectorQuery query; + SearchQuery query; query.output_fields_ = {"id", "name", "age"}; query.topk_ = 200; query.filter_ = "category_set not contain_all ("; @@ -639,7 +639,7 @@ TEST_F(ForwardRecallTest, NotContainAll) { } TEST_F(ForwardRecallTest, ContainAny) { - VectorQuery query; + SearchQuery query; query.output_fields_ = {"id", "name", "age"}; query.topk_ = 200; query.filter_ = "category_set contain_any (98,99,100)"; @@ -667,7 +667,7 @@ TEST_F(ForwardRecallTest, ContainAny) { } TEST_F(ForwardRecallTest, NotContainAny) { - VectorQuery query; + SearchQuery query; query.output_fields_ = {"id", "name", "age"}; query.topk_ = 200; query.filter_ = "category_set not contain_any (98,99,100)"; @@ -696,7 +696,7 @@ TEST_F(ForwardRecallTest, NotContainAny) { } TEST_F(ForwardRecallTest, BoolContainAll) { - VectorQuery query; + SearchQuery query; query.output_fields_ = {"id", "name", "age"}; query.topk_ = 200; query.filter_ = "bool_array contain_all (true, false)"; @@ -721,7 +721,7 @@ TEST_F(ForwardRecallTest, BoolContainAll) { } TEST_F(ForwardRecallTest, BoolContainAny) { - VectorQuery query; + SearchQuery query; query.output_fields_ = {"id", "name", "age"}; query.topk_ = 200; query.filter_ = "bool_array contain_any (true)"; @@ -749,7 +749,7 @@ TEST_F(ForwardRecallTest, BoolContainAny) { } TEST_F(ForwardRecallTest, ContainAllEmptySet) { - VectorQuery query; + SearchQuery query; query.output_fields_ = {"id", "name", "age"}; query.topk_ = 200; query.filter_ = "category_set contain_all ()"; @@ -777,7 +777,7 @@ TEST_F(ForwardRecallTest, ContainAllEmptySet) { } TEST_F(ForwardRecallTest, NotContainAllEmptySet) { - VectorQuery query; + SearchQuery query; query.output_fields_ = {"id", "name", "age"}; query.topk_ = 200; query.filter_ = "category_set not contain_all ()"; @@ -790,7 +790,7 @@ TEST_F(ForwardRecallTest, NotContainAllEmptySet) { } TEST_F(ForwardRecallTest, ContainAnyEmptySet) { - VectorQuery query; + SearchQuery query; query.output_fields_ = {"id", "name", "age"}; query.topk_ = 200; query.filter_ = "category_set contain_any ()"; @@ -803,7 +803,7 @@ TEST_F(ForwardRecallTest, ContainAnyEmptySet) { } TEST_F(ForwardRecallTest, NotContainAnyEmptySet) { - VectorQuery query; + SearchQuery query; query.output_fields_ = {"id", "name", "age"}; query.topk_ = 200; query.filter_ = "category_set not contain_any ()"; @@ -831,7 +831,7 @@ TEST_F(ForwardRecallTest, NotContainAnyEmptySet) { } TEST_F(ForwardRecallTest, BoolEqTrue) { - VectorQuery query; + SearchQuery query; query.output_fields_ = {"id", "name", "age"}; query.topk_ = 200; query.filter_ = "bool = TRuE"; @@ -854,7 +854,7 @@ TEST_F(ForwardRecallTest, BoolEqTrue) { } TEST_F(ForwardRecallTest, BoolEqFalse) { - VectorQuery query; + SearchQuery query; query.output_fields_ = {"id", "name", "age"}; query.topk_ = 200; query.filter_ = "bool = false"; @@ -882,7 +882,7 @@ TEST_F(ForwardRecallTest, BoolEqFalse) { } TEST_F(ForwardRecallTest, ArrayLengthEq) { - VectorQuery query; + SearchQuery query; query.output_fields_ = {"id", "name", "age"}; query.topk_ = 200; query.filter_ = "array_length(category_set) = 32"; @@ -907,7 +907,7 @@ TEST_F(ForwardRecallTest, ArrayLengthEq) { } TEST_F(ForwardRecallTest, ArrayLengthGe) { - VectorQuery query; + SearchQuery query; query.output_fields_ = {"id", "name", "age"}; query.topk_ = 200; query.filter_ = "array_length(category_set) >= 32"; diff --git a/tests/db/sqlengine/invert_recall_test.cc b/tests/db/sqlengine/invert_recall_test.cc index 3095b2e..80235e8 100644 --- a/tests/db/sqlengine/invert_recall_test.cc +++ b/tests/db/sqlengine/invert_recall_test.cc @@ -24,7 +24,7 @@ namespace zvec::sqlengine { class InvertRecallTest : public RecallTest {}; TEST_F(InvertRecallTest, Eq) { - VectorQuery query; + SearchQuery query; query.output_fields_ = {"id", "name", "age"}; query.topk_ = 200; query.filter_ = "invert_age = 1"; @@ -47,7 +47,7 @@ TEST_F(InvertRecallTest, Eq) { } TEST_F(InvertRecallTest, Gt) { - VectorQuery query; + SearchQuery query; query.output_fields_ = {"id", "name", "age"}; query.topk_ = 200; query.filter_ = "invert_id > 1000"; @@ -70,7 +70,7 @@ TEST_F(InvertRecallTest, Gt) { } TEST_F(InvertRecallTest, Ge) { - VectorQuery query; + SearchQuery query; query.output_fields_ = {"id", "name", "age"}; query.topk_ = 200; query.filter_ = "invert_id >= 1000"; @@ -93,7 +93,7 @@ TEST_F(InvertRecallTest, Ge) { } TEST_F(InvertRecallTest, Lt) { - VectorQuery query; + SearchQuery query; query.output_fields_ = {"id", "name", "age"}; query.topk_ = 200; query.filter_ = "invert_id < 100"; @@ -116,7 +116,7 @@ TEST_F(InvertRecallTest, Lt) { } TEST_F(InvertRecallTest, Le) { - VectorQuery query; + SearchQuery query; query.output_fields_ = {"id", "name", "age"}; query.topk_ = 200; query.filter_ = "invert_id <= 100"; @@ -139,7 +139,7 @@ TEST_F(InvertRecallTest, Le) { } TEST_F(InvertRecallTest, And) { - VectorQuery query; + SearchQuery query; query.output_fields_ = {"id", "name", "age"}; query.topk_ = 200; query.filter_ = "invert_id <= 100 and invert_id > 50"; @@ -162,7 +162,7 @@ TEST_F(InvertRecallTest, And) { } TEST_F(InvertRecallTest, Or) { - VectorQuery query; + SearchQuery query; query.output_fields_ = {"id", "name", "age"}; query.topk_ = 200; query.filter_ = "invert_id < 100 or invert_id > 200"; @@ -186,7 +186,7 @@ TEST_F(InvertRecallTest, Or) { } TEST_F(InvertRecallTest, StrEq) { - VectorQuery query; + SearchQuery query; query.output_fields_ = {"id", "name", "age"}; query.topk_ = 200; query.filter_ = "invert_name = 'user_1'"; @@ -209,7 +209,7 @@ TEST_F(InvertRecallTest, StrEq) { } TEST_F(InvertRecallTest, StrGe) { - VectorQuery query; + SearchQuery query; query.output_fields_ = {"id", "name", "age"}; query.topk_ = 200; query.filter_ = "invert_name >= 'user_1'"; @@ -235,7 +235,7 @@ TEST_F(InvertRecallTest, StrGe) { } TEST_F(InvertRecallTest, StrIn) { - VectorQuery query; + SearchQuery query; query.output_fields_ = {"id", "name", "age"}; query.topk_ = 200; query.filter_ = "invert_name IN ('user_1', 'user_2')"; @@ -264,7 +264,7 @@ TEST_F(InvertRecallTest, StrIn) { } TEST_F(InvertRecallTest, StrNotIn) { - VectorQuery query; + SearchQuery query; query.output_fields_ = {"id", "name", "age"}; query.topk_ = 200; query.filter_ = "invert_name NOT IN ('user_1', 'user_2')"; @@ -293,7 +293,7 @@ TEST_F(InvertRecallTest, StrNotIn) { } TEST_F(InvertRecallTest, StrLike) { - VectorQuery query; + SearchQuery query; query.output_fields_ = {"id", "name", "age"}; query.topk_ = 200; query.filter_ = "invert_name like 'user\\_9%'"; @@ -324,7 +324,7 @@ TEST_F(InvertRecallTest, StrLike) { } TEST_F(InvertRecallTest, ContainAll) { - VectorQuery query; + SearchQuery query; query.output_fields_ = {"id", "name", "age"}; query.topk_ = 200; query.filter_ = "invert_category_set contain_all ("; @@ -359,7 +359,7 @@ TEST_F(InvertRecallTest, ContainAll) { } TEST_F(InvertRecallTest, NotContainAll) { - VectorQuery query; + SearchQuery query; query.output_fields_ = {"id", "name", "age"}; query.topk_ = 200; query.filter_ = "invert_category_set not contain_all ("; @@ -395,7 +395,7 @@ TEST_F(InvertRecallTest, NotContainAll) { } TEST_F(InvertRecallTest, ContainAny) { - VectorQuery query; + SearchQuery query; query.output_fields_ = {"id", "name", "age"}; query.topk_ = 200; query.filter_ = "invert_category_set contain_any (98,99,100)"; @@ -423,7 +423,7 @@ TEST_F(InvertRecallTest, ContainAny) { } TEST_F(InvertRecallTest, NotContainAny) { - VectorQuery query; + SearchQuery query; query.output_fields_ = {"id", "name", "age"}; query.topk_ = 200; query.filter_ = "invert_category_set not contain_any (98,99,100)"; @@ -452,7 +452,7 @@ TEST_F(InvertRecallTest, NotContainAny) { } TEST_F(InvertRecallTest, BoolContainAll) { - VectorQuery query; + SearchQuery query; query.output_fields_ = {"id", "name", "age"}; query.topk_ = 200; query.filter_ = "invert_bool_array contain_all (true, false)"; @@ -477,7 +477,7 @@ TEST_F(InvertRecallTest, BoolContainAll) { } TEST_F(InvertRecallTest, BoolContainAny) { - VectorQuery query; + SearchQuery query; query.output_fields_ = {"id", "name", "age"}; query.topk_ = 200; query.filter_ = "invert_bool_array contain_any (true)"; @@ -505,7 +505,7 @@ TEST_F(InvertRecallTest, BoolContainAny) { } TEST_F(InvertRecallTest, ContainAllEmptySet) { - VectorQuery query; + SearchQuery query; query.output_fields_ = {"id", "name", "age"}; query.topk_ = 200; query.filter_ = "invert_category_set contain_all ()"; @@ -533,7 +533,7 @@ TEST_F(InvertRecallTest, ContainAllEmptySet) { } TEST_F(InvertRecallTest, NotContainAllEmptySet) { - VectorQuery query; + SearchQuery query; query.output_fields_ = {"id", "name", "age"}; query.topk_ = 200; query.filter_ = "invert_category_set not contain_all ()"; @@ -546,7 +546,7 @@ TEST_F(InvertRecallTest, NotContainAllEmptySet) { } TEST_F(InvertRecallTest, ContainAnyEmptySet) { - VectorQuery query; + SearchQuery query; query.output_fields_ = {"id", "name", "age"}; query.topk_ = 200; query.filter_ = "invert_category_set contain_any ()"; @@ -559,7 +559,7 @@ TEST_F(InvertRecallTest, ContainAnyEmptySet) { } TEST_F(InvertRecallTest, NotContainAnyEmptySet) { - VectorQuery query; + SearchQuery query; query.output_fields_ = {"id", "name", "age"}; query.topk_ = 200; query.filter_ = "invert_category_set not contain_any ()"; @@ -587,7 +587,7 @@ TEST_F(InvertRecallTest, NotContainAnyEmptySet) { } TEST_F(InvertRecallTest, IsNull) { - VectorQuery query; + SearchQuery query; query.output_fields_ = {"id", "name", "age"}; query.topk_ = 200; query.filter_ = "invert_optional_age is null"; @@ -610,7 +610,7 @@ TEST_F(InvertRecallTest, IsNull) { } TEST_F(InvertRecallTest, IsNotNull) { - VectorQuery query; + SearchQuery query; query.output_fields_ = {"id", "name", "age"}; query.topk_ = 200; query.filter_ = "invert_optional_age is not null"; @@ -636,7 +636,7 @@ TEST_F(InvertRecallTest, IsNotNull) { } TEST_F(InvertRecallTest, BoolEqTrue) { - VectorQuery query; + SearchQuery query; query.output_fields_ = {"id", "name", "age"}; query.topk_ = 200; query.filter_ = "invert_bool = TRuE"; @@ -659,7 +659,7 @@ TEST_F(InvertRecallTest, BoolEqTrue) { } TEST_F(InvertRecallTest, BoolEqFalse) { - VectorQuery query; + SearchQuery query; query.output_fields_ = {"id", "name", "age"}; query.topk_ = 200; query.filter_ = "invert_bool = false"; @@ -687,7 +687,7 @@ TEST_F(InvertRecallTest, BoolEqFalse) { } TEST_F(InvertRecallTest, ArrayLengthGe) { - VectorQuery query; + SearchQuery query; query.output_fields_ = {"id", "name", "age"}; query.topk_ = 200; query.filter_ = "array_length(invert_category_set) >= 32"; @@ -715,7 +715,7 @@ TEST_F(InvertRecallTest, ArrayLengthGe) { } TEST_F(InvertRecallTest, ArrayLengthEq) { - VectorQuery query; + SearchQuery query; query.output_fields_ = {"id", "name", "age"}; query.topk_ = 200; query.filter_ = "array_length(invert_category_set) = 32"; @@ -740,7 +740,7 @@ TEST_F(InvertRecallTest, ArrayLengthEq) { } TEST_F(InvertRecallTest, MultiSegment) { - VectorQuery query; + SearchQuery query; query.output_fields_ = std::vector(); query.topk_ = 200; query.include_vector_ = true; diff --git a/tests/db/sqlengine/like_test.cc b/tests/db/sqlengine/like_test.cc index d1f6bbb..e5ce384 100644 --- a/tests/db/sqlengine/like_test.cc +++ b/tests/db/sqlengine/like_test.cc @@ -96,7 +96,7 @@ class LikeTest : public testing::Test { TEST_F(LikeTest, ForwardLikeAll) { - VectorQuery query; + SearchQuery query; query.output_fields_ = {"name"}; query.topk_ = 200; query.filter_ = "name like '%'"; @@ -112,7 +112,7 @@ TEST_F(LikeTest, ForwardLikeAll) { } TEST_F(LikeTest, InvertLikeAll) { - VectorQuery query; + SearchQuery query; query.output_fields_ = {"name"}; query.topk_ = 200; query.filter_ = "invert_name like '%'"; @@ -128,7 +128,7 @@ TEST_F(LikeTest, InvertLikeAll) { } TEST_F(LikeTest, ForwardPrefixLike) { - VectorQuery query; + SearchQuery query; query.output_fields_ = {"name"}; query.topk_ = 200; query.filter_ = "name like 'user-22%'"; @@ -145,7 +145,7 @@ TEST_F(LikeTest, ForwardPrefixLike) { } TEST_F(LikeTest, InvertPrefixLike) { - VectorQuery query; + SearchQuery query; query.output_fields_ = {"name"}; query.topk_ = 200; query.filter_ = "invert_name like 'user-22%'"; @@ -162,7 +162,7 @@ TEST_F(LikeTest, InvertPrefixLike) { } TEST_F(LikeTest, ForwardSuffixLike) { - VectorQuery query; + SearchQuery query; query.output_fields_ = {"name"}; query.topk_ = 200; query.filter_ = "name like '%ser-22'"; @@ -179,7 +179,7 @@ TEST_F(LikeTest, ForwardSuffixLike) { } TEST_F(LikeTest, NotExtendedInvertSuffixLikeRunAsForward) { - VectorQuery query; + SearchQuery query; query.output_fields_ = {"name"}; query.topk_ = 200; query.filter_ = "invert_name like '%ser-22'"; @@ -196,7 +196,7 @@ TEST_F(LikeTest, NotExtendedInvertSuffixLikeRunAsForward) { } TEST_F(LikeTest, ExtendedInvertSuffixLike) { - VectorQuery query; + SearchQuery query; query.output_fields_ = {"name"}; query.topk_ = 200; query.filter_ = "extended_invert_name like '%ser-22'"; @@ -213,7 +213,7 @@ TEST_F(LikeTest, ExtendedInvertSuffixLike) { } TEST_F(LikeTest, ForwardMiddleLike) { - VectorQuery query; + SearchQuery query; query.output_fields_ = {"name"}; query.topk_ = 200; query.filter_ = "name like 'user%2'"; @@ -232,7 +232,7 @@ TEST_F(LikeTest, ForwardMiddleLike) { } TEST_F(LikeTest, ExtendedInvertMiddleLike) { - VectorQuery query; + SearchQuery query; query.output_fields_ = {"name"}; query.topk_ = 200; query.filter_ = "extended_invert_name like 'user%2'"; @@ -251,7 +251,7 @@ TEST_F(LikeTest, ExtendedInvertMiddleLike) { } TEST_F(LikeTest, UnderScore) { - VectorQuery query; + SearchQuery query; query.output_fields_ = {"name"}; query.topk_ = 200; query.filter_ = "name like 'user-_2'"; @@ -270,7 +270,7 @@ TEST_F(LikeTest, UnderScore) { } TEST_F(LikeTest, InvertUnderScoreRunAsForward) { - VectorQuery query; + SearchQuery query; query.output_fields_ = {"name"}; query.topk_ = 200; query.filter_ = "invert_name like 'user-_2'"; @@ -289,7 +289,7 @@ TEST_F(LikeTest, InvertUnderScoreRunAsForward) { } TEST_F(LikeTest, ForwardEscapePercent) { - VectorQuery query; + SearchQuery query; query.output_fields_ = {"name"}; query.topk_ = 200; query.filter_ = R"(name like 'user-\%%')"; @@ -305,7 +305,7 @@ TEST_F(LikeTest, ForwardEscapePercent) { } TEST_F(LikeTest, InvertEscapePercent) { - VectorQuery query; + SearchQuery query; query.output_fields_ = {"name"}; query.topk_ = 200; query.filter_ = R"(invert_name like 'user-\%%')"; @@ -321,7 +321,7 @@ TEST_F(LikeTest, InvertEscapePercent) { } TEST_F(LikeTest, ForwardEscapeUnderscore) { - VectorQuery query; + SearchQuery query; query.output_fields_ = {"name"}; query.topk_ = 200; query.filter_ = R"(name like 'user-\_%')"; @@ -337,7 +337,7 @@ TEST_F(LikeTest, ForwardEscapeUnderscore) { } TEST_F(LikeTest, InvertEscapeUnderscore) { - VectorQuery query; + SearchQuery query; query.output_fields_ = {"name"}; query.topk_ = 200; query.filter_ = R"(invert_name like 'user-\_%')"; @@ -353,7 +353,7 @@ TEST_F(LikeTest, InvertEscapeUnderscore) { } TEST_F(LikeTest, NoPercentRunAsEqual) { - VectorQuery query; + SearchQuery query; query.output_fields_ = {"name"}; query.topk_ = 200; query.filter_ = R"(invert_name like 'user-22')"; diff --git a/tests/db/sqlengine/optimizer_test.cc b/tests/db/sqlengine/optimizer_test.cc index 4ac0ab4..f726060 100644 --- a/tests/db/sqlengine/optimizer_test.cc +++ b/tests/db/sqlengine/optimizer_test.cc @@ -92,10 +92,10 @@ class OptimizerTest : public testing::Test { TEST_F(OptimizerTest, Basic) { - VectorQuery query; + SearchQuery query; query.output_fields_ = {"*"}; query.topk_ = 11; - query.field_name_ = "face_feature"; + query.target_.field_name_ = "face_feature"; query.include_vector_ = false; query.filter_ = "age > 200"; @@ -115,10 +115,10 @@ TEST_F(OptimizerTest, Basic) { // case 1. invert subroot same as invert cond, do nothing TEST_F(OptimizerTest, Case1) { - VectorQuery query; + SearchQuery query; query.output_fields_ = {"*"}; query.topk_ = 11; - query.field_name_ = "face_feature"; + query.target_.field_name_ = "face_feature"; query.include_vector_ = false; query.filter_ = "age > 12"; @@ -138,10 +138,10 @@ TEST_F(OptimizerTest, Case1) { // case 2.1 invert subroot is not found, all conds are forward cond TEST_F(OptimizerTest, Case2_1) { - VectorQuery query; + SearchQuery query; query.output_fields_ = {"*"}; query.topk_ = 11; - query.field_name_ = "face_feature"; + query.target_.field_name_ = "face_feature"; query.include_vector_ = false; query.filter_ = "age > 100 and age > 101 or age > 102"; @@ -162,10 +162,10 @@ TEST_F(OptimizerTest, Case2_1) { // case 2.2 invert subroot is not found, some conds are forward cond // while left invert cond cannot be invert cond any more TEST_F(OptimizerTest, Case2_2) { - VectorQuery query; + SearchQuery query; query.output_fields_ = {"*"}; query.topk_ = 11; - query.field_name_ = "face_feature"; + query.target_.field_name_ = "face_feature"; query.include_vector_ = false; query.filter_ = "age > 100 or age > 90"; @@ -186,10 +186,10 @@ TEST_F(OptimizerTest, Case2_2) { // case 3.1 subroot is found and be part of invert cond TEST_F(OptimizerTest, Case3_1) { - VectorQuery query; + SearchQuery query; query.output_fields_ = {"*"}; query.topk_ = 11; - query.field_name_ = "face_feature"; + query.target_.field_name_ = "face_feature"; query.include_vector_ = false; query.filter_ = "age > 100 and age > 101 and age > 10"; @@ -210,10 +210,10 @@ TEST_F(OptimizerTest, Case3_1) { // case 3.2 subroot is found and be part of invert cond TEST_F(OptimizerTest, Case3_2) { - VectorQuery query; + SearchQuery query; query.output_fields_ = {"*"}; query.topk_ = 11; - query.field_name_ = "face_feature"; + query.target_.field_name_ = "face_feature"; query.include_vector_ = false; query.filter_ = "age > 10 and age > 11 and age > 100"; @@ -233,10 +233,10 @@ TEST_F(OptimizerTest, Case3_2) { // case 3.3 subroot is found and be part of invert cond TEST_F(OptimizerTest, Case3_3) { - VectorQuery query; + SearchQuery query; query.output_fields_ = {"*"}; query.topk_ = 11; - query.field_name_ = "face_feature"; + query.target_.field_name_ = "face_feature"; query.include_vector_ = false; query.filter_ = "(age > 10 or age > 11) and age > 100"; @@ -257,10 +257,10 @@ TEST_F(OptimizerTest, Case3_3) { // case 3.4 subroot is found and be part of invert cond, but others also have // invert TEST_F(OptimizerTest, Case3_4) { - VectorQuery query; + SearchQuery query; query.output_fields_ = {"*"}; query.topk_ = 11; - query.field_name_ = "face_feature"; + query.target_.field_name_ = "face_feature"; query.include_vector_ = false; query.filter_ = "age > 10 and (age > 101 and (age > 10 and age > 10))"; @@ -281,10 +281,10 @@ TEST_F(OptimizerTest, Case3_4) { // case 4, optimize with in expr TEST_F(OptimizerTest, Case4) { - VectorQuery query; + SearchQuery query; query.output_fields_ = {"*"}; query.topk_ = 11; - query.field_name_ = "face_feature"; + query.target_.field_name_ = "face_feature"; query.include_vector_ = false; query.filter_ = "age in (10, 20)"; diff --git a/tests/db/sqlengine/query_info_test.cc b/tests/db/sqlengine/query_info_test.cc index 56166f5..9beb4e7 100644 --- a/tests/db/sqlengine/query_info_test.cc +++ b/tests/db/sqlengine/query_info_test.cc @@ -86,14 +86,14 @@ class QueryInfoTest : public testing::Test { TEST_F(QueryInfoTest, BasicQueryRequest) { - VectorQuery query; + SearchQuery query; query.output_fields_ = {"*"}; query.topk_ = 11; - query.query_vector_ = "[0.1, 0.2, 0.3, 0.4]"; - query.field_name_ = "face_feature"; + query.target_.set_vector("[0.1, 0.2, 0.3, 0.4]"); + query.target_.field_name_ = "face_feature"; query.include_vector_ = false; - query.query_params_ = std::make_shared(IndexType::FLAT); - query.query_params_->set_radius(0.8F); + query.target_.query_params_ = std::make_shared(IndexType::FLAT); + query.target_.query_params_->set_radius(0.8F); auto engine = std::make_shared(std::make_shared()); auto ret = engine->parse_request(schema, query, nullptr); @@ -116,21 +116,24 @@ TEST_F(QueryInfoTest, BasicQueryRequest) { auto vector_cond = new_query_info->vector_cond_info(); EXPECT_EQ(1, vector_cond->batch()); EXPECT_EQ("face_feature", vector_cond->vector_field_name()); - EXPECT_EQ(query.query_vector_, vector_cond->vector_term()); - EXPECT_EQ(query.query_sparse_indices_, vector_cond->vector_sparse_indices()); - EXPECT_EQ(query.query_sparse_values_, vector_cond->vector_sparse_values()); - EXPECT_EQ(query.query_params_, vector_cond->query_params()); + EXPECT_EQ(std::get(query.target_.clause_).query_vector_, + vector_cond->vector_term()); + EXPECT_EQ(std::get(query.target_.clause_).sparse_indices_, + vector_cond->vector_sparse_indices()); + EXPECT_EQ(std::get(query.target_.clause_).sparse_values_, + vector_cond->vector_sparse_values()); + EXPECT_EQ(query.target_.query_params_, vector_cond->query_params()); } TEST_F(QueryInfoTest, QueryRequestWithFilter) { - VectorQuery query; + SearchQuery query; query.output_fields_ = {"*"}; query.topk_ = 11; - query.query_vector_ = "[0.1, 0.2, 0.3, 0.4]"; - query.field_name_ = "face_feature"; + query.target_.set_vector("[0.1, 0.2, 0.3, 0.4]"); + query.target_.field_name_ = "face_feature"; query.include_vector_ = false; - query.query_params_ = std::make_shared(IndexType::FLAT); - query.query_params_->set_radius(0.8F); + query.target_.query_params_ = std::make_shared(IndexType::FLAT); + query.target_.query_params_->set_radius(0.8F); query.filter_ = "name<3 or name=4 or 1-dash_score_field='test'"; auto engine = std::make_shared(std::make_shared()); @@ -154,10 +157,13 @@ TEST_F(QueryInfoTest, QueryRequestWithFilter) { auto vector_cond = new_query_info->vector_cond_info(); EXPECT_EQ(1, vector_cond->batch()); EXPECT_EQ("face_feature", vector_cond->vector_field_name()); - EXPECT_EQ(query.query_vector_, vector_cond->vector_term()); - EXPECT_EQ(query.query_sparse_indices_, vector_cond->vector_sparse_indices()); - EXPECT_EQ(query.query_sparse_values_, vector_cond->vector_sparse_values()); - EXPECT_EQ(query.query_params_, vector_cond->query_params()); + EXPECT_EQ(std::get(query.target_.clause_).query_vector_, + vector_cond->vector_term()); + EXPECT_EQ(std::get(query.target_.clause_).sparse_indices_, + vector_cond->vector_sparse_indices()); + EXPECT_EQ(std::get(query.target_.clause_).sparse_values_, + vector_cond->vector_sparse_values()); + EXPECT_EQ(query.target_.query_params_, vector_cond->query_params()); EXPECT_TRUE(new_query_info->filter_cond()); // (nullptr) and (xxx) @@ -204,14 +210,14 @@ TEST_F(QueryInfoTest, QueryRequestWithFilter) { } TEST_F(QueryInfoTest, QueryRequestWithIncludeVector) { - VectorQuery query; + SearchQuery query; query.output_fields_ = {"*"}; query.topk_ = 11; - query.query_vector_ = "[0.1, 0.2, 0.3, 0.4]"; - query.field_name_ = "face_feature"; + query.target_.set_vector("[0.1, 0.2, 0.3, 0.4]"); + query.target_.field_name_ = "face_feature"; query.include_vector_ = false; - query.query_params_ = std::make_shared(IndexType::FLAT); - query.query_params_->set_radius(0.8F); + query.target_.query_params_ = std::make_shared(IndexType::FLAT); + query.target_.query_params_->set_radius(0.8F); query.include_vector_ = true; auto engine = std::make_shared(std::make_shared()); @@ -236,21 +242,24 @@ TEST_F(QueryInfoTest, QueryRequestWithIncludeVector) { auto vector_cond = new_query_info->vector_cond_info(); EXPECT_EQ(1, vector_cond->batch()); EXPECT_EQ("face_feature", vector_cond->vector_field_name()); - EXPECT_EQ(query.query_vector_, vector_cond->vector_term()); - EXPECT_EQ(query.query_sparse_indices_, vector_cond->vector_sparse_indices()); - EXPECT_EQ(query.query_sparse_values_, vector_cond->vector_sparse_values()); - EXPECT_EQ(query.query_params_, vector_cond->query_params()); + EXPECT_EQ(std::get(query.target_.clause_).query_vector_, + vector_cond->vector_term()); + EXPECT_EQ(std::get(query.target_.clause_).sparse_indices_, + vector_cond->vector_sparse_indices()); + EXPECT_EQ(std::get(query.target_.clause_).sparse_values_, + vector_cond->vector_sparse_values()); + EXPECT_EQ(query.target_.query_params_, vector_cond->query_params()); } TEST_F(QueryInfoTest, OR_ANCESTOR) { - VectorQuery query; + SearchQuery query; query.output_fields_ = {"*"}; query.topk_ = 11; - query.query_vector_ = "[0.1, 0.2, 0.3, 0.4]"; - query.field_name_ = "face_feature"; + query.target_.set_vector("[0.1, 0.2, 0.3, 0.4]"); + query.target_.field_name_ = "face_feature"; query.include_vector_ = false; - query.query_params_ = std::make_shared(IndexType::FLAT); - query.query_params_->set_radius(0.8F); + query.target_.query_params_ = std::make_shared(IndexType::FLAT); + query.target_.query_params_->set_radius(0.8F); query.filter_ = "name=1 and (name=2 or name=3)"; auto engine = std::make_shared(std::make_shared()); @@ -260,14 +269,14 @@ TEST_F(QueryInfoTest, OR_ANCESTOR) { } TEST_F(QueryInfoTest, QueryRequestWithInFilter) { - VectorQuery query; + SearchQuery query; query.output_fields_ = {"*"}; query.topk_ = 10; - query.query_vector_ = "[0.1, 0.2, 0.3, 0.4]"; - query.field_name_ = "face_feature"; + query.target_.set_vector("[0.1, 0.2, 0.3, 0.4]"); + query.target_.field_name_ = "face_feature"; query.include_vector_ = false; - query.query_params_ = std::make_shared(IndexType::FLAT); - query.query_params_->set_radius(0.8F); + query.target_.query_params_ = std::make_shared(IndexType::FLAT); + query.target_.query_params_->set_radius(0.8F); query.filter_ = "name=3 or name in (1, 2, 3) or category not in (\"a\", \"b\", \"c\")"; @@ -294,7 +303,8 @@ TEST_F(QueryInfoTest, QueryRequestWithInFilter) { EXPECT_EQ(1, vector_cond->batch()); EXPECT_EQ("face_feature", vector_cond->vector_field_name()); std::vector data{1.1, 2.2, 3.3, 4.4}; - EXPECT_EQ(query.query_vector_, vector_cond->vector_term()); + EXPECT_EQ(std::get(query.target_.clause_).query_vector_, + vector_cond->vector_term()); EXPECT_TRUE(new_query_info->filter_cond()); // (nullptr) and (xxx) @@ -350,14 +360,14 @@ TEST_F(QueryInfoTest, QueryRequestWithInFilter) { TEST_F(QueryInfoTest, QueryRequestWithInFilterWrong) { - VectorQuery query; + SearchQuery query; query.output_fields_ = {"*"}; query.topk_ = 11; - query.query_vector_ = "[0.1, 0.2, 0.3, 0.4]"; - query.field_name_ = "face_feature"; + query.target_.set_vector("[0.1, 0.2, 0.3, 0.4]"); + query.target_.field_name_ = "face_feature"; query.include_vector_ = false; - query.query_params_ = std::make_shared(IndexType::FLAT); - query.query_params_->set_radius(0.8F); + query.target_.query_params_ = std::make_shared(IndexType::FLAT); + query.target_.query_params_->set_radius(0.8F); auto engine = std::make_shared(std::make_shared()); auto ret = engine->parse_request(schema, query, nullptr); @@ -381,14 +391,14 @@ TEST_F(QueryInfoTest, QueryRequestWithInFilterWrong) { } TEST_F(QueryInfoTest, QueryRequestWithInFilterNum1024) { - VectorQuery query; + SearchQuery query; query.output_fields_ = {"*"}; query.topk_ = 10; - query.query_vector_ = "[0.1, 0.2, 0.3, 0.4]"; - query.field_name_ = "face_feature"; + query.target_.set_vector("[0.1, 0.2, 0.3, 0.4]"); + query.target_.field_name_ = "face_feature"; query.include_vector_ = false; - query.query_params_ = std::make_shared(IndexType::FLAT); - query.query_params_->set_radius(0.8F); + query.target_.query_params_ = std::make_shared(IndexType::FLAT); + query.target_.query_params_->set_radius(0.8F); std::string filter_str; for (int i = 0; i < 1024; i++) { @@ -425,14 +435,14 @@ TEST_F(QueryInfoTest, QueryRequestWithInFilterNum1024) { TEST_F(QueryInfoTest, QueryRequestWithFilter_contain) { - VectorQuery query; + SearchQuery query; query.output_fields_ = {"*"}; query.topk_ = 10; - query.query_vector_ = "[0.1, 0.2, 0.3, 0.4]"; - query.field_name_ = "face_feature"; + query.target_.set_vector("[0.1, 0.2, 0.3, 0.4]"); + query.target_.field_name_ = "face_feature"; query.include_vector_ = false; - query.query_params_ = std::make_shared(IndexType::FLAT); - query.query_params_->set_radius(0.8F); + query.target_.query_params_ = std::make_shared(IndexType::FLAT); + query.target_.query_params_->set_radius(0.8F); query.filter_ = R"( name_array contain_all (1, 2, 3) and )" R"( (name_array not contain_all (4, 5) or category_array contain_any @@ -582,7 +592,7 @@ TEST_F(QueryInfoTest, QueryRequestWithFilter_contain) { } TEST_F(QueryInfoTest, SelectNonExistField) { - VectorQuery query; + SearchQuery query; query.output_fields_ = {"category_array", "not_exist_field"}; query.topk_ = 11; query.include_vector_ = false; @@ -595,7 +605,7 @@ TEST_F(QueryInfoTest, SelectNonExistField) { } TEST_F(QueryInfoTest, ContainAllExceedLimit) { - VectorQuery query; + SearchQuery query; query.topk_ = 200; query.filter_ = "name_array not contain_all ("; for (int i = 0; i <= 32; i++) { @@ -614,7 +624,7 @@ TEST_F(QueryInfoTest, ContainAllExceedLimit) { } TEST_F(QueryInfoTest, ContainAnyExceedLimit) { - VectorQuery query; + SearchQuery query; query.topk_ = 200; query.filter_ = "name_array not contain_any ("; for (int i = 0; i <= 32; i++) { @@ -633,7 +643,7 @@ TEST_F(QueryInfoTest, ContainAnyExceedLimit) { } TEST_F(QueryInfoTest, ArrayLengthNonExistField) { - VectorQuery query; + SearchQuery query; query.topk_ = 200; query.filter_ = "array_length(not_exist_field) > 1"; auto engine = std::make_shared(std::make_shared()); @@ -644,7 +654,7 @@ TEST_F(QueryInfoTest, ArrayLengthNonExistField) { } TEST_F(QueryInfoTest, ArrayLengthOnNonArrayField) { - VectorQuery query; + SearchQuery query; query.topk_ = 200; query.filter_ = "array_length(name) > 1"; auto engine = std::make_shared(std::make_shared()); @@ -655,7 +665,7 @@ TEST_F(QueryInfoTest, ArrayLengthOnNonArrayField) { } TEST_F(QueryInfoTest, ArrayLengthInvalidArgument) { - VectorQuery query; + SearchQuery query; query.topk_ = 200; query.filter_ = "array_length(name_array) > '1'"; auto engine = std::make_shared(std::make_shared()); @@ -667,7 +677,7 @@ TEST_F(QueryInfoTest, ArrayLengthInvalidArgument) { } TEST_F(QueryInfoTest, ArrayLengthInvalidOp) { - VectorQuery query; + SearchQuery query; query.topk_ = 200; query.filter_ = "array_length(name_array) like '%'"; auto engine = std::make_shared(std::make_shared()); diff --git a/tests/db/sqlengine/simple_rewriter_test.cc b/tests/db/sqlengine/simple_rewriter_test.cc index c2a1d62..ad23107 100644 --- a/tests/db/sqlengine/simple_rewriter_test.cc +++ b/tests/db/sqlengine/simple_rewriter_test.cc @@ -188,7 +188,7 @@ class SimpleRewriterTest : public testing::Test { static void TearDownTestSuite() {} QueryInfo::Ptr parse(const std::string &filter) { - VectorQuery query; + SearchQuery query; query.output_fields_ = {"*"}; query.topk_ = 11; query.include_vector_ = false; diff --git a/tests/db/sqlengine/sqlengine_test.cc b/tests/db/sqlengine/sqlengine_test.cc index 8c4c900..f03a27e 100644 --- a/tests/db/sqlengine/sqlengine_test.cc +++ b/tests/db/sqlengine/sqlengine_test.cc @@ -53,7 +53,7 @@ class SqlEngineTest : public testing::Test { TEST_F(SqlEngineTest, Forward) { std::vector segments = {std::make_shared()}; - VectorQuery query; + SearchQuery query; query.output_fields_ = {"id", "name", "age", "tag_list"}; query.topk_ = 11; // query.filter_ = "id > 3 and score < 0.1"; @@ -76,20 +76,19 @@ TEST_F(SqlEngineTest, Forward) { TEST_F(SqlEngineTest, Vector) { std::vector segments = {std::make_shared()}; - VectorQuery query; + SearchQuery query; query.output_fields_ = {"id", "name", "score"}; query.topk_ = 11; query.filter_ = "id > 3 and score < 0.1"; if (const char *env_var = std::getenv("FILTER"); env_var != nullptr) { query.filter_ = env_var; } - // query.query_vector_ = "[0.1, 0.2, 0.3, 0.4]"; - query.query_sparse_indices_ = "[0, 1, 2, 3]"; - query.query_sparse_values_ = "[0.1, 0.2, 0.3, 0.4]"; - query.field_name_ = "vector"; + // query.target_.set_vector("[0.1, 0.2, 0.3, 0.4]"); + query.target_.set_sparse_vector("[0, 1, 2, 3]", "[0.1, 0.2, 0.3, 0.4]"); + query.target_.field_name_ = "vector"; query.include_vector_ = true; - query.query_params_ = std::make_shared(IndexType::FLAT); - query.query_params_->set_radius(0.8F); + query.target_.query_params_ = std::make_shared(IndexType::FLAT); + query.target_.query_params_->set_radius(0.8F); auto engine = SQLEngine::create(std::make_shared()); auto ret = engine->execute(schema_, query, segments); @@ -101,7 +100,7 @@ TEST_F(SqlEngineTest, Vector) { TEST_F(SqlEngineTest, Invert) { std::vector segments = {std::make_shared()}; - VectorQuery query; + SearchQuery query; query.output_fields_ = {"id", "age", "score"}; query.topk_ = 11; // query.filter_ = "name = 'test_name'"; @@ -122,11 +121,11 @@ TEST_F(SqlEngineTest, Invert) { TEST_F(SqlEngineTest, MultiSegments) { std::vector segments = {std::make_shared(), std::make_shared()}; - VectorQuery query; + SearchQuery query; query.output_fields_ = {"id", "name", "age", "score"}; query.topk_ = 11; - query.query_vector_ = "[0.1, 0.2, 0.3, 0.4]"; - query.field_name_ = "vector"; + query.target_.set_vector("[0.1, 0.2, 0.3, 0.4]"); + query.target_.field_name_ = "vector"; // query.filter_ = "name = 'test_name'"; if (const char *env_var = std::getenv("FILTER"); env_var != nullptr) { query.filter_ = env_var; @@ -151,13 +150,12 @@ TEST_F(SqlEngineTest, GroupBy) { if (const char *env_var = std::getenv("FILTER"); env_var != nullptr) { query.filter_ = env_var; } - // query.query_vector_ = "[0.1, 0.2, 0.3, 0.4]"; - query.query_sparse_indices_ = "[0, 1, 2, 3]"; - query.query_sparse_values_ = "[0.1, 0.2, 0.3, 0.4]"; - query.field_name_ = "vector"; + // query.target_.set_vector("[0.1, 0.2, 0.3, 0.4]"); + query.target_.set_sparse_vector("[0, 1, 2, 3]", "[0.1, 0.2, 0.3, 0.4]"); + query.target_.field_name_ = "vector"; query.include_vector_ = true; - query.query_params_ = std::make_shared(IndexType::FLAT); - query.query_params_->set_radius(0.8F); + query.target_.query_params_ = std::make_shared(IndexType::FLAT); + query.target_.query_params_->set_radius(0.8F); auto engine = SQLEngine::create(std::make_shared()); auto ret = engine->execute_group_by(schema_, query, segments); diff --git a/tests/db/sqlengine/vector_recall_test.cc b/tests/db/sqlengine/vector_recall_test.cc index d3dbccd..f034c71 100644 --- a/tests/db/sqlengine/vector_recall_test.cc +++ b/tests/db/sqlengine/vector_recall_test.cc @@ -23,13 +23,13 @@ namespace zvec::sqlengine { class VectorRecallTest : public RecallTest {}; TEST_F(VectorRecallTest, Basic) { - VectorQuery query; + SearchQuery query; query.output_fields_ = {"id", "name", "age"}; query.topk_ = 200; std::vector feature(4, 0.0); - query.query_vector_.assign((const char *)feature.data(), - feature.size() * sizeof(float)); - query.field_name_ = "dense"; + query.target_.set_vector(std::string((const char *)feature.data(), + feature.size() * sizeof(float))); + query.target_.field_name_ = "dense"; auto engine = SQLEngine::create(std::make_shared()); auto ret = engine->execute(collection_schema_, query, segments_); @@ -52,14 +52,14 @@ TEST_F(VectorRecallTest, Basic) { } TEST_F(VectorRecallTest, HybridInvertFilter) { - VectorQuery query; + SearchQuery query; query.output_fields_ = {"id", "name", "age"}; query.filter_ = "invert_id >= 1"; query.topk_ = 200; std::vector feature(4, 0.0); - query.query_vector_.assign((const char *)feature.data(), - feature.size() * sizeof(float)); - query.field_name_ = "dense"; + query.target_.set_vector(std::string((const char *)feature.data(), + feature.size() * sizeof(float))); + query.target_.field_name_ = "dense"; auto engine = SQLEngine::create(std::make_shared()); auto ret = engine->execute(collection_schema_, query, segments_); @@ -83,14 +83,14 @@ TEST_F(VectorRecallTest, HybridInvertFilter) { } TEST_F(VectorRecallTest, HybridInvertFilterBfByKeys) { - VectorQuery query; + SearchQuery query; query.output_fields_ = {"id", "name", "age"}; query.filter_ = "invert_id < 199"; query.topk_ = 199; std::vector feature(4, 0.0); - query.query_vector_.assign((const char *)feature.data(), - feature.size() * sizeof(float)); - query.field_name_ = "dense"; + query.target_.set_vector(std::string((const char *)feature.data(), + feature.size() * sizeof(float))); + query.target_.field_name_ = "dense"; auto engine = SQLEngine::create(std::make_shared()); auto ret = engine->execute(collection_schema_, query, segments_); @@ -113,14 +113,14 @@ TEST_F(VectorRecallTest, HybridInvertFilterBfByKeys) { } TEST_F(VectorRecallTest, HybridForwardFilter) { - VectorQuery query; + SearchQuery query; query.output_fields_ = {"id", "name", "age"}; query.filter_ = "id >= 1"; query.topk_ = 200; std::vector feature(4, 0.0); - query.query_vector_.assign((const char *)feature.data(), - feature.size() * sizeof(float)); - query.field_name_ = "dense"; + query.target_.set_vector(std::string((const char *)feature.data(), + feature.size() * sizeof(float))); + query.target_.field_name_ = "dense"; auto engine = SQLEngine::create(std::make_shared()); auto ret = engine->execute(collection_schema_, query, segments_); @@ -144,14 +144,14 @@ TEST_F(VectorRecallTest, HybridForwardFilter) { } TEST_F(VectorRecallTest, HybridInvertForwardFilter) { - VectorQuery query; + SearchQuery query; query.output_fields_ = {"name", "age"}; query.filter_ = "invert_id >= 1 and id <= 100"; query.topk_ = 200; std::vector feature(4, 0.0); - query.query_vector_.assign((const char *)feature.data(), - feature.size() * sizeof(float)); - query.field_name_ = "dense"; + query.target_.set_vector(std::string((const char *)feature.data(), + feature.size() * sizeof(float))); + query.target_.field_name_ = "dense"; auto engine = SQLEngine::create(std::make_shared()); auto ret = engine->execute(collection_schema_, query, segments_); @@ -175,16 +175,17 @@ TEST_F(VectorRecallTest, HybridInvertForwardFilter) { } TEST_F(VectorRecallTest, Sparse) { - VectorQuery query; + SearchQuery query; query.output_fields_ = {"id", "name", "age"}; query.topk_ = 200; std::vector feature(4, 1.0); std::vector indices{0, 1, 2, 3}; - query.query_sparse_indices_.assign((const char *)indices.data(), - indices.size() * sizeof(uint32_t)); - query.query_sparse_values_.assign((const char *)feature.data(), - feature.size() * sizeof(float)); - query.field_name_ = "sparse"; + query.target_.set_sparse_vector( + std::string((const char *)indices.data(), + indices.size() * sizeof(uint32_t)), + std::string((const char *)feature.data(), + feature.size() * sizeof(float))); + query.target_.field_name_ = "sparse"; auto engine = SQLEngine::create(std::make_shared()); auto ret = engine->execute(collection_schema_, query, segments_); @@ -218,13 +219,13 @@ TEST_F(VectorRecallTest, DeleteFilter) { segments_[0]->Delete("pk_" + std::to_string(i)); } - VectorQuery query; + SearchQuery query; query.output_fields_ = {"name", "age"}; query.topk_ = 100; std::vector feature(4, 0.0); - query.query_vector_.assign((const char *)feature.data(), - feature.size() * sizeof(float)); - query.field_name_ = "dense"; + query.target_.set_vector(std::string((const char *)feature.data(), + feature.size() * sizeof(float))); + query.target_.field_name_ = "dense"; auto engine = SQLEngine::create(std::make_shared()); auto ret = engine->execute(collection_schema_, query, segments_); @@ -249,14 +250,14 @@ TEST_F(VectorRecallTest, DeleteFilter) { TEST_F(VectorRecallTest, HybridInvertForwardDeleteFilter) { // In previous test, docs[0-4000) has been deleted - VectorQuery query; + SearchQuery query; query.output_fields_ = {"name", "age"}; query.filter_ = "invert_id >= 6000 and id < 6080"; query.topk_ = 100; std::vector feature(4, 0.0); - query.query_vector_.assign((const char *)feature.data(), - feature.size() * sizeof(float)); - query.field_name_ = "dense"; + query.target_.set_vector(std::string((const char *)feature.data(), + feature.size() * sizeof(float))); + query.target_.field_name_ = "dense"; auto engine = SQLEngine::create(std::make_shared()); auto ret = engine->execute(collection_schema_, query, segments_);