FastAPI函数内存占用过高求助:Heroku平台超1GB内存限制
在FastAPI中编写了如下函数,调用时应用内存占用持续增长,超出Heroku平台分配的1GB内存限制。该函数用于读写PostgreSQL表并调用其他服务器API。
program_records_model和ttmp_completion_model是基于declarative_base()的SQLAlchemy模型,db是SQLAlchemy Session实例,通过FastAPI的Depends()注入使用。get_user_info()和get_department_name()均为调用服务器API的函数。使用gunicorn搭配4个uvicorn worker在Heroku上运行应用。
async def update_ttmp_completion(bearer_token, program_records_model, ttmp_completion_model, db: Session): records_to_add = [] records_to_update = [] add_count = 0 update_count = 0 errors = [] # query the first table to get completion times and insert/update into the second table completed_learners = db.query(program_records_model.learner, program_records_model.enroll_at, program_records_model.finished_at, program_records_model.online_seconds).filter(program_records_model.finished_at != None).distinct(program_records_model.learner).all() # gets the distinct learner_id for all of those who have completed any variation of program. for i, (learner_id, enrollment_time, completion_time, online_seconds) in enumerate(completed_learners): if learner_id: existing_learner = db.query(ttmp_completion_model).filter(ttmp_completion_model.user_id == learner_id).first() if existing_learner: # logger.debug(msg=f'Completed learner {learner_id} already exists and will be updated.') existing_learner.enrollment_time = enrollment_time existing_learner.completion_time = completion_time existing_learner.online_seconds = online_seconds records_to_update.append(vars(existing_learner)) # vars() turns SQLAlchemy object into python dictionary that can be used by .bulk_update_mappings() method later. else: # logger.debug(msg=f'Completed learner {learner_id} is new and will be added.') try: learner_info = await get_user_info(bearer_token=bearer_token, user_id=learner_id, user_id_type='open_id') # if the learner does not exist we need to call API to get their information. name = learner_info['en_name'] manager_id = learner_info['manager_id'] email = learner_info['email'] department_id = learner_info['department_id'] department_name = await get_dep_name(bearer_token=bearer_token, department_id=department_id) except Exception as e: errors.append(e) continue # will skip to the next iteration of the for loop. new_record = ttmp_completion_model(user_id=learner_id, name=name, email=email, department_id=department_id, department_name=department_name, completion_time=completion_time, online_seconds=online_seconds, manager_id=manager_id, enrollment_time=enrollment_time) records_to_add.append(new_record) else: logger.error(msg='No learner id was provided for the completed learner!') if len(records_to_add) >= 50 or i == len(completed_learners) - 1: # we break the database updates to batches of n size so that a potential error does not hault the entire update process. try: if records_to_update: db.bulk_update_mappings(ttmp_completion_model, records_to_update) # unlike bulk_save_objects(), this method takes dictionaries. update_count += len(records_to_update) if records_to_add: db.bulk_save_objects(records_to_add) add_count += len(records_to_add) db.commit() logger.info(msg='50 or less new records were added.') records_to_add = [] records_to_update = [] except Exception as e: logger.error(msg=f'Error occurred adding to db: {e}') db.rollback() return {'new_record_count': add_count, 'upate_count': update_count, 'error_count': len(errors)}
已尝试在FastAPI路由中通过await调用该函数,也尝试过直接调用,但内存占用过高的问题仍未解决。
内存优化方案
1. 分批查询核心数据,避免一次性加载全量记录
当前db.query(...).all()会将所有符合条件的记录一次性加载到内存,数据量较大时直接占用大量内存。改用SQLAlchemy的yield_per()实现分批读取,每次从数据库获取固定数量的记录:
# 替换原completed_learners查询逻辑 completed_learners_query = db.query( program_records_model.learner, program_records_model.enroll_at, program_records_model.finished_at, program_records_model.online_seconds ).filter(program_records_model.finished_at != None).distinct(program_records_model.learner) # 每次读取100条(可根据内存情况调整批次大小) for i, (learner_id, enrollment_time, completion_time, online_seconds) in enumerate(completed_learners_query.yield_per(100)): # 后续循环逻辑保持不变
2. 清理SQLAlchemy Session缓存
SQLAlchemy Session会缓存所有查询过的对象,即使处理完成也不会自动释放。每批数据处理完成后,调用db.expire_all()释放缓存对象,必要时触发垃圾回收:
# 在每批commit或rollback后添加 db.expire_all() # 可选:手动触发垃圾回收 import gc gc.collect()
3. 优化更新逻辑,减少不必要的对象复制
原代码使用vars(existing_learner)将SQLAlchemy对象转为字典,会复制所有属性且Session仍持有原对象引用。直接构建更新字典,仅保留必要字段:
# 替换原更新逻辑 if existing_learner: update_data = { 'user_id': learner_id, 'enrollment_time': enrollment_time, 'completion_time': completion_time, 'online_seconds': online_seconds } records_to_update.append(update_data) # 批量更新时指定要更新的字段,提升效率 db.bulk_update_mappings(ttmp_completion_model, records_to_update, ['enrollment_time', 'completion_time', 'online_seconds'])
4. 批量查询用户存在性,减少单次查询次数
原循环中每次单独查询existing_learner,不仅效率低,还会让Session缓存大量对象。改为批量收集learner_id后一次性查询:
# 调整循环逻辑,批量处理用户 batch_size = 100 learner_batch = [] for i, (learner_id, enrollment_time, completion_time, online_seconds) in enumerate(completed_learners_query.yield_per(batch_size)): if learner_id: learner_batch.append((learner_id, enrollment_time, completion_time, online_seconds)) # 积累到指定批次或到末尾时批量处理 if len(learner_batch) >= batch_size or i == len(completed_learners_query) - 1: # 提取当前批次的所有learner_id learner_ids = [item[0] for item in learner_batch] # 批量查询已存在的用户ID existing_user_ids = { user.user_id for user in db.query(ttmp_completion_model.user_id).filter(ttmp_completion_model.user_id.in_(learner_ids)).all() } # 处理每个用户 for learner_id, enrollment_time, completion_time, online_seconds in learner_batch: if learner_id in existing_user_ids: # 构建更新数据 update_data = { 'user_id': learner_id, 'enrollment_time': enrollment_time, 'completion_time': completion_time, 'online_seconds': online_seconds } records_to_update.append(update_data) else: # 调用API获取用户信息并创建新记录,逻辑不变 try: learner_info = await get_user_info(bearer_token=bearer_token, user_id=learner_id, user_id_type='open_id') name = learner_info['en_name'] manager_id = learner_info['manager_id'] email = learner_info['email'] department_id = learner_info['department_id'] department_name = await get_dep_name(bearer_token=bearer_token, department_id=department_id) except Exception as e: errors.append(f"Learner {learner_id}: {type(e).__name__} - {str(e)}") continue new_record = ttmp_completion_model( user_id=learner_id, name=name, email=email, department_id=department_id, department_name=department_name, completion_time=completion_time, online_seconds=online_seconds, manager_id=manager_id, enrollment_time=enrollment_time ) records_to_add.append(new_record) # 执行批量更新/插入 try: if records_to_update: db.bulk_update_mappings(ttmp_completion_model, records_to_update, ['enrollment_time', 'completion_time', 'online_seconds']) update_count += len(records_to_update) if records_to_add: db.bulk_save_objects(records_to_add) add_count += len(records_to_add) db.commit() db.expire_all() # 清理缓存 records_to_add = [] records_to_update = [] learner_batch = [] except Exception as e: logger.error(msg=f'Error occurred adding to db: {e}') db.rollback()
5. 调整gunicorn worker数量
Heroku分配1GB内存,4个worker会平分内存(约250MB/worker),若单个worker处理任务时内存占用超限,可将worker数量调整为2:
# 修改Procfile中的启动命令 web: gunicorn main:app --workers 2 --worker-class uvicorn.workers.UvicornWorker
6. 优化错误记录,减少内存占用
原代码保存完整Exception对象会积累大量内存,改为记录关键错误信息即可:
# 替换原错误记录逻辑 except Exception as e: errors.append(f"Learner {learner_id}: {type(e).__name__} - {str(e)}") continue
内容的提问来源于stack exchange,提问作者Abhijeet

