fastapi yiled与资源管理的异步机制

"""
第六步:async with / 异步生成器 / 异步 yield 依赖

目标:把第四步(yield 依赖)升级成异步版。
async with = with + "进入/退出时可以 await"
异步 yield 依赖 = async def + yield:yield 前 await 准备,yield 后 await 清理。

对应 fastapi:async def 写的 yield 依赖(异步开关数据库连接)。
对应 langchain:astream 返回异步生成器,用 async for 消费。
"""

import asyncio
import inspect


class Depend:
    def __init__(self, dependency):
        self.dependency = dependency


class App:
    def __init__(self):
        self.routes = {}

    def get(self, path):
        def decorator(func):
            self.routes[("GET", path)] = func
            return func
        return decorator

    async def solve(self, func, raw_params, cache, agens):
        """
        agens:收集异步生成器(异步 yield 依赖),供请求结束后清理。
        """
        kwargs = {}
        for name, param in inspect.signature(func).parameters.items():
            if isinstance(param.default, Depend):
                dep_func = param.default.dependency
                if dep_func in cache:
                    kwargs[name] = cache[dep_func]
                elif inspect.isasyncgenfunction(dep_func):
                    # 异步 yield 依赖:async def + yield
                    agen = dep_func()
                    value = await agen.__anext__()   # 推进到 yield,拿到准备好的资源
                    agens.append(agen)               # 记下来,稍后清理
                    cache[dep_func] = kwargs[name] = value
                else:
                    dep_kwargs = await self.solve(dep_func, raw_params, cache, agens)
                    result = dep_func(**dep_kwargs)
                    if inspect.iscoroutine(result):
                        result = await result
                    cache[dep_func] = kwargs[name] = result
            else:
                if name in raw_params:
                    type_ = param.annotation
                    kwargs[name] = type_(raw_params[name]) if type_ is not inspect.Parameter.empty else raw_params[name]
                elif param.default is not inspect.Parameter.empty:
                    kwargs[name] = param.default
                else:
                    raise ValueError(f"缺少参数: {name}")
        return kwargs

    async def handle(self, method, path, raw_params):
        func = self.routes.get((method, path))
        if func is None:
            return "404 Not Found"

        cache = {}
        agens = []
        try:
            kwargs = await self.solve(func, raw_params, cache, agens)
            result = func(**kwargs)
            if inspect.iscoroutine(result):
                result = await result
            return result
        finally:
            # 请求结束(无论成功失败),逆序清理每个异步 yield 依赖
            for agen in reversed(agens):
                try:
                    await agen.__anext__()   # 推进过 yield -> 执行清理 -> 抛 StopAsyncIteration
                except StopAsyncIteration:
                    pass


app = App()


# 异步 yield 依赖:异步地开/关数据库连接
async def get_db():
    print("    [db] 正在异步建立连接...(await)")
    await asyncio.sleep(0.3)
    print("    [db] 连接就绪")
    try:
        yield "db_connection"
    finally:
        print("    [db] 正在异步关闭连接...(await)")
        await asyncio.sleep(0.3)
        print("    [db] 连接已关闭")


@app.get("/query")
async def query(db=Depend(get_db)):
    print(f"    [路由] 用 {db} 查询数据")
    return f"查询结果(via {db})"


async def main():
    print("=== 处理一个请求 ===")
    result = await app.handle("GET", "/query", {})
    print("返回:", result)


if __name__ == "__main__":
    asyncio.run(main())

输出结果:

=== 处理一个请求 ===
    [db] 正在异步建立连接...(await)
    [db] 连接就绪
    [路由] 用 db_connection 查询数据
    [db] 正在异步关闭连接...(await)
    [db] 连接已关闭
返回: 查询结果(via db_connection)
posted @ 2026-07-03 12:49  RolandHe  阅读(3)  评论(0)    收藏  举报