refactor: drop VectorQuery, unify single-target query on SearchQuery (#428)

This commit is contained in:
egolearner 2026-05-29 16:36:09 +08:00 committed by GitHub
parent f539580138
commit 8dcb6cbd7f
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
29 changed files with 830 additions and 718 deletions

View File

@ -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<float> query_vector = std::vector<float>(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;

View File

@ -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<zvec_vector_query_t *>(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<zvec::VectorQuery *>(query);
delete reinterpret_cast<zvec::SearchQuery *>(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<zvec::VectorQuery *>(query);
auto *ptr = reinterpret_cast<zvec::SearchQuery *>(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<const zvec::VectorQuery *>(query);
auto *ptr = reinterpret_cast<const zvec::SearchQuery *>(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<zvec::VectorQuery *>(query);
ptr->field_name_ = field_name ? field_name : "";
auto *ptr = reinterpret_cast<zvec::SearchQuery *>(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<const zvec::VectorQuery *>(query);
return ptr->field_name_.empty() ? nullptr : ptr->field_name_.c_str();
auto *ptr = reinterpret_cast<const zvec::SearchQuery *>(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<zvec::VectorQuery *>(query);
ptr->query_vector_.assign(static_cast<const char *>(data), size);
auto *ptr = reinterpret_cast<zvec::SearchQuery *>(query);
ptr->target_.set_vector(std::string(static_cast<const char *>(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<zvec::VectorQuery *>(query);
auto *ptr = reinterpret_cast<zvec::SearchQuery *>(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<const zvec::VectorQuery *>(query);
auto *ptr = reinterpret_cast<const zvec::SearchQuery *>(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<zvec::VectorQuery *>(query);
auto *ptr = reinterpret_cast<zvec::SearchQuery *>(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<const zvec::VectorQuery *>(query);
auto *ptr = reinterpret_cast<const zvec::SearchQuery *>(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<zvec::VectorQuery *>(query);
auto *ptr = reinterpret_cast<zvec::SearchQuery *>(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<const zvec::VectorQuery *>(query);
auto *ptr = reinterpret_cast<const zvec::SearchQuery *>(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<zvec::VectorQuery *>(query);
auto *ptr = reinterpret_cast<zvec::SearchQuery *>(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<const zvec::VectorQuery *>(query);
auto *ptr = reinterpret_cast<const zvec::SearchQuery *>(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<zvec::VectorQuery *>(query);
auto *query_ptr = reinterpret_cast<zvec::SearchQuery *>(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<zvec::QueryParams *>(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<zvec::VectorQuery *>(query);
auto *query_ptr = reinterpret_cast<zvec::SearchQuery *>(query);
auto *params_ptr = reinterpret_cast<zvec::HnswQueryParams *>(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<zvec::VectorQuery *>(query);
auto *query_ptr = reinterpret_cast<zvec::SearchQuery *>(query);
auto *params_ptr = reinterpret_cast<zvec::IVFQueryParams *>(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<zvec::VectorQuery *>(query);
auto *query_ptr = reinterpret_cast<zvec::SearchQuery *>(query);
auto *params_ptr = reinterpret_cast<zvec::FlatQueryParams *>(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<zvec::GroupByVectorQuery *>(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<const zvec::GroupByVectorQuery *>(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<zvec::GroupByVectorQuery *>(query);
ptr->query_vector_.assign(static_cast<const char *>(data), size);
ptr->target_.set_vector(std::string(static_cast<const char *>(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<zvec::GroupByVectorQuery *>(query);
auto *params_ptr = reinterpret_cast<zvec::QueryParams *>(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<zvec::GroupByVectorQuery *>(query);
auto *params_ptr = reinterpret_cast<zvec::HnswQueryParams *>(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<zvec::GroupByVectorQuery *>(query);
auto *params_ptr = reinterpret_cast<zvec::IVFQueryParams *>(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<zvec::GroupByVectorQuery *>(query);
auto *params_ptr = reinterpret_cast<zvec::FlatQueryParams *>(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<const std::shared_ptr<zvec::Collection> *>(
collection);
// Cast zvec_vector_query_t* to zvec::VectorQuery* directly
// zvec_vector_query_t wraps zvec::SearchQuery internally.
auto *internal_query =
reinterpret_cast<const zvec::VectorQuery *>(query);
reinterpret_cast<const zvec::SearchQuery *>(query);
auto result = (*coll_ptr)->Query(*internal_query);
zvec_error_code_t error_code = handle_expected_result(result);

View File

@ -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<VectorClause>(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_<VectorQuery>(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_<SearchQuery>(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<float>(
static_cast<const float *>(buf.ptr), buf.size);
self.target_.set_vector(serialize_vector<float>(
static_cast<const float *>(buf.ptr), buf.size));
return;
}
case DataType::VECTOR_FP64: {
self.query_vector_ = serialize_vector<double>(
static_cast<const double *>(buf.ptr), buf.size);
self.target_.set_vector(serialize_vector<double>(
static_cast<const double *>(buf.ptr), buf.size));
return;
}
case DataType::VECTOR_INT8: {
self.query_vector_ = serialize_vector<int8_t>(
static_cast<const int8_t *>(buf.ptr), buf.size);
self.target_.set_vector(serialize_vector<int8_t>(
static_cast<const int8_t *>(buf.ptr), buf.size));
return;
}
case DataType::VECTOR_FP16: {
self.query_vector_ = serialize_vector<uint16_t>(
static_cast<const uint16_t *>(buf.ptr), buf.size);
self.target_.set_vector(serialize_vector<uint16_t>(
static_cast<const uint16_t *>(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<const uint32_t *>(
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<const float *>(
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<const uint16_t *>(
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<int>();
obj.field_name_ = t[1].cast<std::string>();
obj.query_vector_ = t[2].cast<std::string>();
obj.query_sparse_indices_ = t[3].cast<std::string>();
obj.query_sparse_values_ = t[4].cast<std::string>();
obj.target_.field_name_ = t[1].cast<std::string>();
obj.target_.clause_ =
VectorClause{t[2].cast<std::string>(), t[3].cast<std::string>(),
t[4].cast<std::string>()};
obj.filter_ = t[5].cast<std::string>();
obj.include_vector_ = t[6].cast<bool>();
obj.output_fields_ = t[7].cast<std::vector<std::string>>();
if (!t[8].is_none()) {
obj.query_params_ = t[8].cast<QueryParams::Ptr>();
obj.target_.query_params_ = t[8].cast<QueryParams::Ptr>();
}
return obj;
}));
}
} // namespace zvec
} // namespace zvec

View File

@ -252,7 +252,7 @@ void ZVecPyCollection::bind_dml_methods(
void ZVecPyCollection::bind_dql_methods(
py::class_<Collection, Collection::Ptr> &col) {
col.def("Query",
[](const Collection &self, const VectorQuery &query) {
[](const Collection &self, const SearchQuery &query) {
Result<DocPtrList> result;
{
py::gil_scoped_release release;

View File

@ -118,7 +118,7 @@ class CollectionImpl : public Collection {
Status DeleteByFilter(const std::string &filter) override;
Result<DocPtrList> Query(const VectorQuery &query) const override;
Result<DocPtrList> Query(const SearchQuery &query) const override;
Result<DocPtrList> 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<std::string>{};
@ -1584,14 +1584,14 @@ Status CollectionImpl::DeleteByFilter(const std::string &filter) {
return Status::OK();
}
Result<DocPtrList> CollectionImpl::Query(const VectorQuery &query) const {
Result<DocPtrList> 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<DocPtrList> CollectionImpl::Query(const MultiQuery &query) const {
return DocPtrList();
}
// Convert SubVectorQuery to VectorQuery and validate
// Convert each SubQuery to a SearchQuery and validate.
std::set<std::string> seen_fields;
std::vector<VectorQuery> converted_queries;
std::vector<SearchQuery> converted_queries;
converted_queries.reserve(query.queries.size());
for (const auto &sub : query.queries) {
@ -1640,31 +1640,26 @@ Result<DocPtrList> 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<VectorClause>(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<std::future<Result<DocPtrList>>> 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<Profiler>());
return engine->execute(schema_, vq, segments);
return engine->execute(schema_, sq, segments);
}));
}
@ -1674,7 +1669,8 @@ Result<DocPtrList> 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

View File

@ -187,45 +187,6 @@ template <class... Ts>
overloaded(Ts...) -> overloaded<Ts...>;
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<size_t> 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<uint32_t> sorted_indices(n);
std::vector<char> 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<uint32_t *>(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

View File

@ -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 <cstdint>
#include <zvec/db/query.h>
#include <zvec/db/schema.h>
#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<uint32_t *>(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

View File

@ -13,10 +13,53 @@
// limitations under the License.
#include "type_helper.h"
#include <algorithm>
#include <cstring>
#include <numeric>
#include <vector>
#include <zvec/core/framework/index_meta.h>
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<size_t> 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<uint32_t> sorted_indices(n);
std::vector<char> 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:

View File

@ -14,12 +14,19 @@
#pragma once
#include <cstddef>
#include <cstdint>
#include <zvec/core/framework/index_meta.h>
#include <zvec/db/type.h>
#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

View File

@ -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<Node>(NodeOp::T_EQ);
rel_exp->set_left(std::make_shared<IDNode>(request.field_name_));
rel_exp->set_left(std::make_shared<IDNode>(request.target_.field_name_));
rel_exp->set_right(std::make_shared<VectorMatrixNode>(
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<SelectedElemInfo>();
@ -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<GroupBy> 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;

View File

@ -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<GroupBy> group_by,
sqlengine::SQLInfo::Ptr *sql_info,
std::string *err_msg);

View File

@ -27,7 +27,7 @@ class SQLEngine {
virtual ~SQLEngine();
virtual Result<DocPtrList> execute(
CollectionSchema::Ptr collection, const VectorQuery &query,
CollectionSchema::Ptr collection, const SearchQuery &query,
const std::vector<Segment::Ptr> &segments) = 0;
virtual Result<GroupResults> execute_group_by(

View File

@ -50,7 +50,7 @@ SQLEngineImpl::SQLEngineImpl(zvec::Profiler::Ptr profiler)
: profiler_(std::move(profiler)) {}
Result<DocPtrList> SQLEngineImpl::execute(
CollectionSchema::Ptr collection, const VectorQuery &query,
CollectionSchema::Ptr collection, const SearchQuery &query,
const std::vector<Segment::Ptr> &segments) {
if (segments.empty()) {
return DocPtrList{};
@ -76,18 +76,14 @@ Result<DocPtrList> 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<GroupResults> SQLEngineImpl::execute_group_by(
@ -97,7 +93,7 @@ Result<GroupResults> 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<GroupBy>(group_by_query.group_by_field_name_,
@ -135,7 +131,7 @@ Result<QueryInfo::Ptr> SQLEngineImpl::parse_sql_info(
}
Result<QueryInfo::Ptr> SQLEngineImpl::parse_request(
CollectionSchema::Ptr collection, const VectorQuery &request,
CollectionSchema::Ptr collection, const SearchQuery &request,
std::shared_ptr<GroupBy> group_by) {
profiler_->open_stage("message_to_sqlinfo");
sqlengine::SQLInfo::Ptr sql_info;

View File

@ -34,7 +34,7 @@ class SQLEngineImpl : public SQLEngine {
//! Parse pb request
Result<QueryInfo::Ptr> parse_request(CollectionSchema::Ptr collection,
const VectorQuery &request,
const SearchQuery &request,
std::shared_ptr<GroupBy> group_by);
//! Perform search with given query_info, segments and index filter
@ -44,7 +44,7 @@ class SQLEngineImpl : public SQLEngine {
std::vector<sqlengine::QueryInfo::Ptr> *query_infos);
Result<DocPtrList> execute(
CollectionSchema::Ptr collection, const VectorQuery &query,
CollectionSchema::Ptr collection, const SearchQuery &query,
const std::vector<Segment::Ptr> &segments) override;
Result<GroupResults> execute_group_by(

View File

@ -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;

View File

@ -97,7 +97,7 @@ class Collection {
virtual Status DeleteByFilter(const std::string &filter) = 0;
virtual Result<DocPtrList> Query(const VectorQuery &query) const = 0;
virtual Result<DocPtrList> Query(const SearchQuery &query) const = 0;
virtual Result<DocPtrList> Query(const MultiQuery &query) const = 0;

View File

@ -13,6 +13,7 @@
// limitations under the License.
#pragma once
#include <cstdint>
#include <memory>
#include <optional>
#include <string>
@ -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<std::vector<std::string>> 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<std::vector<std::string>> 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<VectorClause, FtsClause> 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<VectorClause>(&clause_);
}
const VectorClause *get_vector_clause() const {
return std::get_if<VectorClause>(&clause_);
}
private:
// Resets clause_ to an empty VectorClause unless it already holds one.
VectorClause &ensure_vector_clause() {
if (!std::holds_alternative<VectorClause>(clause_)) {
clause_ = VectorClause{};
}
return std::get<VectorClause>(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<std::vector<std::string>> 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<std::vector<std::string>> 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<Doc> docs_;
};
using GroupResults = std::vector<GroupResult>;
//! 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<std::vector<std::string>> output_fields;
std::shared_ptr<Reranker> reranker{nullptr};
};
struct GroupResult {
std::string group_by_value_;
std::vector<Doc> docs_;
};
using GroupResults = std::vector<GroupResult>;
} // namespace zvec

View File

@ -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<ailego::Float16>(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<float>(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<int8_t>(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<std::vector<float>>("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<std::vector<float>>(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::pair<std::vector<uint32_t>, std::vector<float>>>(
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<std::vector<float>>(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::pair<std::vector<uint32_t>, std::vector<float>>>(
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<std::string>>(
std::vector<std::string>(1025));
@ -3406,18 +3406,18 @@ TEST_F(CollectionTest, Feature_Query_Validate) {
if (field_scheama->is_dense_vector()) {
auto vector = query_doc.get<std::vector<float>>(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::pair<std::vector<uint32_t>, std::vector<float>>>(
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<std::vector<float>>(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::pair<std::vector<uint32_t>, std::vector<float>>>(
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<std::vector<float>>("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<std::string>{"int32"};

View File

@ -192,12 +192,12 @@ TEST_F(OptimizeRecoveryTest, CrashDuringOptimize) {
}
}
VectorQuery query;
SearchQuery query;
query.topk_ = 10;
std::vector<float> 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<float> 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();

View File

@ -823,8 +823,7 @@ TEST_F(DocDetailedTest, ValidateAndSanitization) {
auto schema = test::TestHelper::CreateNormalSchema(false);
std::vector<std::string> 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<float> query_vector = {1.0f, 2.0f, 3.0f, 4.0f};
std::string query_vector_str =
std::string(reinterpret_cast<char *>(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<std::string>(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<float> query_vector = {1.0f, 2.0f, 3.0f, 4.0f};
std::string query_vector_str =
std::string(reinterpret_cast<char *>(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<uint32_t> query_indices(16385);
std::vector<float> query_values(16385);
query.query_sparse_indices_ =
query.target_.set_sparse_vector(
std::string(reinterpret_cast<char *>(query_indices.data()),
query_indices.size() * sizeof(uint32_t));
query.query_sparse_values_ =
query_indices.size() * sizeof(uint32_t)),
std::string(reinterpret_cast<char *>(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<char *>(&one_index), sizeof(uint32_t));
query.query_sparse_values_ =
std::string(reinterpret_cast<char *>(&one_value), sizeof(float));
query.target_.set_sparse_vector(
std::string(reinterpret_cast<char *>(&one_index), sizeof(uint32_t)),
std::string(reinterpret_cast<char *>(&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<uint32_t>{3u, 7u, 42u, 99u, 128u}));
EXPECT_EQ(decode_val(query.query_sparse_values_),
(std::vector<float>{0.4f, 0.2f, 0.1f, 0.5f, 0.3f}));
EXPECT_EQ(
decode_idx(
std::get<VectorClause>(query.target_.clause_).sparse_indices_),
(std::vector<uint32_t>{3u, 7u, 42u, 99u, 128u}));
EXPECT_EQ(
decode_val(
std::get<VectorClause>(query.target_.clause_).sparse_values_),
(std::vector<float>{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<VectorClause>(query.target_.clause_).sparse_indices_,
pack_idx({1u, 2u, 3u, 4u}));
EXPECT_EQ(std::get<VectorClause>(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<VectorClause>(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<float> query_vector(128, 1.0f);
query.query_vector_ =
query.target_.set_vector(
std::string(reinterpret_cast<char *>(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<HnswIndexParams>(MetricType::L2));
query.query_params_ = std::make_shared<HnswQueryParams>(150);
query.target_.query_params_ = std::make_shared<HnswQueryParams>(150);
auto s = query.validate_and_sanitize(&schema);
EXPECT_TRUE(s.ok());
query.query_params_ = std::make_shared<IVFQueryParams>(50);
query.target_.query_params_ = std::make_shared<IVFQueryParams>(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());
}

View File

@ -124,7 +124,7 @@ class ContainTest : public testing::Test {
TEST_F(ContainTest, ContainAllInt32) {
VectorQuery query;
SearchQuery query;
query.output_fields_ = std::vector<std::string>{};
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<std::string>{};
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<std::string>{};
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<std::string>{};
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<std::string>{};
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<std::string>{};
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<std::string>{};
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<std::string>{};
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<std::string>{};
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<std::string>{};
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<std::string>{};
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<std::string>{};
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<std::string>{};
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<std::string>{};
query.topk_ = 200;
query.filter_ = "str_array contain_any ('name98','name99','name100')";

View File

@ -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<std::string>{};
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<std::string>();
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";

View File

@ -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<std::string>();
query.topk_ = 200;
query.include_vector_ = true;

View File

@ -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')";

View File

@ -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)";

View File

@ -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<QueryParams>(IndexType::FLAT);
query.query_params_->set_radius(0.8F);
query.target_.query_params_ = std::make_shared<QueryParams>(IndexType::FLAT);
query.target_.query_params_->set_radius(0.8F);
auto engine = std::make_shared<SQLEngineImpl>(std::make_shared<Profiler>());
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<VectorClause>(query.target_.clause_).query_vector_,
vector_cond->vector_term());
EXPECT_EQ(std::get<VectorClause>(query.target_.clause_).sparse_indices_,
vector_cond->vector_sparse_indices());
EXPECT_EQ(std::get<VectorClause>(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<QueryParams>(IndexType::FLAT);
query.query_params_->set_radius(0.8F);
query.target_.query_params_ = std::make_shared<QueryParams>(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<SQLEngineImpl>(std::make_shared<Profiler>());
@ -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<VectorClause>(query.target_.clause_).query_vector_,
vector_cond->vector_term());
EXPECT_EQ(std::get<VectorClause>(query.target_.clause_).sparse_indices_,
vector_cond->vector_sparse_indices());
EXPECT_EQ(std::get<VectorClause>(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<QueryParams>(IndexType::FLAT);
query.query_params_->set_radius(0.8F);
query.target_.query_params_ = std::make_shared<QueryParams>(IndexType::FLAT);
query.target_.query_params_->set_radius(0.8F);
query.include_vector_ = true;
auto engine = std::make_shared<SQLEngineImpl>(std::make_shared<Profiler>());
@ -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<VectorClause>(query.target_.clause_).query_vector_,
vector_cond->vector_term());
EXPECT_EQ(std::get<VectorClause>(query.target_.clause_).sparse_indices_,
vector_cond->vector_sparse_indices());
EXPECT_EQ(std::get<VectorClause>(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<QueryParams>(IndexType::FLAT);
query.query_params_->set_radius(0.8F);
query.target_.query_params_ = std::make_shared<QueryParams>(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<SQLEngineImpl>(std::make_shared<Profiler>());
@ -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<QueryParams>(IndexType::FLAT);
query.query_params_->set_radius(0.8F);
query.target_.query_params_ = std::make_shared<QueryParams>(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<float> data{1.1, 2.2, 3.3, 4.4};
EXPECT_EQ(query.query_vector_, vector_cond->vector_term());
EXPECT_EQ(std::get<VectorClause>(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<QueryParams>(IndexType::FLAT);
query.query_params_->set_radius(0.8F);
query.target_.query_params_ = std::make_shared<QueryParams>(IndexType::FLAT);
query.target_.query_params_->set_radius(0.8F);
auto engine = std::make_shared<SQLEngineImpl>(std::make_shared<Profiler>());
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<QueryParams>(IndexType::FLAT);
query.query_params_->set_radius(0.8F);
query.target_.query_params_ = std::make_shared<QueryParams>(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<QueryParams>(IndexType::FLAT);
query.query_params_->set_radius(0.8F);
query.target_.query_params_ = std::make_shared<QueryParams>(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<SQLEngineImpl>(std::make_shared<Profiler>());
@ -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<SQLEngineImpl>(std::make_shared<Profiler>());
@ -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<SQLEngineImpl>(std::make_shared<Profiler>());
@ -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<SQLEngineImpl>(std::make_shared<Profiler>());

View File

@ -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;

View File

@ -53,7 +53,7 @@ class SqlEngineTest : public testing::Test {
TEST_F(SqlEngineTest, Forward) {
std::vector<Segment::Ptr> segments = {std::make_shared<MockSegment>()};
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<Segment::Ptr> segments = {std::make_shared<MockSegment>()};
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<QueryParams>(IndexType::FLAT);
query.query_params_->set_radius(0.8F);
query.target_.query_params_ = std::make_shared<QueryParams>(IndexType::FLAT);
query.target_.query_params_->set_radius(0.8F);
auto engine = SQLEngine::create(std::make_shared<Profiler>());
auto ret = engine->execute(schema_, query, segments);
@ -101,7 +100,7 @@ TEST_F(SqlEngineTest, Vector) {
TEST_F(SqlEngineTest, Invert) {
std::vector<Segment::Ptr> segments = {std::make_shared<MockSegment>()};
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<Segment::Ptr> segments = {std::make_shared<MockSegment>(),
std::make_shared<MockSegment>()};
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<QueryParams>(IndexType::FLAT);
query.query_params_->set_radius(0.8F);
query.target_.query_params_ = std::make_shared<QueryParams>(IndexType::FLAT);
query.target_.query_params_->set_radius(0.8F);
auto engine = SQLEngine::create(std::make_shared<Profiler>());
auto ret = engine->execute_group_by(schema_, query, segments);

View File

@ -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<float> 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<Profiler>());
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<float> 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<Profiler>());
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<float> 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<Profiler>());
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<float> 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<Profiler>());
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<float> 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<Profiler>());
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<float> feature(4, 1.0);
std::vector<uint32_t> 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<Profiler>());
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<float> 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<Profiler>());
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<float> 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<Profiler>());
auto ret = engine->execute(collection_schema_, query, segments_);