如何用enable_if实现带不定参数的C++模板特化
C++20带可变参数的类模板针对基类派生类的特化实现
需求目标
实现如下形式的类模板:
template<typename T, typename ...Args>
其中T为类类型,当T是BaseA或其任意派生类时,使用模板的特化版本(基于C++20标准)。
遇到的问题
当模板第二个参数为固定类型时,通过std::enable_if结合默认模板参数可以轻松实现特化:
template <typename T, typename U, typename Enable = void> class Factory { public: Factory(U arg) { m_spInner = std::make_shared<T>(arg); } void print() const { printf("Factory normal: "); m_spInner->print(); } private: std::shared_ptr<T> m_spInner; }; template <typename T, typename U> class Factory<T, U, typename std::enable_if<std::is_base_of<BaseA, T>::value>::type> { public: Factory(U arg) { m_spInner = std::make_shared<T>(arg); } void print() const { printf("Factory BaseA: "); m_spInner->print(); } private: std::shared_ptr<T> m_spInner; };
但当第二个参数改为可变参数包...Args时,会遇到编译问题:
- 若将
Enable放在...Args之后,编译器报错“可变参数必须是最后一个模板参数”; - 若将
Enable放在...Args之前,调用模板时无法自动推导Enable参数,导致匹配失败。
解决方案
方法1:使用C++20 Concepts(推荐)
Concepts是C++20引入的特性,能直接对模板参数进行约束,写法简洁且可读性高:
首先定义一个Concept,用于判断类型是否派生自BaseA:
template<typename T> concept DerivedFromBaseA = std::is_base_of_v<BaseA, T>;
然后编写主模板和特化版本:
// 普通版本:适用于非BaseA派生类的类型 template<typename T, typename... Args> class Factory { public: Factory(Args&&... args) : m_spInner(std::make_shared<T>(std::forward<Args>(args)...)) {} void print() const { printf("Factory normal: "); m_spInner->print(); } private: std::shared_ptr<T> m_spInner; }; // 特化版本:仅适用于BaseA或其派生类 template<DerivedFromBaseA T, typename... Args> class Factory<T, Args...> { public: Factory(Args&&... args) : m_spInner(std::make_shared<T>(std::forward<Args>(args)...)) {} void print() const { printf("Factory BaseA: "); m_spInner->print(); } private: std::shared_ptr<T> m_spInner; };
方法2:结合std::enable_if与模板参数默认值
如果不想使用Concepts,也可以通过调整模板参数顺序,将enable_if的结果作为一个带有默认值的模板参数,放在可变参数包之后:
// 主模板:默认处理非BaseA派生类 template<typename T, typename... Args, typename = std::enable_if_t<!std::is_base_of_v<BaseA, T>>> class Factory { public: Factory(Args&&... args) : m_spInner(std::make_shared<T>(std::forward<Args>(args)...)) {} void print() const { printf("Factory normal: "); m_spInner->print(); } private: std::shared_ptr<T> m_spInner; }; // 特化版本:匹配BaseA派生类 template<typename T, typename... Args> class Factory<T, Args..., std::enable_if_t<std::is_base_of_v<BaseA, T>>> { public: Factory(Args&&... args) : m_spInner(std::make_shared<T>(std::forward<Args>(args)...)) {} void print() const { printf("Factory BaseA: "); m_spInner->print(); } private: std::shared_ptr<T> m_spInner; };
完整测试代码
#include <type_traits> #include <cstdio> #include <memory> #include <string> using namespace std; class BaseA { public: BaseA(int id): m_id(id) {} virtual void print() const { printf("BaseA\n"); } private: int m_id; }; class BaseB { public: BaseB(char id): m_id(id) {} virtual void print() const { printf("BaseB\n"); } private: char m_id; }; class DerivedA1 : public BaseA { public: DerivedA1(int id) : BaseA(id) {} void print() const override { printf("DerivedA1\n"); } }; class DerivedA2 : public BaseA { public: DerivedA2(int id) : BaseA(id) {} void print() const override { printf("DerivedA2\n"); } }; class DerivedB : public BaseB { public: DerivedB(char id) : BaseB(id) {} void print() const override { printf("DerivedB\n"); } }; class C { public: C(string id) : m_id(id) {} void print() const { printf("C\n"); } private: string m_id; }; // 以下使用Concepts方案,若要使用enable_if方案替换上述Factory定义即可 template<typename T> concept DerivedFromBaseA = std::is_base_of_v<BaseA, T>; template<typename T, typename... Args> class Factory { public: Factory(Args&&... args) : m_spInner(std::make_shared<T>(std::forward<Args>(args)...)) {} void print() const { printf("Factory normal: "); m_spInner->print(); } private: std::shared_ptr<T> m_spInner; }; template<DerivedFromBaseA T, typename... Args> class Factory<T, Args...> { public: Factory(Args&&... args) : m_spInner(std::make_shared<T>(std::forward<Args>(args)...)) {} void print() const { printf("Factory BaseA: "); m_spInner->print(); } private: std::shared_ptr<T> m_spInner; }; int main() { Factory<DerivedA1, int> factoryA1(1); factoryA1.print(); // 输出:Factory BaseA: DerivedA1 Factory<DerivedA2, int> factoryA2(2); factoryA2.print(); // 输出:Factory BaseA: DerivedA2 Factory<DerivedB, char> factoryB('B'); factoryB.print(); // 输出:Factory normal: DerivedB Factory<C, string> factoryC("C"); factoryC.print(); // 输出:Factory normal: C return 0; }
内容的提问来源于stack exchange,提问作者BookSword
相关产品推荐
相关产品推荐

