S SmartDocs
Serie: C++ cpp 384 líneas · Actualizado 2026-04-03

trailing_return.cpp

C++/Part5_進階主題/Ch21_型別推導/trailing_return.cpp

// trailing_return.cpp
// 編譯指令:g++ -std=c++17 -Wall trailing_return.cpp -o trailing_return
//
// 本程式示範尾端回傳型別(trailing return type)的用法與應用場景

#include <iostream>
#include <string>
#include <vector>
#include <array>
#include <type_traits>
#include <typeinfo>
#include <cxxabi.h>
#include <cmath>

// 輔助函式:取得可讀的型別名稱
template<typename T>
std::string type_name() {
    int status;
    char* demangled = abi::__cxa_demangle(typeid(T).name(), nullptr, nullptr, &status);
    std::string result = (status == 0) ? demangled : typeid(T).name();
    free(demangled);
    return result;
}

// ============================================================
// 第一部分:基本語法
// ============================================================

// 傳統寫法
int traditional_add(int a, int b) {
    return a + b;
}

// 尾端回傳型別寫法 — 功能相同,但語法不同
auto trailing_add(int a, int b) -> int {
    return a + b;
}

// ============================================================
// 第二部分:模板中回傳型別依賴參數
// ============================================================

// 問題:回傳型別取決於 T 和 U 的運算結果
// 傳統方式在 C++11 無法直接寫出回傳型別(a、b 尚未宣告)
// 使用尾端回傳型別可以引用參數名
template<typename T, typename U>
auto add(T a, U b) -> decltype(a + b) {
    return a + b;
}

// C++14 起可以省略尾端回傳型別,讓編譯器自動推導
template<typename T, typename U>
auto multiply(T a, U b) {
    return a * b;
}

void demo_basic_trailing() {
    std::cout << "========================================\n";
    std::cout << "  基本尾端回傳型別\n";
    std::cout << "========================================\n\n";

    std::cout << "traditional_add(3, 4) = " << traditional_add(3, 4) << "\n";
    std::cout << "trailing_add(3, 4)    = " << trailing_add(3, 4) << "\n\n";

    // 泛型 add:不同型別相加
    auto r1 = add(1, 2.5);       // int + double -> double
    auto r2 = add(1.0f, 2);      // float + int -> float
    auto r3 = add(1L, 2);        // long + int -> long

    std::cout << "add(1, 2.5)   = " << r1 << " (型別: " << type_name<decltype(r1)>() << ")\n";
    std::cout << "add(1.0f, 2)  = " << r2 << " (型別: " << type_name<decltype(r2)>() << ")\n";
    std::cout << "add(1L, 2)    = " << r3 << " (型別: " << type_name<decltype(r3)>() << ")\n";

    auto m1 = multiply(3, 4.5);
    std::cout << "multiply(3, 4.5) = " << m1 << " (型別: " << type_name<decltype(m1)>() << ")\n";
    std::cout << "\n";
}

// ============================================================
// 第三部分:讓複雜回傳型別更易讀
// ============================================================

// 不使用尾端回傳型別 — 回傳型別在最前面,難以閱讀
std::vector<std::pair<std::string, int>>::const_iterator
find_student_traditional(
    const std::vector<std::pair<std::string, int>>& students,
    const std::string& name)
{
    for (auto it = students.cbegin(); it != students.cend(); ++it) {
        if (it->first == name) return it;
    }
    return students.cend();
}

// 使用尾端回傳型別 — 先看函式名和參數,再看回傳型別
auto find_student_trailing(
    const std::vector<std::pair<std::string, int>>& students,
    const std::string& name)
    -> std::vector<std::pair<std::string, int>>::const_iterator
{
    for (auto it = students.cbegin(); it != students.cend(); ++it) {
        if (it->first == name) return it;
    }
    return students.cend();
}

void demo_readability() {
    std::cout << "========================================\n";
    std::cout << "  改善可讀性\n";
    std::cout << "========================================\n\n";

    std::vector<std::pair<std::string, int>> students = {
        {"Alice", 95}, {"Bob", 87}, {"Charlie", 92}
    };

    auto it1 = find_student_traditional(students, "Bob");
    auto it2 = find_student_trailing(students, "Charlie");

    if (it1 != students.cend())
        std::cout << "找到 " << it1->first << ": " << it1->second << " 分\n";
    if (it2 != students.cend())
        std::cout << "找到 " << it2->first << ": " << it2->second << " 分\n";

    std::cout << "\n兩種寫法功能相同,但尾端回傳型別讓函式簽名更清楚\n";
    std::cout << "\n";
}

// ============================================================
// 第四部分:decltype 在尾端回傳型別中的應用
// ============================================================

// 泛型容器存取:根據容器的 operator[] 推導回傳型別
template<typename Container, typename Index>
auto get_element(Container& c, Index i) -> decltype(c[i]) {
    return c[i];
}

// 泛型容器大小比較
template<typename C1, typename C2>
auto size_diff(const C1& c1, const C2& c2) -> decltype(static_cast<long>(c1.size()) - static_cast<long>(c2.size())) {
    return static_cast<long>(c1.size()) - static_cast<long>(c2.size());
}

void demo_decltype_trailing() {
    std::cout << "========================================\n";
    std::cout << "  decltype 與尾端回傳型別\n";
    std::cout << "========================================\n\n";

    std::vector<int> vec = {10, 20, 30};
    std::string str = "Hello";

    // get_element 可用於不同容器
    auto& v_elem = get_element(vec, 1);
    auto& s_elem = get_element(str, 0);

    std::cout << "vec[1] = " << v_elem << "\n";
    std::cout << "str[0] = " << s_elem << "\n";

    // 透過參考修改原始容器
    v_elem = 99;
    s_elem = 'h';
    std::cout << "\n修改後:\n";
    std::cout << "vec[1] = " << vec[1] << "\n";
    std::cout << "str = " << str << "\n";

    // 大小差異
    std::vector<int> v1 = {1, 2, 3, 4, 5};
    std::vector<int> v2 = {1, 2};
    std::cout << "\nsize_diff({1,2,3,4,5}, {1,2}) = " << size_diff(v1, v2) << "\n";
    std::cout << "\n";
}

// ============================================================
// 第五部分:矩陣運算範例
// ============================================================

template<typename T, std::size_t Rows, std::size_t Cols>
struct Matrix {
    std::array<std::array<T, Cols>, Rows> data{};

    auto at(std::size_t r, std::size_t c) -> T& {
        return data[r][c];
    }

    auto at(std::size_t r, std::size_t c) const -> const T& {
        return data[r][c];
    }

    auto row_count() const -> std::size_t { return Rows; }
    auto col_count() const -> std::size_t { return Cols; }
};

// 矩陣相加:兩個不同元素型別的矩陣相加
template<typename T, typename U, std::size_t Rows, std::size_t Cols>
auto matrix_add(const Matrix<T, Rows, Cols>& a, const Matrix<U, Rows, Cols>& b)
    -> Matrix<decltype(std::declval<T>() + std::declval<U>()), Rows, Cols>
{
    Matrix<decltype(std::declval<T>() + std::declval<U>()), Rows, Cols> result;
    for (std::size_t r = 0; r < Rows; ++r) {
        for (std::size_t c = 0; c < Cols; ++c) {
            result.at(r, c) = a.at(r, c) + b.at(r, c);
        }
    }
    return result;
}

// 矩陣純量乘法
template<typename T, typename Scalar, std::size_t Rows, std::size_t Cols>
auto matrix_scale(const Matrix<T, Rows, Cols>& m, Scalar s)
    -> Matrix<decltype(std::declval<T>() * std::declval<Scalar>()), Rows, Cols>
{
    Matrix<decltype(std::declval<T>() * std::declval<Scalar>()), Rows, Cols> result;
    for (std::size_t r = 0; r < Rows; ++r) {
        for (std::size_t c = 0; c < Cols; ++c) {
            result.at(r, c) = m.at(r, c) * s;
        }
    }
    return result;
}

template<typename T, std::size_t Rows, std::size_t Cols>
void print_matrix(const std::string& name, const Matrix<T, Rows, Cols>& m) {
    std::cout << name << " (" << Rows << "x" << Cols
              << ", 元素型別: " << type_name<T>() << "):\n";
    for (std::size_t r = 0; r < Rows; ++r) {
        std::cout << "  [";
        for (std::size_t c = 0; c < Cols; ++c) {
            if (c > 0) std::cout << ", ";
            std::cout << m.at(r, c);
        }
        std::cout << "]\n";
    }
}

void demo_matrix_operations() {
    std::cout << "========================================\n";
    std::cout << "  矩陣運算(尾端回傳型別應用)\n";
    std::cout << "========================================\n\n";

    Matrix<int, 2, 3> intMatrix;
    intMatrix.at(0, 0) = 1; intMatrix.at(0, 1) = 2; intMatrix.at(0, 2) = 3;
    intMatrix.at(1, 0) = 4; intMatrix.at(1, 1) = 5; intMatrix.at(1, 2) = 6;

    Matrix<double, 2, 3> dblMatrix;
    dblMatrix.at(0, 0) = 0.1; dblMatrix.at(0, 1) = 0.2; dblMatrix.at(0, 2) = 0.3;
    dblMatrix.at(1, 0) = 0.4; dblMatrix.at(1, 1) = 0.5; dblMatrix.at(1, 2) = 0.6;

    print_matrix("整數矩陣 A", intMatrix);
    std::cout << "\n";
    print_matrix("浮點矩陣 B", dblMatrix);
    std::cout << "\n";

    // int + double -> double(由尾端回傳型別的 decltype 自動推導)
    auto sum = matrix_add(intMatrix, dblMatrix);
    print_matrix("A + B", sum);
    std::cout << "\n";

    // int * double -> double
    auto scaled = matrix_scale(intMatrix, 2.5);
    print_matrix("A * 2.5", scaled);
    std::cout << "\n";
}

// ============================================================
// 第六部分:進階應用 — SFINAE 與條件回傳型別
// ============================================================

// 安全除法:整數回傳 double,浮點數保持原型別
template<typename T>
auto safe_divide(T a, T b) -> std::conditional_t<std::is_integral_v<T>, double, T> {
    if (b == T{}) {
        std::cout << "  警告:除以零!\n";
        return {};
    }
    return static_cast<std::conditional_t<std::is_integral_v<T>, double, T>>(a) /
           static_cast<std::conditional_t<std::is_integral_v<T>, double, T>>(b);
}

// 條件回傳:根據型別選擇不同運算
template<typename T>
auto transform(T value) -> decltype(auto) {
    if constexpr (std::is_integral_v<T>) {
        return value * value;       // 整數:平方
    } else if constexpr (std::is_floating_point_v<T>) {
        return std::sqrt(value);    // 浮點數:開根號
    } else {
        return value;
    }
}

void demo_advanced_trailing() {
    std::cout << "========================================\n";
    std::cout << "  進階:條件回傳型別\n";
    std::cout << "========================================\n\n";

    // safe_divide:整數除法自動提升為 double
    auto r1 = safe_divide(7, 2);     // int -> double
    auto r2 = safe_divide(7.0, 2.0); // double -> double
    auto r3 = safe_divide(7.0f, 2.0f); // float -> float

    std::cout << "safe_divide(7, 2)       = " << r1
              << " (型別: " << type_name<decltype(r1)>() << ")\n";
    std::cout << "safe_divide(7.0, 2.0)   = " << r2
              << " (型別: " << type_name<decltype(r2)>() << ")\n";
    std::cout << "safe_divide(7.0f, 2.0f) = " << r3
              << " (型別: " << type_name<decltype(r3)>() << ")\n";

    // transform:根據型別執行不同運算
    std::cout << "\ntransform(5)    = " << transform(5) << " (整數平方)\n";
    std::cout << "transform(25.0) = " << transform(25.0) << " (浮點數開根號)\n";
    std::cout << "\n";
}

// ============================================================
// 第七部分:成員函式的尾端回傳型別
// ============================================================

class Calculator {
    double result_ = 0.0;

public:
    // 鏈式呼叫:回傳 *this 的參考
    auto add(double val) -> Calculator& {
        result_ += val;
        return *this;
    }

    auto subtract(double val) -> Calculator& {
        result_ -= val;
        return *this;
    }

    auto multiply_by(double val) -> Calculator& {
        result_ *= val;
        return *this;
    }

    auto get_result() const -> double {
        return result_;
    }

    auto reset() -> Calculator& {
        result_ = 0.0;
        return *this;
    }
};

void demo_member_trailing() {
    std::cout << "========================================\n";
    std::cout << "  成員函式的尾端回傳型別\n";
    std::cout << "========================================\n\n";

    Calculator calc;
    double result = calc.add(10)
                        .multiply_by(3)
                        .subtract(5)
                        .get_result();

    std::cout << "(10 * 3) - 5 = " << result << "\n";

    calc.reset().add(100).subtract(50);
    std::cout << "100 - 50 = " << calc.get_result() << "\n";
    std::cout << "\n";
}

// ============================================================
// 主程式
// ============================================================

int main() {
    std::cout << "╔══════════════════════════════════════════╗\n";
    std::cout << "║  C++17 尾端回傳型別(Trailing Return Type)║\n";
    std::cout << "╚══════════════════════════════════════════╝\n\n";

    demo_basic_trailing();
    demo_readability();
    demo_decltype_trailing();
    demo_matrix_operations();
    demo_advanced_trailing();
    demo_member_trailing();

    std::cout << "=== 程式結束 ===\n";
    return 0;
}

Artículos relacionados