bk-tree 0.1.4
Header-only Burkhard-Keller tree library
Loading...
Searching...
No Matches
bktree.hpp
1//
2// bk-tree Header-only Burkhard-Keller tree library
3// Copyright (C) 2020-2026 John Law
4//
5// This file is part of bk-tree.
6//
7// bk-tree is free software: you can redistribute it and/or modify
8// it under the terms of the GNU General Public License as published by
9// the Free Software Foundation, either version 3 of the License, or
10// (at your option) any later version.
11//
12// bk-tree is distributed in the hope that it will be useful,
13// but WITHOUT ANY WARRANTY; without even the implied warranty of
14// MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the
15// GNU General Public License for more details.
16//
17// You should have received a copy of the GNU General Public License
18// along with bk-tree. If not, see <https://www.gnu.org/licenses/>.
19//
20
21#pragma once
22
23#ifndef BK_MATRIX_INITIAL_SIZE
24#define BK_MATRIX_INITIAL_SIZE 0
25#endif
26#ifndef BK_LCS_MATRIX_INITIAL_SIZE
27#define BK_LCS_MATRIX_INITIAL_SIZE BK_MATRIX_INITIAL_SIZE
28#endif
29#ifndef BK_ED_MATRIX_INITIAL_SIZE
30#define BK_ED_MATRIX_INITIAL_SIZE BK_MATRIX_INITIAL_SIZE
31#endif
32#ifndef BK_LEE_ALPHABET_SIZE
33#define BK_LEE_ALPHABET_SIZE 26
34#endif
35#ifndef BK_TREE_INITIAL_SIZE
36#define BK_TREE_INITIAL_SIZE 0
37#endif
38#include <algorithm>
39#include <cstddef>
40#include <iterator>
41#include <limits>
42#include <map>
43#include <memory>
44#include <queue>
45#include <string>
46#include <type_traits>
47#include <utility>
48#include <vector>
49
53namespace bk_tree {
54
55using integer_type = std::uint64_t;
56
60namespace metrics {
61
65template <typename Metric>
66class Distance {
67public:
68 integer_type operator()(std::string_view s, std::string_view t) const {
69 return (static_cast<Metric const *>(this))->compute_distance(s, t);
70 }
71};
72
80class LengthDistance final : public Distance<LengthDistance> {
81public:
82 explicit LengthDistance() {};
83 integer_type compute_distance(std::string_view s, std::string_view t) const noexcept {
84 return s.length() > t.length() ? s.length() - t.length() : t.length() - s.length();
85 }
86};
87
93class IdentityDistance final : public Distance<IdentityDistance> {
94public:
95 explicit IdentityDistance() {};
96 integer_type compute_distance(std::string_view, std::string_view) const noexcept {
97 return integer_type{1};
98 }
99};
100
108class LeeDistance final : public Distance<LeeDistance> {
109 integer_type m_alphabet_size;
110
111public:
112 explicit LeeDistance(integer_type alphabet_size = BK_LEE_ALPHABET_SIZE)
113 : m_alphabet_size(alphabet_size) {};
114 integer_type compute_distance(std::string_view s, std::string_view t) const noexcept {
115 const integer_type M = s.length(), N = t.length();
116 if (M != N) {
117 return std::numeric_limits<integer_type>::max();
118 }
119 const integer_type m_comparison_size = M;
120 integer_type counter = 0, diff;
121 for (integer_type i = 0; i < m_comparison_size; ++i) {
122 diff = std::abs(s[i] - t[i]);
123 counter += std::min(diff, m_alphabet_size - diff);
124 }
125 return counter;
126 }
127};
128
144class LCSubseqDistance final : public Distance<LCSubseqDistance> {
145 mutable std::vector<integer_type> m_current, m_previous;
146
147public:
148 explicit LCSubseqDistance(size_t initial_size = BK_LCS_MATRIX_INITIAL_SIZE)
149 : m_current(initial_size), m_previous(initial_size) {};
150 integer_type compute_distance(std::string_view s, std::string_view t) const noexcept {
151 const integer_type M = s.length(), N = t.length();
152 if (M == 0 || N == 0) {
153 return M + N;
154 }
155 if (m_current.size() <= N || m_previous.size() <= N) {
156 m_current.resize(N + 1);
157 m_previous.resize(N + 1);
158 }
159 std::fill(m_previous.begin(), m_previous.end(), 0);
160 for (integer_type i = 1; i <= M; ++i) {
161 for (integer_type j = 1; j <= N; ++j) {
162 if (s[i - 1] == t[j - 1]) {
163 m_current[j] = m_previous[j - 1] + 1;
164 } else {
165 m_current[j] = std::max(m_previous[j], m_current[j - 1]);
166 }
167 }
168 m_previous = m_current;
169 }
170 return M + N - 2 * m_previous[N];
171 }
172};
173
181class HammingDistance final : public Distance<HammingDistance> {
182public:
183 explicit HammingDistance() = default;
184 integer_type compute_distance(std::string_view s, std::string_view t) const noexcept {
185 const integer_type M = s.length(), N = t.length();
186 if (M != N) {
187 return std::numeric_limits<integer_type>::max();
188 }
189 const integer_type m_comparison_size = M;
190 integer_type counter = 0;
191 for (integer_type i = 0; i < m_comparison_size; ++i) {
192 counter += (s[i] != t[i]);
193 }
194 return counter;
195 }
196};
197
222class EditDistance final : public Distance<EditDistance> {
223 mutable std::vector<std::vector<integer_type>> m_matrix;
224
225public:
226 explicit EditDistance(size_t initial_size = BK_ED_MATRIX_INITIAL_SIZE)
227 : m_matrix(initial_size, std::vector<integer_type>(initial_size)) {};
228 integer_type compute_distance(std::string_view s, std::string_view t) const noexcept {
229 const integer_type M = s.length(), N = t.length();
230 if (M == 0 || N == 0) {
231 return N + M;
232 }
233 if (m_matrix.size() <= M || m_matrix[0].size() <= N) {
234 std::vector<std::vector<integer_type>> a_matrix(M + 1,
235 std::vector<integer_type>(N + 1));
236 m_matrix.swap(a_matrix);
237 }
238 for (integer_type i = 1; i <= M; ++i) {
239 m_matrix[i][0] = i;
240 }
241 for (integer_type j = 1; j <= N; ++j) {
242 m_matrix[0][j] = j;
243 }
244 for (integer_type j = 1; j <= N; ++j) {
245 for (integer_type i = 1; i <= M; ++i) {
246 m_matrix[i][j] = std::min(
247 {m_matrix[i][j - 1] + 1 /*Insertion*/, m_matrix[i - 1][j] + 1 /*Deletion*/,
248 m_matrix[i - 1][j - 1] + (s[i - 1] != t[j - 1]) /*Substitution*/});
249 }
250 }
251 return m_matrix[M][N];
252 }
253};
254
260class DamerauLevenshteinDistance final : public Distance<DamerauLevenshteinDistance> {
261 mutable std::vector<std::vector<integer_type>> m_matrix;
262
263public:
264 explicit DamerauLevenshteinDistance(size_t initial_size = BK_MATRIX_INITIAL_SIZE)
265 : m_matrix(initial_size, std::vector<integer_type>(initial_size)) {};
266 integer_type compute_distance(std::string_view s, std::string_view t) const noexcept {
267 const integer_type M = s.length(), N = t.length();
268 if (M == 0 || N == 0) {
269 return N + M;
270 }
271 if (m_matrix.size() <= M || m_matrix[0].size() <= N) {
272 std::vector<std::vector<integer_type>> a_matrix(M + 1,
273 std::vector<integer_type>(N + 1));
274 m_matrix.swap(a_matrix);
275 }
276 for (integer_type i = 1; i <= M; ++i) {
277 m_matrix[i][0] = i;
278 }
279 for (integer_type j = 1; j <= N; ++j) {
280 m_matrix[0][j] = j;
281 }
282 for (integer_type j = 1; j <= N; ++j) {
283 for (integer_type i = 1; i <= M; ++i) {
284 m_matrix[i][j] = std::min(
285 {m_matrix[i][j - 1] + 1 /*Insertion*/, m_matrix[i - 1][j] + 1 /*Deletion*/,
286 m_matrix[i - 1][j - 1] + (s[i - 1] == t[j - 1] ? 0 : 1) /*Substitution*/});
287 if (i > 1 && j > 1 && s[i - 1] == t[j - 2] && s[i - 2] == t[j - 1]) {
288 m_matrix[i][j] =
289 std::min(m_matrix[i][j], m_matrix[i - 2][j - 2] + 1 /*Transposition*/);
290 }
291 }
292 }
293 return m_matrix[M][N];
294 }
295};
296
297} // namespace metrics
298
299namespace helpers {
300std::false_type is_metric_impl(...);
301
302template <typename Metric>
303std::true_type is_metric_impl(const volatile metrics::Distance<Metric> &);
304
305template <typename Metric>
306using is_metric = decltype(is_metric_impl(std::declval<Metric &>()));
307} // namespace helpers
308
309template <typename Metric>
310class BKTree;
311template <typename Metric>
312class BKTreeNode;
313
314using ResultEntry = std::pair<std::string, integer_type>;
315using ResultList = std::vector<ResultEntry>;
316
317template <typename Metric>
318class BKTreeNode {
319 friend class BKTree<Metric>;
320 using metric_type = Metric;
321 using node_type = BKTreeNode<metric_type>;
322
323 BKTreeNode(std::string_view value) : m_word(value) {}
324 bool _insert(std::string_view value, const metric_type &distance);
325 bool _erase(std::string_view value, const metric_type &distance);
326 void _find(ResultList &output, std::string_view value, integer_type limit,
327 const metric_type &metric) const;
328 ResultList _find_wrapper(std::string_view value, integer_type limit,
329 const metric_type &metric) const;
330
331 std::map<integer_type, std::unique_ptr<node_type>> m_children;
332 std::string m_word;
333
334 friend std::ostream &operator<<(std::ostream &oss, const BKTreeNode &node) {
335 oss << node.m_word;
336 return oss;
337 }
338
339public:
340 std::string_view word() const noexcept { return m_word; }
341};
342
346template <typename Metric>
347class BKTree {
348 static_assert(helpers::is_metric<Metric>::value, "Metric must be of type Distance");
349
350 using metric_type = Metric;
351 using node_type = typename BKTreeNode<metric_type>::node_type;
352
353public:
357 class Iterator {
358 public:
359 using iterator_category = std::forward_iterator_tag;
360 using difference_type = std::ptrdiff_t;
361 using value_type = std::unique_ptr<node_type>;
362 using pointer = std::unique_ptr<node_type> *;
363 using reference = std::unique_ptr<node_type> &;
364
365 public:
366 Iterator() = default;
367 Iterator(pointer ptr) : m_pointer(ptr) {}
368
369 pointer operator->() { return m_pointer; }
370
371 reference operator*() const { return *m_pointer; }
372
373 Iterator &operator++() {
374 if (m_pointer == nullptr) {
375 throw std::out_of_range("No more tree node");
376 }
377 for (auto &[_, child] : (*m_pointer)->m_children) {
378 m_queue.push(&child);
379 }
380 if (m_queue.empty()) {
381 m_pointer = nullptr;
382 } else {
383 m_pointer = m_queue.front();
384 m_queue.pop();
385 }
386 return *this;
387 }
388
389 Iterator operator++(int) {
390 Iterator tmp{*this};
391 ++(*this);
392 return tmp;
393 }
394
395 friend bool operator==(const Iterator &a, const Iterator &b) {
396 return a.m_pointer == b.m_pointer;
397 };
398
399 friend bool operator!=(const Iterator &a, const Iterator &b) {
400 return a.m_pointer != b.m_pointer;
401 };
402
403 private:
404 pointer m_pointer;
405 std::queue<pointer> m_queue;
406 };
407
408public:
409 BKTree(const metric_type &distance = Metric())
410 : m_root(nullptr), m_metric(distance), m_tree_size(BK_TREE_INITIAL_SIZE) {}
411
412 BKTree(std::initializer_list<std::string_view> list)
413 : m_root(nullptr), m_metric(Metric()), m_tree_size(BK_TREE_INITIAL_SIZE) {
414 for (auto &str : list) {
415 insert(str);
416 }
417 }
418
419 BKTree(const BKTree &other) : BKTree(other.m_metric) {
420 if (other.m_root == nullptr) {
421 return;
422 }
423 std::queue<std::unique_ptr<node_type> const *> bq;
424 bq.push(&(other.m_root));
425 while (!bq.empty()) {
426 auto *nptr = bq.front();
427 bq.pop();
428 this->insert((*nptr)->m_word);
429 for (auto &[_, child_node] : (*nptr)->m_children) {
430 bq.push(&child_node);
431 }
432 }
433 }
434
435 BKTree(BKTree &&other) noexcept
436 : m_root(std::exchange(other.m_root, nullptr)), m_tree_size(other.m_tree_size) {}
437
438 BKTree &operator=(const BKTree &other) {
439 if (this == &other) {
440 return *this;
441 }
442 BKTree temp(other);
443 std::swap(m_root, temp.m_root);
444 std::swap(m_tree_size, temp.m_tree_size);
445 return *this;
446 }
447
448 BKTree &operator=(BKTree &&other) noexcept {
449 std::swap(m_root, other.m_root);
450 std::swap(m_tree_size, other.m_tree_size);
451 return *this;
452 }
453
454 ~BKTree() = default;
455
456 bool insert(std::string_view value);
457 bool erase(std::string_view value);
458 size_t size() const noexcept { return m_tree_size; }
459 bool empty() const noexcept { return m_tree_size == 0; }
460 [[nodiscard]] ResultList find(std::string_view value, integer_type limit) const;
461 [[nodiscard]] ResultList find(std::string_view value, int limit) const;
462 template <typename Integer>
463 requires std::is_integral_v<Integer>
464 [[nodiscard]] ResultList find(std::string_view value, Integer limit) const {
465 if constexpr (std::is_signed_v<Integer>) {
466 if (limit < 0) {
467 return ResultList{};
468 }
469 }
470 return find(value, static_cast<integer_type>(limit));
471 }
472
473 Iterator begin() { return Iterator(&m_root); }
474 Iterator end() { return Iterator(); }
475
476private:
477 std::unique_ptr<node_type> m_root;
478 const metric_type m_metric;
479 size_t m_tree_size;
480};
481
482template <typename Metric>
483bool BKTreeNode<Metric>::_insert(std::string_view value,
484 const metric_type &distance_metric) {
485 const integer_type distance_between = distance_metric(value, m_word);
486 if (distance_between == std::numeric_limits<integer_type>::max()) {
487 return false;
488 }
489 auto it = m_children.find(distance_between);
490 if (it == m_children.end()) {
491 m_children.emplace(std::make_pair(
492 distance_between, std::unique_ptr<node_type>(new node_type(value))));
493 return true;
494 }
495 return it->second->_insert(value, distance_metric);
496}
497
498template <typename Metric>
499bool BKTreeNode<Metric>::_erase(std::string_view value,
500 const metric_type &distance_metric) {
501 bool erased = false;
502 const integer_type distance_between = distance_metric(value, m_word);
503 auto it = m_children.find(distance_between);
504 if (it != m_children.end()) {
505 if (it->second->m_word == value) {
506 auto node = std::move(it->second);
507 m_children.erase(it);
508 std::queue<std::unique_ptr<node_type> const *> bq;
509 for (auto const &[_, child_node] : node->m_children) {
510 bq.push(&child_node);
511 }
512 while (!bq.empty()) {
513 auto *node = bq.front();
514 bq.pop();
515 for (auto const &[_, child_node] : (*node)->m_children) {
516 bq.push(&child_node);
517 }
518 _insert((*node)->m_word, distance_metric);
519 }
520 erased = true;
521 } else {
522 erased = it->second->_erase(value, distance_metric);
523 }
524 } else {
525 for (auto const &[_, child] : m_children) {
526 if (child->_erase(value, distance_metric)) {
527 return true;
528 }
529 }
530 }
531 return erased;
532}
533
534template <typename Metric>
535void BKTreeNode<Metric>::_find(ResultList &output, std::string_view value,
536 integer_type limit, const metric_type &metric) const {
537 const integer_type distance = metric(value, m_word);
538 if (distance != std::numeric_limits<integer_type>::max() && distance <= limit) {
539 output.push_back({m_word, distance});
540 }
541 for (auto const &[dist, node] : m_children) {
542 const integer_type difference = dist > distance ? dist - distance : distance - dist;
543 if (difference <= limit) {
544 node->_find(output, value, limit, metric);
545 }
546 }
547}
548
549template <typename Metric>
550ResultList BKTreeNode<Metric>::_find_wrapper(std::string_view value, integer_type limit,
551 const metric_type &metric) const {
552 ResultList output;
553 _find(output, value, limit, metric);
554 return output;
555}
556
557template <typename Metric>
558bool BKTree<Metric>::insert(std::string_view value) {
559 bool inserted = false;
560 if (m_root == nullptr) {
561 m_root = std::unique_ptr<node_type>(new node_type(value));
562 ++m_tree_size;
563 inserted = true;
564 } else if (m_root->_insert(value, m_metric)) {
565 ++m_tree_size;
566 inserted = true;
567 }
568 return inserted;
569}
570
571template <typename Metric>
572bool BKTree<Metric>::erase(std::string_view value) {
573 bool erased = false;
574 if (m_root == nullptr) {
575 erased = false;
576 } else if (m_root->m_word == value) {
577 if (m_tree_size > 1) {
578 auto &replacement_node = m_root->m_children.begin()->second;
579 std::queue<std::unique_ptr<node_type> const *> bq;
580 for (bool first = true; auto const &[_, node] : m_root->m_children) {
581 if (first) {
582 first = false;
583 continue;
584 }
585 bq.push(&node);
586 }
587 while (!bq.empty()) {
588 auto node = bq.front();
589 bq.pop();
590 for (auto const &[_, child] : (*node)->m_children) {
591 bq.push(&child);
592 }
593 replacement_node->_insert((*node)->m_word, m_metric);
594 }
595 m_root = std::move(replacement_node);
596 } else {
597 m_root.reset(nullptr);
598 }
599 --m_tree_size;
600 erased = true;
601 } else if (m_root->_erase(value, m_metric)) {
602 --m_tree_size;
603 erased = true;
604 }
605 return erased;
606}
607
608template <typename Metric>
609ResultList BKTree<Metric>::find(std::string_view value, integer_type limit) const {
610 if (m_root == nullptr) {
611 return ResultList{};
612 }
613 return m_root->_find_wrapper(value, limit, m_metric);
614}
615
616template <typename Metric>
617ResultList BKTree<Metric>::find(std::string_view value, int limit) const {
618 if (limit < 0) {
619 return ResultList{};
620 }
621 return find(value, static_cast<integer_type>(limit));
622}
623
624} // namespace bk_tree
BK-tree class iterator.
Definition bktree.hpp:357
BK-tree template class.
Definition bktree.hpp:347
Damerau–Levenshtein metric.
Definition bktree.hpp:260
Metric interface for string distances.
Definition bktree.hpp:66
Edit distance metric.
Definition bktree.hpp:222
Hamming distance metric.
Definition bktree.hpp:181
Identity metric.
Definition bktree.hpp:93
Longest Common Subsequence distance metric.
Definition bktree.hpp:144
Lee distance metric.
Definition bktree.hpp:108
Length metric.
Definition bktree.hpp:80
BK-tree namespace.
Definition bktree.hpp:53