前言
上一篇文章实现了哈希表HashTable,但它只能存pair<const K, V>。而标准库的unordered_set只存键,unordered_map存键值对 ——它们内部用的是同一套哈希表。
怎么做到"一份哈希表,封装出两个容器"?这就是本文的主题。
关键技巧是KeyOfValue仿函数:让哈希表只知道"怎么从存储的元素里取出键",而不关心元素到底是K还是pair<K,V>。
一、先说结论:设计骨架
┌─────────────────┐ │ HashTable │ ← 底层,存任意类型 T │ <K, T, │ 只需要知道: │ KeyOfValue, │ ① 键 K 是什么类型 │ Hash> │ ② 怎么从 T 里取出 K └────────┬────────┘ │ ┌─────────────┴─────────────┐ │ │ ┌───────▼────────┐ ┌─────────▼──────────┐ │ myunordered_set│ │ myunordered_map │ │ 存储元素: K │ │ 存储元素: pair<K,V>│ │ KeyOfValue: │ │ KeyOfValue: │ │ 返回自身 │ │ 返回 .first │ └────────────────┘ └────────────────────┘三个模板参数的分工:
| 参数 | 含义 | set 中 | map 中 |
|---|---|---|---|
K | 键的类型 | K | K |
T | 实际存储的元素类型 | K | pair<const K, V> |
KeyOfValue | 从T取出K的仿函数 | 返回自身 | 返回.first |
二、KeyOfValue 仿函数
这是整个设计的核心,就两个小类:
// set 用:元素本身就是键 template <typename K> struct Identity { const K& operator()(const K& key) const { return key; } }; // map 用:元素是 pair,键在 first template <typename K, typename V> struct SelectFirst { const K& operator()(const std::pair<const K, V>& kv) const { return kv.first; } };就这么简单。但正是这两行,让哈希表不用为 set 和 map 各写一份。
三、改造哈希表
在上一篇的基础上,把硬编码的kv.first换成KeyOfValue调用,并补上迭代器。
#include <vector> #include <list> #include <utility> #include <functional> #include <iostream> template <typename K, typename T, typename KeyOfValue, typename Hash = std::hash<K>> class HashTable { public: using value_type = T; using Bucket = std::list<T>; private: std::vector<Bucket> buckets_; size_t size_ = 0; float max_load_ = 1.0f; Hash hash_{}; KeyOfValue key_of_{}; // ★ 取键的仿函数 size_t bucketIndex(const K& key) const { return hash_(key) % buckets_.size(); } const K& getKey(const T& val) const { return key_of_(val); // ★ 统一入口,不再写 .first } public: explicit HashTable(size_t bucket_count = 16, float max_load = 1.0f) : buckets_(bucket_count), max_load_(max_load) {} /* ================= 迭代器 ================= */ template <bool IsConst> class Iterator { using BucketPtr = std::conditional_t<IsConst, const Bucket*, Bucket*>; using ListIter = std::conditional_t<IsConst, typename Bucket::const_iterator, typename Bucket::iterator>; BucketPtr table_ = nullptr; // 指向桶数组 size_t idx_ = 0; // 当前桶下标 ListIter it_{}; // 桶内位置 size_t bucket_count_ = 0; void skipEmpty() { while (table_ && idx_ < bucket_count_ && it_ == table_[idx_].end()) { ++idx_; if (idx_ < bucket_count_) it_ = table_[idx_].begin(); } } public: using iterator_category = std::forward_iterator_tag; using value_type = T; using difference_type = std::ptrdiff_t; using pointer = std::conditional_t<IsConst, const T*, T*>; using reference = std::conditional_t<IsConst, const T&, T&>; Iterator() = default; Iterator(BucketPtr t, size_t idx, size_t n) : table_(t), idx_(idx), bucket_count_(n) { if (table_ && idx_ < bucket_count_) it_ = table_[idx_].begin(); skipEmpty(); } Iterator(BucketPtr t, size_t idx, ListIter it, size_t n) : table_(t), idx_(idx), it_(it), bucket_count_(n) {} // 允许 iterator → const_iterator 转换 template <bool OtherConst, typename = std::enable_if_t<IsConst && !OtherConst>> Iterator(const Iterator<OtherConst>& o) : table_(o.table_), idx_(o.idx_), it_(o.it_), bucket_count_(o.bucket_count_) {} reference operator*() const { return *it_; } pointer operator->() const { return &(*it_); } Iterator& operator++() { ++it_; skipEmpty(); return *this; } Iterator operator++(int) { Iterator t = *this; ++(*this); return t; } template <bool C> bool operator==(const Iterator<C>& o) const { return it_ == o.it_; } template <bool C> bool operator!=(const Iterator<C>& o) const { return !(*this == o); } template <bool> friend class Iterator; friend class HashTable; }; using iterator = Iterator<false>; using const_iterator = Iterator<true>; iterator begin() { return iterator(buckets_.data(), 0, buckets_.size()); } iterator end() { return iterator(buckets_.data(), buckets_.size(), buckets_.size()); } const_iterator begin() const { return const_iterator(buckets_.data(), 0, buckets_.size()); } const_iterator end() const { return const_iterator(buckets_.data(), buckets_.size(), buckets_.size()); } /* ================= 基本操作 ================= */ size_t size() const { return size_; } bool empty() const { return size_ == 0; } std::pair<iterator, bool> insert(const T& val) { if (size_ + 1 > static_cast<size_t>(max_load_ * buckets_.size())) { rehash(buckets_.size() * 2); } const K& key = getKey(val); Bucket& b = buckets_[bucketIndex(key)]; for (auto it = b.begin(); it != b.end(); ++it) { if (getKey(*it) == key) { return {iterator(buckets_.data(), bucketIndex(key), it, buckets_.size()), false}; } } b.push_back(val); ++size_; auto it = b.end(); --it; return {iterator(buckets_.data(), bucketIndex(key), it, buckets_.size()), true}; } iterator find(const K& key) { size_t idx = bucketIndex(key); auto& b = buckets_[idx]; for (auto it = b.begin(); it != b.end(); ++it) { if (getKey(*it) == key) { return iterator(buckets_.data(), idx, it, buckets_.size()); } } return end(); } size_t erase(const K& key) { size_t idx = bucketIndex(key); auto& b = buckets_[idx]; for (auto it = b.begin(); it != b.end(); ++it) { if (getKey(*it) == key) { b.erase(it); --size_; return 1; } } return 0; } void rehash(size_t new_count) { if (new_count < 1) new_count = 1; std::vector<Bucket> old; old.swap(buckets_); buckets_.resize(new_count); for (auto& b : old) { for (auto& val : b) { buckets_[hash_(getKey(val)) % new_count].push_back(std::move(val)); } } } void reserve(size_t n) { size_t need = static_cast<size_t>(n / max_load_) + 1; if (need > buckets_.size()) rehash(need); } };改动只有三处:加了key_of_成员、getKey()统一入口、把原来写死的kv.first全换成getKey(...)。
四、封装myunordered_set
template <typename K, typename Hash = std::hash<K>> class myunordered_set { using HT = HashTable<K, K, Identity<K>, Hash>; HT ht_; public: using iterator = typename HT::iterator; using const_iterator = typename HT::const_iterator; myunordered_set() = default; std::pair<iterator, bool> insert(const K& key) { return ht_.insert(key); } iterator find(const K& key) { return ht_.find(key); } size_t erase(const K& key) { return ht_.erase(key); } size_t size() const { return ht_.size(); } bool empty() const { return ht_.empty(); } void reserve(size_t n) { ht_.reserve(n); } iterator begin() { return ht_.begin(); } iterator end() { return ht_.end(); } const_iterator begin() const { return ht_.begin(); } const_iterator end() const { return ht_.end(); } // 便捷接口 bool contains(const K& key) { return ht_.find(key) != ht_.end(); } size_t count(const K& key) const { return const_cast<myunordered_set*>(this)->contains(key) ? 1 : 0; } };注意Identity<K>—— 告诉哈希表"元素本身就是键"。
测试
int main() { myunordered_set<std::string> s; s.insert("apple"); s.insert("banana"); s.insert("apple"); // 重复,不会插入 std::cout << "大小: " << s.size() << '\n'; // 2 std::cout << "banana 在吗? " << (s.contains("banana") ? "在" : "不在") << '\n'; s.erase("banana"); std::cout << "删除后大小: " << s.size() << '\n'; // 1 std::cout << "遍历: "; for (const auto& x : s) std::cout << x << ' '; std::cout << '\n'; return 0; }五、封装myunordered_map
template <typename K, typename V, typename Hash = std::hash<K>> class myunordered_map { using Pair = std::pair<const K, V>; using HT = HashTable<K, Pair, SelectFirst<K, V>, Hash>; HT ht_; public: using iterator = typename HT::iterator; using const_iterator = typename HT::const_iterator; myunordered_map() = default; std::pair<iterator, bool> insert(const Pair& kv) { return ht_.insert(kv); } iterator find(const K& key) { return ht_.find(key); } size_t erase(const K& key) { return ht_.erase(key); } size_t size() const { return ht_.size(); } bool empty() const { return ht_.empty(); } void reserve(size_t n) { ht_.reserve(n); } iterator begin() { return ht_.begin(); } iterator end() { return ht_.end(); } const_iterator begin() const { return ht_.begin(); } const_iterator end() const { return ht_.end(); } /* ★ operator[]:不存在则插入默认值 */ V& operator[](const K& key) { auto it = ht_.find(key); if (it != ht_.end()) return it->second; // 插入默认构造的 V auto [newIt, ok] = ht_.insert(Pair(key, V{})); return newIt->second; } /* ★ at():不存在抛异常 */ V& at(const K& key) { auto it = ht_.find(key); if (it == ht_.end()) throw std::out_of_range("key not found"); return it->second; } };operator[]的实现要点
V& operator[](const K& key) { auto it = ht_.find(key); if (it != ht_.end()) return it->second; // 已存在,直接返回引用 auto [newIt, ok] = ht_.insert(Pair(key, V{})); // 不存在,插入默认值 return newIt->second; // 必须返回插入后元素的引用 }⚠️绝对不能返回局部变量的引用,必须返回哈希表内部元素的引用:
// ❌ 悬垂引用 V& operator[](const K& key) { V v{}; ht_.insert({key, v}); return v; // v 已析构! }测试
int main() { myunordered_map<std::string, int> m; // operator[] 自动插入 m["apple"] = 1; m["banana"] = 2; m["cherry"] = 3; std::cout << "apple = " << m["apple"] << '\n'; // 访问不存在的键 → 插入默认值 0 std::cout << "访问不存在的 pear: " << m["pear"] << '\n'; std::cout << "大小变为: " << m.size() << '\n'; // 4 // 修改值 m["apple"] = 100; std::cout << "修改后 apple = " << m["apple"] << '\n'; // 遍历 std::cout << "遍历:\n"; for (const auto& [k, v] : m) { std::cout << " " << k << " -> " << v << '\n'; } // at() 越界 try { m.at("notexist"); } catch (const std::out_of_range& e) { std::cout << "at 抛异常: " << e.what() << '\n'; } return 0; }六、关键坑点
坑 1:KeyOfValue返回引用,别返回值
// ❌ 返回拷贝,每次比较都构造临时对象 K operator()(const K& key) const { return key; } // ✅ 返回 const 引用 const K& operator()(const K& key) const { return key; }坑 2:iterator与const_iterator的转换
标准库允许iterator隐式转换为const_iterator,但反之不行。上面代码用enable_if_t<IsConst && !OtherConst>精确控制了这个方向。
如果写成无条件模板构造,会出现const_iterator转iterator的漏洞 —— 等于绕过了 const 保护。
坑 3:迭代器的skipEmpty要处理"最后一个桶"
void skipEmpty() { while (table_ && idx_ < bucket_count_ && it_ == table_[idx_].end()) { ++idx_; if (idx_ < bucket_count_) it_ = table_[idx_].begin(); // ★ 边界检查 } }++idx_之后可能已经越界(idx_ == bucket_count_),此时不能再解引用table_[idx_]。漏掉这个判断在空表或末尾时必然崩溃。
坑 4:end()迭代器的构造
iterator end() { return iterator(buckets_.data(), buckets_.size(), buckets_.size()); }idx_ == bucket_count_,skipEmpty的循环条件立刻为假,不会解引用越界下标。这是刻意的设计。
坑 5:operator[]无法用于const对象
const myunordered_map<std::string,int> m; int x = m["a"]; // ❌ 编译错误(operator[] 会插入,必须非 const) int y = m.at("a"); // ✅这是正确行为,不是 bug。标准库同理。
坑 6:Pair里的键必须const
using Pair = std::pair<const K, V>; // ^^^^^ 不能省省了会导致rehash时std::move把键改掉,元素在桶里的位置就和哈希值不一致了。
七、为什么标准库要这么设计
你可能会问:为什么不直接写两个独立的容器,非要抽象出HashTable?
原因有三:
- 代码复用:哈希函数、冲突解决、扩容、迭代器遍历 —— 这些逻辑完全一样,写两遍就是两倍维护成本。
- 行为一致:
unordered_set和unordered_map的迭代器失效规则、扩容时机、复杂度保证完全统一,因为它们真的是同一份代码。 - 易于扩展:再加一个
unordered_multimap(允许重复键),只需要改insert的重复判断,其余全部复用。
同样的设计也用在std::map/std::set上—— 它们共享同一份红黑树实现,靠的也是KeyOfValue这个技巧(标准库里叫_KeyOfValue或KeyExtract)。
八、总结
| 要点 | 结论 |
|---|---|
| 核心技巧 | KeyOfValue仿函数,让底层容器不关心元素类型 |
| set 的 KeyOfValue | Identity,返回元素自身 |
| map 的 KeyOfValue | SelectFirst,返回pair.first |
| 模板参数 | <K, T, KeyOfValue, Hash>,T 是实际存储类型 |
operator[] | 必须返回容器内元素的引用,不能返回局部变量 |
| 迭代器转换 | 只能iterator→const_iterator,方向不能反 |
skipEmpty | 必须做idx_ < bucket_count_边界检查 |
| 底层 pair | 键必须是const |
这个设计的价值不在于"造轮子",而在于理解 STL 的设计哲学:用最小的抽象(一个仿函数)换取最大的复用。看懂之后,你再看std::map、std::set、std::multimap为什么共享实现,就一目了然了。
代码在 GCC 13 / C++17 下编译测试通过。觉得有帮助的话,点赞收藏支持一下。