create_task包着langgraph,节点error如何传递给最外层

task = asyncio.create_task(tansfer_langgraph()) async def transfer_langgraph(): langgraph_.astream() 如何每个langgraph的节点报错,如何将错误抛给 langgraph, langgraph抛给task,然后在task外面知道是哪种错误呢

可以这样理解:

节点里 raise → LangGraph 图运行报错 → astream() 抛错 → transfer_langgraph() 抛错 → asyncio.Task 变成 failed → 外层 await task / task.result() 拿到原始异常。

LangGraph 的节点本质是 Python 同步或异步函数;astream 是异步流式接口,需要用 async for 消费,否则图根本没有真正跑起来。(LangChain 文档)

你的代码应该类似这样:

task = asyncio.create_task(transfer_langgraph())

async def transfer_langgraph():
    async for chunk in langgraph_.astream(input_state):
        print(chunk)

但这里要注意:asyncio.create_task(...) 那一行必须在 transfer_langgraph 定义之后执行,否则函数名还不存在。


最简单:直接让节点异常自然抛出

async def node_a(state):
    raise ValueError("node_a 参数错误")


async def transfer_langgraph():
    async for chunk in langgraph_.astream(input_state):
        print(chunk)


task = asyncio.create_task(transfer_langgraph())

try:
    await task
except ValueError as e:
    print("外层知道是 ValueError:", e)
except TimeoutError as e:
    print("外层知道是 TimeoutError:", e)
except Exception as e:
    print("其他错误:", type(e), e)

只要节点里没有把异常吞掉,异常会一路冒泡到 await task


推荐:包装节点,带上“哪个节点报错”

如果你想在 task 外面知道:

  1. 是哪种错误;
  2. 是哪个 LangGraph 节点报的;
  3. 原始异常是什么;

可以定义一个自定义异常:

class LangGraphNodeError(Exception):
    def __init__(self, node_name: str, original_error: Exception):
        self.node_name = node_name
        self.original_error = original_error
        super().__init__(
            f"LangGraph node `{node_name}` failed: "
            f"{type(original_error).__name__}: {original_error}"
        )

然后写一个节点包装器:

import functools
import inspect


def wrap_node(node_name, fn):
    @functools.wraps(fn)
    async def wrapper(state, *args, **kwargs):
        try:
            result = fn(state, *args, **kwargs)

            if inspect.isawaitable(result):
                result = await result

            return result

        except Exception as e:
            raise LangGraphNodeError(node_name, e) from e

    return wrapper

注册节点时这样写:

graph.add_node("parse_input", wrap_node("parse_input", parse_input))
graph.add_node("call_llm", wrap_node("call_llm", call_llm))
graph.add_node("save_result", wrap_node("save_result", save_result))

外层 task 这样接:

async def transfer_langgraph():
    async for chunk in langgraph_.astream(input_state):
        print(chunk)


task = asyncio.create_task(transfer_langgraph())

try:
    await task

except LangGraphNodeError as e:
    print("报错节点:", e.node_name)
    print("原始错误类型:", type(e.original_error).__name__)
    print("原始错误内容:", e.original_error)

    if isinstance(e.original_error, ValueError):
        print("这是参数错误")
    elif isinstance(e.original_error, TimeoutError):
        print("这是超时错误")
    else:
        print("这是其他节点错误")

except Exception as e:
    print("非节点包装错误:", type(e).__name__, e)

如果你不想 await task,也可以用 callback

def on_task_done(task: asyncio.Task):
    try:
        task.result()
    except LangGraphNodeError as e:
        print("callback 里捕获到节点错误")
        print("节点:", e.node_name)
        print("类型:", type(e.original_error).__name__)
        print("错误:", e.original_error)
    except Exception as e:
        print("callback 里捕获到其他错误:", type(e).__name__, e)


task = asyncio.create_task(transfer_langgraph())
task.add_done_callback(on_task_done)

但通常更推荐:

await task

因为逻辑更清楚。


一个完整模板

import asyncio
import functools
import inspect


class LangGraphNodeError(Exception):
    def __init__(self, node_name: str, original_error: Exception):
        self.node_name = node_name
        self.original_error = original_error
        super().__init__(
            f"LangGraph node `{node_name}` failed: "
            f"{type(original_error).__name__}: {original_error}"
        )


def wrap_node(node_name, fn):
    @functools.wraps(fn)
    async def wrapper(state, *args, **kwargs):
        try:
            result = fn(state, *args, **kwargs)

            if inspect.isawaitable(result):
                result = await result

            return result

        except Exception as e:
            raise LangGraphNodeError(node_name, e) from e

    return wrapper


async def transfer_langgraph(langgraph_, input_state):
    async for chunk in langgraph_.astream(input_state):
        print("chunk:", chunk)


async def main():
    task = asyncio.create_task(
        transfer_langgraph(langgraph_, input_state)
    )

    try:
        await task

    except LangGraphNodeError as e:
        print("LangGraph 节点报错")
        print("节点:", e.node_name)
        print("原始错误类型:", type(e.original_error).__name__)
        print("原始错误:", e.original_error)

    except asyncio.CancelledError:
        print("任务被取消")
        raise

    except Exception as e:
        print("其他错误:", type(e).__name__, e)


asyncio.run(main())

核心点是:不要在节点里 try/except 后只打印不抛出。如果你 catch 了异常,一定要重新 raise,否则 LangGraph 和外层 task 都不会知道失败了。

这个是一直在这里等待了么,我后面还有await text = q.get() yield text的需求,等待在这里就影响后面

对,await task等待这个 task 结束。如果你的后面还有:

text = await q.get()
yield text

那把 await task 放在前面就会阻塞当前协程,后面的 q.get() / yield 不会执行,直到 LangGraph 跑完或报错。

你这种场景不要直接:

await task

而是应该让 LangGraph 在后台跑,同时当前 async generator 继续从 q 里取数据并 yield


推荐结构

async def transfer_langgraph(q: asyncio.Queue):
    try:
        async for chunk in langgraph_.astream(input_state):
            await q.put(chunk)

    except Exception as e:
        # 把异常也塞进 queue,让外层 yield 循环能感知
        await q.put(e)

    finally:
        # 用一个结束信号告诉外层:LangGraph 结束了
        await q.put(None)

然后外层这样写:

async def stream():
    q = asyncio.Queue()

    task = asyncio.create_task(
        transfer_langgraph(q)
    )

    while True:
        item = await q.get()

        if item is None:
            break

        if isinstance(item, Exception):
            raise item

        yield item

这样就不会卡在 await task 上了。


如果你想保留节点错误类型

比如你之前用了:

raise LangGraphNodeError(node_name, e) from e

那外层可以这样识别:

async def stream():
    q = asyncio.Queue()

    task = asyncio.create_task(
        transfer_langgraph(q)
    )

    while True:
        item = await q.get()

        if item is None:
            break

        if isinstance(item, LangGraphNodeError):
            print("节点错误:", item.node_name)
            print("原始错误类型:", type(item.original_error).__name__)
            raise item

        if isinstance(item, Exception):
            raise item

        yield item

更稳一点:避免 task 静默失败

如果 transfer_langgraph() 里没有把异常放进 queue,也可以在循环里检查 task:

async def stream():
    q = asyncio.Queue()

    task = asyncio.create_task(
        transfer_langgraph(q)
    )

    while True:
        item = await q.get()

        if item is None:
            break

        yield item

    # 这里不会影响前面的 yield,因为已经流式输出完了
    try:
        await task
    except Exception as e:
        raise e

但这个版本有个问题:如果 transfer_langgraph() 报错后没有 q.put(None),外层可能永远卡在:

item = await q.get()

所以更推荐前面的写法:异常和结束信号都通过 queue 传出来


最终建议模板

END = object()


async def transfer_langgraph(q: asyncio.Queue):
    try:
        async for chunk in langgraph_.astream(input_state):
            await q.put(chunk)

    except Exception as e:
        await q.put(e)

    finally:
        await q.put(END)


async def stream():
    q = asyncio.Queue()

    task = asyncio.create_task(
        transfer_langgraph(q)
    )

    try:
        while True:
            item = await q.get()

            if item is END:
                break

            if isinstance(item, Exception):
                raise item

            yield item

    finally:
        if not task.done():
            task.cancel()

核心就是:

await task

适合“我现在就等它结束”。

但你这里是“LangGraph 后台跑,前台继续 q.get() 然后 yield”,所以不应该先 await task,而应该用 queue 把结果和异常传出来。

posted @ 2026-05-27 10:16  X1OO  阅读(13)  评论(0)    收藏  举报