如何用aiohttp和aiofiles实现生产者-消费者模型?附代码求建议
代码改进建议
原代码的结构思路没问题,但有几个核心细节需要修正,同时可以做一些优化来提升稳定性和效率,具体改进如下:
1. 补全生产者的参数与异步操作
原生产者没传入session、url、queue这些必要参数,根本无法正常运行;另外r.read()是异步方法必须加await,还要增加响应状态校验,避免把无效的错误响应内容塞进队列。
修改后的生产者:
async def producer(session: aiohttp.ClientSession, url: str, queue: asyncio.Queue): try: async with session.get(url) as r: r.raise_for_status() # 遇到HTTP错误直接抛出异常 content = await r.read() # 异步读取响应内容 await queue.put((url, content)) # 同时传入url和内容,方便消费者命名文件 except Exception as e: print(f"下载{url}失败: {str(e)}")
2. 让消费者循环处理任务直到收到停止信号
原消费者仅执行一次就结束,无法处理大量文件。需要改成循环从队列取任务,直到收到None这类停止标记;同时要确保无论写入成功或失败,都调用queue.task_done(),避免队列阻塞。
修改后的消费者:
async def consumer(queue: asyncio.Queue, output_dir: str): await aiofiles.os.makedirs(output_dir, exist_ok=True) # 确保输出目录存在,避免报错 while True: item = await queue.get() if item is None: # 收到停止信号,退出循环 queue.task_done() break url, content = item filename = url.split("/")[-1] # 从URL提取文件名 file_path = os.path.join(output_dir, filename) try: async with aiofiles.open(file_path, "wb") as f: await f.write(content) print(f"已保存文件: {file_path}") except Exception as e: print(f"写入{file_path}失败: {str(e)}") finally: queue.task_done() # 无论结果如何,标记任务完成
3. 优化主流程的任务调度与协同逻辑
- 先启动消费者任务,避免生产者先塞满队列导致内存压力过大
- 给队列设置最大容量,限制内存占用
- 所有生产者任务完成后,向队列发送与消费者数量相等的停止信号,通知消费者退出
- 调整
await顺序,确保所有任务正常收尾
修改后的主函数:
async def main(urls: list[str], number_of_consumers: int, output_dir: str = "./downloads"): queue = asyncio.Queue(maxsize=10) # 设置队列最大容量,防止内存溢出 tasks = [] # 先启动所有消费者任务 for _ in range(number_of_consumers): consumer_task = asyncio.create_task(consumer(queue, output_dir)) tasks.append(consumer_task) async with aiohttp.ClientSession() as session: # 启动生产者任务 for url in urls: producer_task = asyncio.create_task(producer(session, url, queue)) tasks.append(producer_task) # 等待所有生产者任务完成 await asyncio.gather(*[t for t in tasks if "producer" in str(t)], return_exceptions=True) # 向每个消费者发送停止信号 for _ in range(number_of_consumers): await queue.put(None) # 等待队列所有任务处理完成,再等待消费者退出 await queue.join() await asyncio.gather(*tasks, return_exceptions=True)
额外实用优化建议
- 限制并发请求数:aiohttp默认无严格并发限制,容易被目标服务器拒绝服务,可通过
TCPConnector控制并发:connector = aiohttp.TCPConnector(limit=15) # 限制同时15个并发请求 async with aiohttp.ClientSession(connector=connector) as session: # 生产者逻辑 - 避免文件名重复:用URL最后一段命名可能重复,可结合哈希生成唯一文件名:
import hashlib ext = os.path.splitext(url.split("/")[-1])[-1] filename = hashlib.md5(url.encode()).hexdigest() + ext - 添加进度监控:用计数器跟踪下载和完成数量,方便掌握进度:
# 主函数中定义锁和计数器 progress_lock = asyncio.Lock() downloaded_count = 0 completed_count = 0 # 生产者下载成功后更新 async with progress_lock: nonlocal downloaded_count downloaded_count +=1 print(f"已下载 {downloaded_count}/{len(urls)} 个文件") # 消费者写入成功后更新 async with progress_lock: nonlocal completed_count completed_count +=1 print(f"已完成 {completed_count}/{len(urls)} 个文件")
内容的提问来源于stack exchange,提问作者Round
相关产品推荐
相关产品推荐

