You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

C++树遍历延迟返回结果:实现find生成器函数

问题

我有如下C++树结构实现:

template<typename T>
class Tree : public std::enable_shared_from_this<Tree<T>>
{
public:
    Tree(T data);
    Tree(T data, std::vector<std::shared_ptr<Tree<T>>> children);

    void add_child(std::shared_ptr<Tree<T>>& child);
    void add_children(std::vector<std::shared_ptr<Tree<T>>> children);

    void set_parent(std::shared_ptr<Tree<T>> parent)
    { this->parent = parent; }
    
    const T get_data() const
    { return this->data; }
    const std::shared_ptr<Tree<T>>& get_parent() const
    { return this->parent; }
    const std::vector<std::shared_ptr<Tree<T>>>& get_children() const
    { return this->children; }

private:
    T data;
    std::shared_ptr<Tree<T>> parent=nullptr;
    std::vector<std::shared_ptr<Tree<T>>> children;
};


template<typename T>
Tree<T>::Tree(T data): data(data)
{ }


template<typename T>
Tree<T>::Tree(T data, std::vector<std::shared_ptr<Tree<T>>> children): data(data)
{ 
    this->add_children(children);
}


template<typename T>
void Tree<T>::add_child(std::shared_ptr<Tree<T>>& child)
{
    this->children.push_back(child);
    child->set_parent(this->shared_from_this());
}


template<typename T>
void Tree<T>::add_children(std::vector<std::shared_ptr<Tree<T>>> children)
{
    for (auto&& child : children)
    {
        this->children.push_back(child);
        child->set_parent(this->shared_from_this());
    }
}

我希望实现一个find函数,签名如下:

template<typename T>
generator<std::shared_ptr<Tree<T>>> find(const std::shared_ptr<Tree<T>>& t, T value);

这个函数接收一棵树和一个目标值,返回一个生成器,用于遍历所有根节点数据等于目标值的子树。我想知道怎么在C++中实现这个生成器?我了解过协程,但觉得它对这个场景(乃至几乎所有场景)都过于复杂。我也尝试过用静态辅助变量实现搜索,但想不出不会遗漏节点的写法。

以下是创建测试树的辅助代码:

#include <vector>
#include <memory>
#include <iostream>

#include "tree.hpp"

std::shared_ptr<Tree<int>> create_tree()
{
    // init tree and nodes
    std::shared_ptr<Tree<int>> tree(new Tree<int>(1));
    std::shared_ptr<Tree<int>> leaf1(new Tree<int>(5));
    std::shared_ptr<Tree<int>> leaf2(new Tree<int>(4));
    std::shared_ptr<Tree<int>> leaf11(new Tree<int>(3));

    // construct tree
    leaf1->add_child(leaf11);
    std::vector<std::shared_ptr<Tree<int>>> v;
    v.push_back(leaf1);
    v.push_back(leaf2);
    tree->add_children(v);

    return tree;
}


int main() 
{
    std::shared_ptr<Tree<int>> tree = create_tree();
    return 0;
}

解决方案

方案1:基于回调的遍历(无需生成器)

如果不想用生成器,最简单的方式是用回调函数。遍历树的过程中,每找到匹配的节点就调用回调,直接处理结果:

template<typename T>
void find(const std::shared_ptr<Tree<T>>& node, T value, const std::function<void(std::shared_ptr<Tree<T>>)>& callback)
{
    if (!node) return;

    // 当前节点匹配则调用回调
    if (node->get_data() == value)
    {
        callback(node);
    }

    // 递归遍历所有子节点
    for (const auto& child : node->get_children())
    {
        find(child, value, callback);
    }
}

使用示例:

int main() 
{
    std::shared_ptr<Tree<int>> tree = create_tree();
    find(tree, 5, [](auto node) {
        std::cout << "Found node with value: " << node->get_data() << std::endl;
    });
    return 0;
}

方案2:手动实现简单生成器(可迭代)

如果一定要返回类似生成器的可迭代对象,可以自己实现一个类,用栈保存遍历状态,支持范围for循环:

template<typename T>
class TreeGenerator
{
public:
    TreeGenerator(std::shared_ptr<Tree<T>> root, T value)
        : m_root(root), m_value(value)
    {
        if (m_root)
        {
            m_stack.push(m_root);
        }
    }

    // 迭代器类
    class Iterator
    {
    public:
        using value_type = std::shared_ptr<Tree<T>>;
        using reference = const value_type&;
        using pointer = const value_type*;
        using iterator_category = std::input_iterator_tag;

        Iterator(std::stack<std::shared_ptr<Tree<T>>> stack, T value)
            : m_stack(std::move(stack)), m_value(value)
        {
            // 找到第一个匹配的节点
            advance_to_match();
        }

        reference operator*() const { return m_current; }
        pointer operator->() const { return &m_current; }

        Iterator& operator++()
        {
            advance();
            advance_to_match();
            return *this;
        }

        bool operator==(const Iterator& other) const
        {
            return m_stack.empty() == other.m_stack.empty() && m_current == other.m_current;
        }

        bool operator!=(const Iterator& other) const
        {
            return !(*this == other);
        }

    private:
        void advance()
        {
            if (m_stack.empty()) return;

            auto node = m_stack.top();
            m_stack.pop();

            // 把子节点逆序入栈,保证遍历顺序和递归一致
            auto children = node->get_children();
            for (auto it = children.rbegin(); it != children.rend(); ++it)
            {
                m_stack.push(*it);
            }
        }

        void advance_to_match()
        {
            while (!m_stack.empty())
            {
                auto node = m_stack.top();
                if (node->get_data() == m_value)
                {
                    m_current = node;
                    return;
                }
                advance();
            }
            m_current.reset();
        }

        std::stack<std::shared_ptr<Tree<T>>> m_stack;
        T m_value;
        std::shared_ptr<Tree<T>> m_current;
    };

    Iterator begin() const
    {
        return Iterator(m_stack, m_value);
    }

    Iterator end() const
    {
        return Iterator({}, m_value);
    }

private:
    std::shared_ptr<Tree<T>> m_root;
    T m_value;
    mutable std::stack<std::shared_ptr<Tree<T>>> m_stack; // mutable因为begin会复制栈
};

// 包装成find函数
template<typename T>
TreeGenerator<T> find(const std::shared_ptr<Tree<T>>& t, T value)
{
    return TreeGenerator<T>(t, value);
}

使用示例:

int main() 
{
    std::shared_ptr<Tree<int>> tree = create_tree();
    for (auto node : find(tree, 5))
    {
        std::cout << "Found node: " << node->get_data() << std::endl;
    }
    return 0;
}

方案3:C++20协程实现(简洁版)

虽然你觉得协程复杂,但针对树遍历场景,协程的代码其实非常直观,不需要手动管理状态栈:

首先需要定义一个简单的生成器类型(C++20没有内置generator,这里给出极简版):

#include <coroutine>

template<typename T>
class generator
{
public:
    struct promise_type
    {
        T current_value;
        std::suspend_always yield_value(T value)
        {
            current_value = value;
            return {};
        }
        std::suspend_never initial_suspend() { return {}; }
        std::suspend_never final_suspend() noexcept { return {}; }
        generator get_return_object() { return generator{this}; }
        void return_void() {}
        void unhandled_exception() { std::terminate(); }
    };

    struct iterator
    {
        std::coroutine_handle<promise_type> coro;
        bool done;

        iterator(std::coroutine_handle<promise_type> h, bool d) : coro(h), done(d) {}

        T operator*() const { return coro.promise().current_value; }
        iterator& operator++()
        {
            coro.resume();
            done = coro.done();
            return *this;
        }

        bool operator==(const iterator& other) const { return done == other.done; }
        bool operator!=(const iterator& other) const { return !(*this == other); }
    };

    generator(generator&& other) : coro(other.coro) { other.coro = nullptr; }
    ~generator() { if (coro) coro.destroy(); }

    iterator begin()
    {
        coro.resume();
        return iterator{coro, coro.done()};
    }

    iterator end() { return iterator{coro, true}; }

private:
    explicit generator(promise_type* p) : coro(std::coroutine_handle<promise_type>::from_promise(*p)) {}
    std::coroutine_handle<promise_type> coro;
};

然后实现find协程函数:

template<typename T>
generator<std::shared_ptr<Tree<T>>> find(const std::shared_ptr<Tree<T>>& node, T value)
{
    if (!node) co_return;

    // 当前节点匹配则yield
    if (node->get_data() == value)
    {
        co_yield node;
    }

    // 遍历子节点
    for (const auto& child : node->get_children())
    {
        // 协程可以直接嵌套调用另一个生成器,自动遍历其结果
        for (auto match : find(child, value))
        {
            co_yield match;
        }
    }
}

使用示例和方案2一致,直接用范围for循环即可。


内容的提问来源于stack exchange,提问作者Samuel

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.04 11:10:50