From c35d24e21505fea63db53fb035d75236a5dce9da Mon Sep 17 00:00:00 2001 From: egolearner Date: Tue, 7 Jul 2026 14:14:13 +0800 Subject: [PATCH] feat(fts): add stemmer token filter based on Snowball 3.1.1 (#513) --- .gitmodules | 3 + python/zvec/model/param/__init__.pyi | 52 +++++- .../python/model/param/python_param.cc | 58 +++++- src/db/CMakeLists.txt | 1 + src/db/index/CMakeLists.txt | 1 + .../tokenizer/stemmer_token_filter.cc | 80 +++++++++ .../tokenizer/stemmer_token_filter.h | 46 +++++ .../fts_column/tokenizer/token_filter.h | 10 ++ .../fts_column/tokenizer/tokenizer_factory.cc | 10 ++ .../fts_column/tokenizer/tokenizer_factory.h | 5 +- src/db/index/common/schema.cc | 28 +++ src/include/zvec/c_api.h | 31 +++- src/include/zvec/db/index_params.h | 23 ++- .../fts_column/fts_column_indexer_test.cc | 86 +++++++++ .../fts_column/stemmer_token_filter_test.cc | 167 ++++++++++++++++++ tests/db/index/common/schema_test.cc | 21 +++ thirdparty/CMakeLists.txt | 1 + thirdparty/snowball/CMakeLists.txt | 100 +++++++++++ thirdparty/snowball/snowball-3.1.1 | 1 + 19 files changed, 703 insertions(+), 21 deletions(-) create mode 100644 src/db/index/column/fts_column/tokenizer/stemmer_token_filter.cc create mode 100644 src/db/index/column/fts_column/tokenizer/stemmer_token_filter.h create mode 100644 tests/db/index/column/fts_column/stemmer_token_filter_test.cc create mode 100644 thirdparty/snowball/CMakeLists.txt create mode 160000 thirdparty/snowball/snowball-3.1.1 diff --git a/.gitmodules b/.gitmodules index afd11c0..f919c73 100644 --- a/.gitmodules +++ b/.gitmodules @@ -56,3 +56,6 @@ [submodule "thirdparty/utf8proc/utf8proc-2.11.3"] path = thirdparty/utf8proc/utf8proc-2.11.3 url = https://github.com/JuliaStrings/utf8proc.git +[submodule "thirdparty/snowball/snowball-3.1.1"] + path = thirdparty/snowball/snowball-3.1.1 + url = https://github.com/snowballstem/snowball.git diff --git a/python/zvec/model/param/__init__.pyi b/python/zvec/model/param/__init__.pyi index fd087bc..2653d2c 100644 --- a/python/zvec/model/param/__init__.pyi +++ b/python/zvec/model/param/__init__.pyi @@ -718,9 +718,30 @@ class FtsIndexParam(IndexParam): "whitespace"). Default is "standard". filters (list[str]): List of token filter names applied after tokenization. - Supported filters are "lowercase" and "ascii_folding". Default is - ["lowercase"]. - extra_params (str): Additional parameters passed to the tokenizer. + Supported filters are "lowercase", "ascii_folding", and "stemmer". + Default is ["lowercase"]. + extra_params (str): Additional tokenizer/filter parameters as an empty + string or JSON object string. Supported keys are grouped by component: + Tokenizers: + standard: + - "max_token_length" (positive integer). + jieba: + - "jieba_dict_dir" (directory containing jieba.dict.utf8 and + hmm_model.utf8). + - "user_dict_path" (user dictionary path). + - "cut_mode" ("search", "mix", "full", or "hmm"; default + "search"). + whitespace: + - no extra_params. + Filters: + lowercase: + - no extra_params. + ascii_folding: + - no extra_params. + stemmer: + - "stemmer_lang" (Snowball language/algorithm; default + "english"), for example {"stemmer_lang":"porter"} for ES + behaviour. Default is "". Examples: @@ -744,8 +765,29 @@ class FtsIndexParam(IndexParam): Args: tokenizer_name (str, optional): Tokenizer name. Defaults to "standard". filters (list[str], optional): Token filter names. Supports - "lowercase" and "ascii_folding". Defaults to ["lowercase"]. - extra_params (str, optional): Extra tokenizer parameters. Defaults to "". + "lowercase", "ascii_folding", and "stemmer". Defaults to + ["lowercase"]. + extra_params (str, optional): Extra tokenizer/filter parameters as an + empty string or JSON object string. Supported keys: + Tokenizers: + standard: + - "max_token_length" (positive integer). + jieba: + - "jieba_dict_dir". + - "user_dict_path". + - "cut_mode" ("search", "mix", "full", or "hmm"; + default "search"). + whitespace: + - no extra_params. + Filters: + lowercase: + - no extra_params. + ascii_folding: + - no extra_params. + stemmer: + - "stemmer_lang" (Snowball language/algorithm; default + "english"). + Defaults to "". """ def __repr__(self) -> str: ... diff --git a/src/binding/python/model/param/python_param.cc b/src/binding/python/model/param/python_param.cc index 2fa10ac..500455c 100644 --- a/src/binding/python/model/param/python_param.cc +++ b/src/binding/python/model/param/python_param.cc @@ -261,9 +261,30 @@ Attributes: "whitespace"). Default is "standard". filters (list[str]): List of token filter names applied after tokenization. - Supported filters are "lowercase" and "ascii_folding". Default is - ["lowercase"]. - extra_params (str): Additional parameters passed to the tokenizer. + Supported values include "lowercase", "ascii_folding", and "stemmer". + Default is ["lowercase"]. + extra_params (str): Additional tokenizer/filter parameters as an empty + string or JSON object string. Supported keys are grouped by component: + Tokenizers: + standard: + - "max_token_length" (positive integer). + jieba: + - "jieba_dict_dir" (directory containing jieba.dict.utf8 and + hmm_model.utf8). + - "user_dict_path" (user dictionary path). + - "cut_mode" ("search", "mix", "full", or "hmm"; default + "search"). + whitespace: + - no extra_params. + Filters: + lowercase: + - no extra_params. + ascii_folding: + - no extra_params. + stemmer: + - "stemmer_lang" (Snowball language/algorithm; default + "english"), for example {"stemmer_lang":"porter"} for ES + behaviour. Default is "". Examples: @@ -272,6 +293,11 @@ Examples: ... ) >>> print(params.tokenizer_name) jieba + >>> params = FtsIndexParam( + ... tokenizer_name="standard", + ... filters=["lowercase", "stemmer"], + ... extra_params='{"stemmer_lang":"porter"}', + ... ) )pbdoc"); fts_index_params .def(py::init, std::string>(), @@ -283,9 +309,29 @@ Constructs an FtsIndexParam instance. Args: tokenizer_name (str, optional): Tokenizer name. Defaults to "standard". - filters (list[str], optional): Token filter names. Supports "lowercase" and - "ascii_folding". Defaults to ["lowercase"]. - extra_params (str, optional): Extra tokenizer parameters. Defaults to "". + filters (list[str], optional): Token filter names. Supports "lowercase", + "ascii_folding", and "stemmer". Defaults to ["lowercase"]. + extra_params (str, optional): Extra tokenizer/filter parameters as an empty + string or JSON object string. Supported keys: + Tokenizers: + standard: + - "max_token_length" (positive integer). + jieba: + - "jieba_dict_dir". + - "user_dict_path". + - "cut_mode" ("search", "mix", "full", or "hmm"; default + "search"). + whitespace: + - no extra_params. + Filters: + lowercase: + - no extra_params. + ascii_folding: + - no extra_params. + stemmer: + - "stemmer_lang" (Snowball language/algorithm; default + "english"). + Defaults to "". )pbdoc") .def_property_readonly("tokenizer_name", &FtsIndexParams::tokenizer_name, "str: Name of the tokenizer.") diff --git a/src/db/CMakeLists.txt b/src/db/CMakeLists.txt index 361680d..69426d1 100644 --- a/src/db/CMakeLists.txt +++ b/src/db/CMakeLists.txt @@ -45,6 +45,7 @@ cc_library( libprotobuf FastPFOR cppjieba + snowball Arrow::arrow_static Arrow::parquet_static Arrow::arrow_compute diff --git a/src/db/index/CMakeLists.txt b/src/db/index/CMakeLists.txt index 08f472e..f383bd1 100644 --- a/src/db/index/CMakeLists.txt +++ b/src/db/index/CMakeLists.txt @@ -29,6 +29,7 @@ cc_library( Arrow::arrow_compute Arrow::arrow_dataset cppjieba + snowball FastPFOR utf8proc INCS . ${PROJECT_ROOT_DIR}/src diff --git a/src/db/index/column/fts_column/tokenizer/stemmer_token_filter.cc b/src/db/index/column/fts_column/tokenizer/stemmer_token_filter.cc new file mode 100644 index 0000000..a52ef32 --- /dev/null +++ b/src/db/index/column/fts_column/tokenizer/stemmer_token_filter.cc @@ -0,0 +1,80 @@ +// 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 "stemmer_token_filter.h" +#include +#include + +extern "C" { +#include +} + +namespace zvec::fts { + +struct ThreadLocalStemmerCache { + std::unordered_map stemmers; + + ~ThreadLocalStemmerCache() { + for (auto &[_, s] : stemmers) { + sb_stemmer_delete(s); + } + } + + struct sb_stemmer *get(const std::string &lang) { + auto it = stemmers.find(lang); + if (it != stemmers.end()) { + return it->second; + } + auto *s = sb_stemmer_new(lang.c_str(), nullptr); + if (s) { + stemmers[lang] = s; + } + return s; + } +}; + +bool StemmerTokenFilter::init(const ailego::JsonObject &config) { + std::string lang; + if (config.get("stemmer_lang", &lang) && !lang.empty()) { + language_ = lang; + } + auto *test_stemmer = sb_stemmer_new(language_.c_str(), nullptr); + if (!test_stemmer) { + LOG_ERROR("[StemmerTokenFilter] failed to create stemmer for language: %s", + language_.c_str()); + return false; + } + sb_stemmer_delete(test_stemmer); + return true; +} + +std::vector StemmerTokenFilter::filter(std::vector tokens) const { + static thread_local ThreadLocalStemmerCache tls_cache; + auto *stemmer = tls_cache.get(language_); + if (!stemmer) { + return tokens; + } + for (auto &token : tokens) { + const auto *result = sb_stemmer_stem( + stemmer, reinterpret_cast(token.text.data()), + static_cast(token.text.size())); + if (result) { + int len = sb_stemmer_length(stemmer); + token.text.assign(reinterpret_cast(result), len); + } + } + return tokens; +} + +} // namespace zvec::fts diff --git a/src/db/index/column/fts_column/tokenizer/stemmer_token_filter.h b/src/db/index/column/fts_column/tokenizer/stemmer_token_filter.h new file mode 100644 index 0000000..c17271f --- /dev/null +++ b/src/db/index/column/fts_column/tokenizer/stemmer_token_filter.h @@ -0,0 +1,46 @@ +// 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. + +#pragma once + +#include +#include +#include "token_filter.h" + +namespace zvec::fts { + +/*! Snowball stemmer token filter. + * Configure the language with extra_params key "stemmer_lang". When omitted, + * the language is "english". + */ +class StemmerTokenFilter : public TokenFilter { + public: + StemmerTokenFilter() = default; + ~StemmerTokenFilter() override = default; + + StemmerTokenFilter(const StemmerTokenFilter &) = delete; + StemmerTokenFilter &operator=(const StemmerTokenFilter &) = delete; + + bool init(const ailego::JsonObject &config) override; + std::vector filter(std::vector tokens) const override; + + const char *name() const override { + return "stemmer"; + } + + private: + std::string language_{"english"}; +}; + +} // namespace zvec::fts diff --git a/src/db/index/column/fts_column/tokenizer/token_filter.h b/src/db/index/column/fts_column/tokenizer/token_filter.h index fbacdfa..0cc939a 100644 --- a/src/db/index/column/fts_column/tokenizer/token_filter.h +++ b/src/db/index/column/fts_column/tokenizer/token_filter.h @@ -17,6 +17,7 @@ #include #include #include +#include #include "tokenizer.h" namespace zvec::fts { @@ -29,6 +30,15 @@ class TokenFilter { public: virtual ~TokenFilter() = default; + /*! Initialise the filter from a JSON configuration object. + * Must be called once before filter(). + * \param config JSON object containing filter-specific parameters. + * \return true on success, false on failure. + */ + virtual bool init(const ailego::JsonObject & /*config*/) { + return true; + } + /*! Filter/transform a list of tokens. * \param tokens input token list (may be modified in place) * \return processed token list diff --git a/src/db/index/column/fts_column/tokenizer/tokenizer_factory.cc b/src/db/index/column/fts_column/tokenizer/tokenizer_factory.cc index bfb30fd..41927c3 100644 --- a/src/db/index/column/fts_column/tokenizer/tokenizer_factory.cc +++ b/src/db/index/column/fts_column/tokenizer/tokenizer_factory.cc @@ -18,6 +18,7 @@ #include "ascii_folding_token_filter.h" #include "jieba_tokenizer.h" #include "standard_tokenizer.h" +#include "stemmer_token_filter.h" #include "whitespace_tokenizer.h" namespace zvec::fts { @@ -56,6 +57,11 @@ TokenizerPipelinePtr TokenizerFactory::create(const FtsIndexParams ¶ms) { filter_name.c_str()); return nullptr; } + if (!filter->init(extra_json)) { + LOG_ERROR("[TokenizerFactory] failed to init filter: %s", + filter_name.c_str()); + return nullptr; + } filters.push_back(std::move(filter)); } @@ -99,6 +105,10 @@ TokenFilterPtr TokenizerFactory::create_filter(const std::string &filter_name) { return std::make_shared(); } else if (filter_name == "ascii_folding") { return std::make_shared(); + } else if (filter_name == "stemmer") { + // The stemmer filter uses Snowball and defaults to "english" unless + // extra_params overrides stemmer_lang. + return std::make_shared(); } LOG_ERROR("[TokenizerFactory] unknown filter name: %s", filter_name.c_str()); return nullptr; diff --git a/src/db/index/column/fts_column/tokenizer/tokenizer_factory.h b/src/db/index/column/fts_column/tokenizer/tokenizer_factory.h index 646212c..22a3481 100644 --- a/src/db/index/column/fts_column/tokenizer/tokenizer_factory.h +++ b/src/db/index/column/fts_column/tokenizer/tokenizer_factory.h @@ -44,15 +44,14 @@ using TokenizerPipelinePtr = std::shared_ptr; /*! Tokenizer factory * Create TokenizerPipeline based on FtsIndexParams configuration. - * Supported tokenizers: standard, jieba, whitespace. - * Supported filters: lowercase, ascii_folding. */ class TokenizerFactory { public: /*! Create tokenizer pipeline from FtsIndexParams. * \param params FTS index parameters containing tokenizer_name, filters, * and extra_params (JSON string for tokenizer-specific - * configuration). + * configuration). The stemmer filter reads stemmer_lang from + * extra_params and uses Snowball English by default. * \return Tokenizer pipeline, returns nullptr on failure */ static TokenizerPipelinePtr create(const FtsIndexParams ¶ms); diff --git a/src/db/index/common/schema.cc b/src/db/index/common/schema.cc index 9471eee..06958cc 100644 --- a/src/db/index/common/schema.cc +++ b/src/db/index/common/schema.cc @@ -25,6 +25,8 @@ #include "db/common/constants.h" #include "db/common/typedef.h" #include "db/common/utils.h" +#include "db/index/column/fts_column/fts_types.h" +#include "db/index/column/fts_column/tokenizer/tokenizer_factory.h" #include "db/index/common/type_helper.h" namespace zvec { @@ -61,6 +63,28 @@ std::unordered_set support_dense_vector_index = { std::unordered_set support_sparse_vector_index = {IndexType::FLAT, IndexType::HNSW}; +static Status validate_fts_index_params(const FieldSchema &field) { + auto params = std::dynamic_pointer_cast(field.index_params()); + if (!params) { + return Status::InvalidArgument( + "schema validate failed: FTS index requires FtsIndexParams, but field[", + field.name(), "] has incompatible index params"); + } + + fts::FtsIndexParams internal_params; + internal_params.tokenizer_name = params->tokenizer_name(); + internal_params.filters = params->filters(); + internal_params.extra_params = params->extra_params(); + + auto pipeline = fts::TokenizerFactory::create(internal_params); + if (!pipeline) { + return Status::InvalidArgument( + "schema validate failed: invalid FTS index params for field[", + field.name(), "]"); + } + return Status::OK(); +} + Status FieldSchema::validate() const { if (data_type_ == DataType::UNDEFINED) { return Status::InvalidArgument("schema validate failed: field[", name_, @@ -261,6 +285,10 @@ Status FieldSchema::validate() const { "but field[", name_, "]'s data_type is ", DataTypeCodeBook::AsString(data_type_)); } + if (index_params_->type() == IndexType::FTS) { + auto s = validate_fts_index_params(*this); + CHECK_RETURN_STATUS(s); + } } } return Status::OK(); diff --git a/src/include/zvec/c_api.h b/src/include/zvec/c_api.h index 98828d3..baeb531 100644 --- a/src/include/zvec/c_api.h +++ b/src/include/zvec/c_api.h @@ -1089,12 +1089,31 @@ ZVEC_EXPORT zvec_error_code_t ZVEC_CALL zvec_index_params_set_invert_params( /** * @brief Set FTS index specific parameters * @param params Index parameters (must be FTS type) - * @param tokenizer_name Tokenizer name: "standard", "jieba", or "whitespace" - * (NULL keeps current value) - * @param filters Token filter names: "lowercase" and/or "ascii_folding" - * (NULL keeps current value) - * @param extra_params Additional tokenizer parameters (NULL keeps current - * value) + * @param tokenizer_name Tokenizer pipeline name (NULL keeps current value). + * Supported values are "standard", "jieba", and "whitespace". + * @param filters Token filter names (NULL keeps current value). Supported + * values are "lowercase", "ascii_folding", and "stemmer". + * @param extra_params Additional tokenizer/filter parameters (NULL keeps + * current value). Must be empty or a JSON object string. Supported keys by + * tokenizer/filter: + * Tokenizers: + * standard: + * - "max_token_length" (positive integer). + * jieba: + * - "jieba_dict_dir" (directory containing jieba.dict.utf8 and + * hmm_model.utf8). + * - "user_dict_path" (user dictionary path). + * - "cut_mode" ("search", "mix", "full", or "hmm"; default "search"). + * whitespace: + * - no extra_params. + * Filters: + * lowercase: + * - no extra_params. + * ascii_folding: + * - no extra_params. + * stemmer: + * - "stemmer_lang" (Snowball language/algorithm; default "english"), + * for example {"stemmer_lang":"porter"} for ES behaviour. * @return ZVEC_OK on success, error code on failure */ ZVEC_EXPORT zvec_error_code_t ZVEC_CALL zvec_index_params_set_fts_params( diff --git a/src/include/zvec/db/index_params.h b/src/include/zvec/db/index_params.h index 1a79e4b..31ec5b6 100644 --- a/src/include/zvec/db/index_params.h +++ b/src/include/zvec/db/index_params.h @@ -718,7 +718,28 @@ class VamanaIndexParams : public VectorIndexParams { /* * FTS (Full-Text Search) index params * Supported tokenizers: "standard", "jieba", "whitespace". - * Supported filters: "lowercase", "ascii_folding". + * Supported filters: "lowercase", "ascii_folding", "stemmer". + * + * extra_params must be either empty or a JSON object string. Supported keys are + * grouped by tokenizer/filter: + * Tokenizers: + * standard: + * - "max_token_length" (positive integer). + * jieba: + * - "jieba_dict_dir" (directory containing jieba.dict.utf8 and + * hmm_model.utf8). + * - "user_dict_path" (user dictionary path). + * - "cut_mode" ("search", "mix", "full", or "hmm"; default "search"). + * whitespace: + * - no extra_params. + * Filters: + * lowercase: + * - no extra_params. + * ascii_folding: + * - no extra_params. + * stemmer: + * - "stemmer_lang" (Snowball language/algorithm; default "english"), + * for example {"stemmer_lang":"porter"} for ES behaviour. * * Not copyable. Use shared_ptr for shared ownership. */ diff --git a/tests/db/index/column/fts_column/fts_column_indexer_test.cc b/tests/db/index/column/fts_column/fts_column_indexer_test.cc index 5bce2c5..e9b816e 100644 --- a/tests/db/index/column/fts_column/fts_column_indexer_test.cc +++ b/tests/db/index/column/fts_column/fts_column_indexer_test.cc @@ -1832,3 +1832,89 @@ TEST_F(FtsColumnIndexerTest, FilterPushdownNullFilterUnchanged) { EXPECT_FLOAT_EQ(baseline[i].score, with_null[i].score); } } + +// ============================================================ +// Stemmer token filter end-to-end tests +// ============================================================ + +static zvec::fts::TokenizerPipelinePtr make_stemmer_pipeline() { + zvec::fts::FtsIndexParams params; + params.tokenizer_name = "standard"; + params.filters = {"lowercase", "stemmer"}; + return zvec::fts::TokenizerFactory::create(params); +} + +class FtsStemmerIndexerTest : public FtsColumnIndexerTest { + protected: + std::unique_ptr make_stemmer_indexer( + const std::string &field_name = "content") { + auto fts_params = std::make_shared( + "standard", std::vector{"lowercase", "stemmer"}, ""); + auto field_meta = make_test_field_meta(field_name, fts_params); + auto indexer = std::make_unique(); + auto ret = indexer->open(field_meta, &db_, postings_cf_, positions_cf_, + term_freq_cf_, max_tf_cf_, doc_len_cf_, stat_cf_); + EXPECT_TRUE(ret.has_value()); + return indexer; + } +}; + +TEST_F(FtsStemmerIndexerTest, StemmedTermMatchesMorphologicalVariants) { + auto indexer = make_stemmer_indexer(); + EXPECT_TRUE(indexer->insert(0, "the cats are running quickly").has_value()); + EXPECT_TRUE(indexer->insert(1, "a dog runs slowly").has_value()); + EXPECT_TRUE(indexer->insert(2, "birds fly high").has_value()); + + auto pipeline = make_stemmer_pipeline(); + + // "running" stems to "run", matches doc 0 ("running") and doc 1 ("runs") + std::vector results; + EXPECT_TRUE(search_ok(*indexer, "running", 10, &results, pipeline)); + EXPECT_EQ(results.size(), 2u); + + // "cats" stems to "cat", matches only doc 0 + results.clear(); + EXPECT_TRUE(search_ok(*indexer, "cats", 10, &results, pipeline)); + EXPECT_EQ(results.size(), 1u); + EXPECT_EQ(results[0].doc_id, 0ull); +} + +TEST_F(FtsStemmerIndexerTest, QueryWithBaseFormMatchesVariants) { + auto indexer = make_stemmer_indexer(); + EXPECT_TRUE(indexer->insert(0, "connected connections").has_value()); + EXPECT_TRUE(indexer->insert(1, "connecting wires").has_value()); + EXPECT_TRUE(indexer->insert(2, "unrelated text").has_value()); + + auto pipeline = make_stemmer_pipeline(); + + // "connect" is already a stem, should match doc 0 and doc 1 + std::vector results; + EXPECT_TRUE(search_ok(*indexer, "connect", 10, &results, pipeline)); + EXPECT_EQ(results.size(), 2u); +} + +TEST_F(FtsStemmerIndexerTest, StemmerWithAndQuery) { + auto indexer = make_stemmer_indexer(); + EXPECT_TRUE(indexer->insert(0, "dogs running fast").has_value()); + EXPECT_TRUE(indexer->insert(1, "cats running slow").has_value()); + EXPECT_TRUE(indexer->insert(2, "dogs sleeping").has_value()); + + auto pipeline = make_stemmer_pipeline(); + + // "dogs AND running" -> stems to "dog AND run" -> doc 0 only + std::vector results; + EXPECT_TRUE(search_ok(*indexer, "dogs AND running", 10, &results, pipeline)); + EXPECT_EQ(results.size(), 1u); + EXPECT_EQ(results[0].doc_id, 0ull); +} + +TEST_F(FtsStemmerIndexerTest, StemmerNoMatchAfterStemming) { + auto indexer = make_stemmer_indexer(); + EXPECT_TRUE(indexer->insert(0, "hello world").has_value()); + + auto pipeline = make_stemmer_pipeline(); + + std::vector results; + EXPECT_TRUE(search_ok(*indexer, "nonexistent", 10, &results, pipeline)); + EXPECT_TRUE(results.empty()); +} diff --git a/tests/db/index/column/fts_column/stemmer_token_filter_test.cc b/tests/db/index/column/fts_column/stemmer_token_filter_test.cc new file mode 100644 index 0000000..b350c77 --- /dev/null +++ b/tests/db/index/column/fts_column/stemmer_token_filter_test.cc @@ -0,0 +1,167 @@ +// Copyright 2025-present the zvec project +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +#include +#include +#include +#include "db/index/column/fts_column/fts_types.h" +#include "db/index/column/fts_column/tokenizer/tokenizer_factory.h" + +using namespace zvec::fts; + +// ============================================================ +// Helpers +// ============================================================ + +static FtsIndexParams make_stemmer_params( + const std::string &lang = "", + const std::vector &filters = {"lowercase", "stemmer"}) { + FtsIndexParams params; + params.tokenizer_name = "standard"; + params.filters = filters; + if (!lang.empty()) { + params.extra_params = R"({"stemmer_lang":")" + lang + R"("})"; + } + return params; +} + +// ============================================================ +// Pipeline creation +// ============================================================ + +TEST(StemmerTokenFilterTest, CreatePipelineDefaultEnglish) { + auto pipeline = TokenizerFactory::create(make_stemmer_params()); + ASSERT_NE(pipeline, nullptr); +} + +TEST(StemmerTokenFilterTest, CreatePipelineExplicitLanguage) { + auto pipeline = TokenizerFactory::create(make_stemmer_params("german")); + ASSERT_NE(pipeline, nullptr); +} + +TEST(StemmerTokenFilterTest, CreatePipelineInvalidLanguageFails) { + auto pipeline = + TokenizerFactory::create(make_stemmer_params("nonexistent_lang")); + EXPECT_EQ(pipeline, nullptr); +} + +// ============================================================ +// English stemming +// ============================================================ + +TEST(StemmerTokenFilterTest, EnglishStemming) { + auto pipeline = TokenizerFactory::create(make_stemmer_params()); + ASSERT_NE(pipeline, nullptr); + + auto tokens = pipeline->process("running cats easily connection"); + ASSERT_EQ(tokens.size(), 4u); + EXPECT_EQ(tokens[0].text, "run"); + EXPECT_EQ(tokens[1].text, "cat"); + EXPECT_EQ(tokens[2].text, "easili"); + EXPECT_EQ(tokens[3].text, "connect"); +} + +TEST(StemmerTokenFilterTest, AlreadyStemmedWordsUnchanged) { + auto pipeline = TokenizerFactory::create(make_stemmer_params()); + ASSERT_NE(pipeline, nullptr); + + auto tokens = pipeline->process("run cat"); + ASSERT_EQ(tokens.size(), 2u); + EXPECT_EQ(tokens[0].text, "run"); + EXPECT_EQ(tokens[1].text, "cat"); +} + +TEST(StemmerTokenFilterTest, EmptyInput) { + auto pipeline = TokenizerFactory::create(make_stemmer_params()); + ASSERT_NE(pipeline, nullptr); + + auto tokens = pipeline->process(""); + EXPECT_TRUE(tokens.empty()); +} + +TEST(StemmerTokenFilterTest, PreservesOffsetAndPosition) { + auto pipeline = TokenizerFactory::create(make_stemmer_params()); + ASSERT_NE(pipeline, nullptr); + + auto tokens = pipeline->process("running dogs"); + ASSERT_EQ(tokens.size(), 2u); + EXPECT_EQ(tokens[0].position, 0u); + EXPECT_EQ(tokens[1].position, 1u); + EXPECT_EQ(tokens[0].offset, 0u); + EXPECT_EQ(tokens[1].offset, 8u); +} + +// ============================================================ +// Lowercase + stemmer chain +// ============================================================ + +TEST(StemmerTokenFilterTest, LowercaseThenStem) { + auto pipeline = TokenizerFactory::create(make_stemmer_params()); + ASSERT_NE(pipeline, nullptr); + + auto tokens = pipeline->process("Running Cats EASILY"); + ASSERT_EQ(tokens.size(), 3u); + EXPECT_EQ(tokens[0].text, "run"); + EXPECT_EQ(tokens[1].text, "cat"); + EXPECT_EQ(tokens[2].text, "easili"); +} + +// ============================================================ +// Stemmer-only (no lowercase) +// ============================================================ + +TEST(StemmerTokenFilterTest, StemmerOnlyNoLowercase) { + auto pipeline = + TokenizerFactory::create(make_stemmer_params("", {"stemmer"})); + ASSERT_NE(pipeline, nullptr); + + auto tokens = pipeline->process("running"); + ASSERT_EQ(tokens.size(), 1u); + EXPECT_EQ(tokens[0].text, "run"); +} + +// ============================================================ +// Non-English language +// ============================================================ + +TEST(StemmerTokenFilterTest, GermanStemming) { + auto pipeline = TokenizerFactory::create(make_stemmer_params("german")); + ASSERT_NE(pipeline, nullptr); + + auto tokens = pipeline->process("laufen"); + ASSERT_EQ(tokens.size(), 1u); + EXPECT_EQ(tokens[0].text, "lauf"); +} + +// ============================================================ +// ISO code as language +// ============================================================ + +TEST(StemmerTokenFilterTest, LanguageByISOCode) { + auto pipeline = TokenizerFactory::create(make_stemmer_params("en")); + ASSERT_NE(pipeline, nullptr); + + auto tokens = pipeline->process("running"); + ASSERT_EQ(tokens.size(), 1u); + EXPECT_EQ(tokens[0].text, "run"); +} + +TEST(StemmerTokenFilterTest, PorterAlgorithm) { + auto pipeline = TokenizerFactory::create(make_stemmer_params("porter")); + ASSERT_NE(pipeline, nullptr); + + auto tokens = pipeline->process("running"); + ASSERT_EQ(tokens.size(), 1u); + EXPECT_EQ(tokens[0].text, "run"); +} diff --git a/tests/db/index/common/schema_test.cc b/tests/db/index/common/schema_test.cc index e45f221..6984b60 100644 --- a/tests/db/index/common/schema_test.cc +++ b/tests/db/index/common/schema_test.cc @@ -343,6 +343,27 @@ TEST(FieldSchemaTest, Validate) { EXPECT_TRUE(status.ok()); } + { + auto fts_params = std::make_shared( + "standard", std::vector{"lowercase", "stemmer"}, + R"({"stemmer_lang":"english"})"); + FieldSchema field("fts_field", DataType::STRING, false, fts_params); + auto status = field.validate(); + EXPECT_TRUE(status.ok()); + } + + { + auto fts_params = std::make_shared( + "standard", std::vector{"lowercase", "stemmer"}, + R"({"stemmer_lang":"nonexistent_lang"})"); + FieldSchema field("fts_field", DataType::STRING, false, fts_params); + auto status = field.validate(); + EXPECT_FALSE(status.ok()); + EXPECT_EQ(status.code(), StatusCode::INVALID_ARGUMENT); + EXPECT_NE(status.message().find("invalid FTS index params"), + std::string::npos); + } + { FieldSchema field("simple_field", DataType::STRING); auto status = field.validate(); diff --git a/thirdparty/CMakeLists.txt b/thirdparty/CMakeLists.txt index a5384c4..b719043 100644 --- a/thirdparty/CMakeLists.txt +++ b/thirdparty/CMakeLists.txt @@ -31,3 +31,4 @@ add_subdirectory(FastPFOR FastPFOR EXCLUDE_FROM_ALL) add_subdirectory(limonp limonp EXCLUDE_FROM_ALL) add_subdirectory(cppjieba cppjieba EXCLUDE_FROM_ALL) add_subdirectory(utf8proc utf8proc EXCLUDE_FROM_ALL) +add_subdirectory(snowball snowball EXCLUDE_FROM_ALL) diff --git a/thirdparty/snowball/CMakeLists.txt b/thirdparty/snowball/CMakeLists.txt new file mode 100644 index 0000000..92d6756 --- /dev/null +++ b/thirdparty/snowball/CMakeLists.txt @@ -0,0 +1,100 @@ +include(ExternalProject) + +set(SNOWBALL_SOURCE_DIR "${CMAKE_CURRENT_SOURCE_DIR}/snowball-3.1.1") +set(SNOWBALL_BUILD_DIR "${CMAKE_CURRENT_BINARY_DIR}/snowball-codegen") +set(SNOWBALL_HOST_CC "" CACHE STRING + "Optional host C compiler for building the Snowball code generator") +find_program(_SNOWBALL_MAKE NAMES make gmake REQUIRED) + +# --------------------------------------------------------------------------- +# Parse modules.txt → UTF-8 algorithm list +# --------------------------------------------------------------------------- +set(_snowball_gen_srcs) +set(_snowball_gen_hdrs) +set(_snowball_make_targets) +file(STRINGS "${SNOWBALL_SOURCE_DIR}/libstemmer/modules.txt" _lines) +foreach(_line IN LISTS _lines) + if(_line MATCHES "^#" OR _line MATCHES "^[ \t]*$") + continue() + endif() + if(_line MATCHES "^([a-z_]+)[ \t]+([A-Z_0-9,]+)") + set(_alg "${CMAKE_MATCH_1}") + list(APPEND _snowball_gen_srcs + "${SNOWBALL_BUILD_DIR}/src_c/stem_UTF_8_${_alg}.c") + list(APPEND _snowball_gen_hdrs + "${SNOWBALL_BUILD_DIR}/src_c/stem_UTF_8_${_alg}.h") + list(APPEND _snowball_make_targets + "src_c/stem_UTF_8_${_alg}.c") + endif() +endforeach() + +set(_snowball_make_args "CFLAGS=-O2") +if(NOT SNOWBALL_HOST_CC STREQUAL "") + list(APPEND _snowball_make_args "CC=${SNOWBALL_HOST_CC}") +endif() + +# --------------------------------------------------------------------------- +# Phase 1 (host): build snowball compiler & generate UTF-8 sources only +# --------------------------------------------------------------------------- +# Copy source tree into the build directory so the original stays clean. +# Request only the UTF-8 stemmer sources, the utf8 libstemmer entry point, +# and the utf8 modules header — no ISO-8859/KOI8 stemmers, no host .a. +# Each src_c/stem_UTF_8_*.c target implicitly builds the snowball compiler +# (host executable) as a dependency. +# By default make uses system `cc`; set SNOWBALL_HOST_CC to override when +# the environment CC points to a cross-compiler. +ExternalProject_Add(snowball_codegen + DOWNLOAD_COMMAND ${CMAKE_COMMAND} -E copy_directory + ${SNOWBALL_SOURCE_DIR} ${SNOWBALL_BUILD_DIR} + SOURCE_DIR ${SNOWBALL_BUILD_DIR} + CONFIGURE_COMMAND "" + BUILD_COMMAND ${_SNOWBALL_MAKE} + libstemmer/libstemmer_utf8.c + libstemmer/modules_utf8.h + ${_snowball_make_targets} + ${_snowball_make_args} + BUILD_IN_SOURCE TRUE + INSTALL_COMMAND "" + BUILD_BYPRODUCTS + ${SNOWBALL_BUILD_DIR}/runtime/api.c + ${SNOWBALL_BUILD_DIR}/runtime/utilities.c + ${SNOWBALL_BUILD_DIR}/libstemmer/libstemmer_utf8.c + ${SNOWBALL_BUILD_DIR}/libstemmer/modules_utf8.h + ${_snowball_gen_srcs} + ${_snowball_gen_hdrs} +) + +# --------------------------------------------------------------------------- +# Phase 2 (target): compile generated sources with the project toolchain +# --------------------------------------------------------------------------- +set(_snowball_target_srcs + ${SNOWBALL_BUILD_DIR}/runtime/api.c + ${SNOWBALL_BUILD_DIR}/runtime/utilities.c + ${SNOWBALL_BUILD_DIR}/libstemmer/libstemmer_utf8.c + ${_snowball_gen_srcs} +) + +set_source_files_properties(${_snowball_target_srcs} + PROPERTIES GENERATED TRUE) + +if(NOT TARGET snowball) + add_library(snowball STATIC ${_snowball_target_srcs}) + add_dependencies(snowball snowball_codegen) + # Public include points to the SOURCE directory — libstemmer.h exists at + # configure time and does not depend on the codegen step. + target_include_directories(snowball SYSTEM PUBLIC + ${SNOWBALL_SOURCE_DIR}/include + ) + # Private includes for generated headers (modules_utf8.h, stem_*.h). + target_include_directories(snowball PRIVATE + ${SNOWBALL_BUILD_DIR} + ${SNOWBALL_BUILD_DIR}/libstemmer + ${SNOWBALL_BUILD_DIR}/src_c + ) + set_target_properties(snowball PROPERTIES + POSITION_INDEPENDENT_CODE ON + C_STANDARD 99 + ) +endif() + +set(snowball_FOUND TRUE PARENT_SCOPE) diff --git a/thirdparty/snowball/snowball-3.1.1 b/thirdparty/snowball/snowball-3.1.1 new file mode 160000 index 0000000..cd195b5 --- /dev/null +++ b/thirdparty/snowball/snowball-3.1.1 @@ -0,0 +1 @@ +Subproject commit cd195b51e948a902a4312f023f4a14392516a543