单线程下实现生成器向外部回调生成器yield数据的技术问询
我需要实现一个从远程数据库加载数据并写入服务器存储的服务。客户端发起RPC调用import_data(datasource="mysql:host/db/table", storage="backend_storage"),期望定期获取处理进度响应。
服务器端RPC回调签名为def callback(request) -> Iterator[response],服务器会提供线程执行该回调,并迭代其返回值尽快发送响应。服务器是第三方框架,无法修改,伪代码如下:
def service_executor(connection): request = connection.get_request() for response in callback(request): connection.send_response(response)
我期望的实现方式如下:
def callback(request): def loading_data_generator(): for batch in loading_from_database(): yield batch # 能否在当前线程内向callback yield响应 # 将"processed a new batch"作为callback的yield结果 # write_to_storage为第三方库,无法修改 write_to_stroage(wrapped_iterator(loading_data_generator())
由于无法在loading_data_generator内部向callback进行yield,当前实现会在write_to_storage执行完成后才返回,线程在write_to_storage执行期间无法回到service_executor。
我可以将write_to_storage安排到单独线程,通过队列将进度传递到当前回调以yield给service_executor,但所有操作都是串行执行的,能否在单线程下实现该需求?
预期控制流如下:
service_executor -> callback -> write_to_storage -> loading_data_generator -> yield data -> write_to_storage[写入] -> loading_data_generator -> [向service_executor yield进度] -> service_executor[发送] -> .....
极简代码:
def write_to_storage(data_generator): for i in data_generator: print(f"writing {i} to storage") def callback(request): def data_generator(): for i in range(request): yield i # 此处期望向service_executor yield # yield f"{i} consumed" write_to_storage(data_generator()) def service_executor(request): response = callback(request) if isinstance(response, Iterable): for item in response: print(f"show response {response}") if __name__ == '__main__': service_executor(5)
预期输出:
writing 0 to storage show response 0 consumed writing 1 to storage show response 1 consumed writing 2 to storage show response 2 consumed writing 3 to storage show response 3 consumed writing 4 to storage show response 4 consumed
第一次尝试
我尝试将整个回调逻辑包装为异步函数,在当前线程执行事件循环,但不知道如何在任务完成时yield异步任务结果。
def callback(request): loop = asyncio.new_event_loop() asyncio.set_event_loop(loop) q = asyncio.Queue() def load_from_database(): for batch in loading_data_generator(): yield batch asyncio.get_running_loop().create_task(q.put("process a new batch")) async def wrap_writing_into_async(): write_to_storage(wrap_iterator(load_from_database()) loop.create_task(wrap_writing_into_async()) # 尝试从队列获取并yield while True: task = loop.create_task(q.get()) loop.run_until_complete(task) yield task.result()
但未达到预期效果,线程会卡在loop.run_until_complete直到wrap_writing_into_async执行完成。
第二次尝试
我发现了一个名为greenback的库,它可以从同步上下文调用异步协程,同时有方法可以将异步生成器转换为同步生成器。
通过该方式我大致实现了需求:
def iter_over_async(ait, loop): ait = ait.__aiter__() async def get_next(): try: obj = await ait.__anext__() return False, obj except StopAsyncIteration: return True, None while True: done, obj = loop.run_until_complete(get_next()) if done: break yield obj def write_to_storage(data_generator): for i in data_generator: print(f"writing {i} to storage") def callback(request): loop = asyncio.get_event_loop() q = asyncio.Queue(1) def g(): for i in range(10): yield i greenback.await_(q.put(f"we have processed {i}")) greenback.await_(q.join()) async def wrap_consumer(): await greenback.ensure_portal() write_to_storage(g()) t = loop.create_task(wrap_consumer()) async def q_get(): while True: yield await q.get() q.task_done() yield from iter_over_async(q_get(), loop) loop.run_until_complete(t) def service_executor(request): response = callback(request) if isinstance(response, Iterable): for item in response: print(f"show response {item}") if __name__ == "__main__": service_executor(5)
但队列并未限制为仅1个元素,write_to_storage总是会先迭代3次后q_get才会yield结果。
输出结果:
writing 0 to storage writing 1 to storage writing 2 to storage show response we have processed 0 writing 3 to storage show response we have processed 1 writing 4 to storage show response we have processed 2 writing 5 to storage show response we have processed 3 writing 6 to storage show response we have processed 4 writing 7 to storage show response we have processed 5 writing 8 to storage show response we have processed 6 writing 9 to storage show response we have processed 7 show response we have processed 8 show response we have processed 9
此外,q_get需要通过某个唯一对象来标识迭代终止。
内容的提问来源于stack exchange,提问作者LvLng

