如何遍历SqlAlchemy实例列表/字典,实现重复运行仅插入新数据?
我原本以为这是一项简单的任务:使用flask-sqlalchemy创建一个可重复运行的种子数据脚本,仅在数据库中不存在对应数据时插入新值。
由于部分模型需要多步骤构建正确的模型和关系,我的思路是采用分步骤流程:用字典保存后续模型关系属性需要引用的实例声明,再编写commit_objects()函数,每次接收一个列表或字典,将其中的对象添加到session并提交,之后处理下一个字典或列表。以下是我尝试实现的示例代码:
def initial_seed(): # Step 1 - Objects with no dependencies ref_objects = {} ref_objects["tmro_cat_t"] = AssessmentCategory(name="Technology", abbreviation="T") ref_objects["rl1"] = AssessmentScore(score=1) ref_objects["rl_type_trl"] = AssessmentType(name="TRL", sort_order=2) commit_objects(ref_objects, "Importing Reference Objects - Step 1", True) #Step 2 - objects with dependencies from step 1 non_ref_objects_2 = [ AssessmentOption(assessment_score=ref_objects["rl1"], assessment_type=ref_objects["rl_type_trl"], name="Basic principles"), ActivityType(name="Program", abbreviation="PGM"), ] commit_objects(non_ref_objects_2, "Importing Non-Reference Objects - Step 2", True)
后续还需要4个左右的步骤来定义所有依赖关系和实例,以完成初始种子数据的构建,并可按需向列表或字典中添加更多数据。
我原本以为只需将列表或字典传入函数,判断对象是否已存在于数据库中,若所有属性完全匹配则跳过,否则插入即可。但实际操作中遇到了两个问题:
- 无法添加新对象并提交到session;
- 无法多次遍历列表识别已存在的对象并跳过。
以下是我目前编写的代码,但不确定是否方向正确:
def is_collection(attr): return hasattr(attr.property, "direction") and attr.property.direction.name in ("ONETOMANY", "MANYTOMANY") def record_exists(obj): # Get all columns for the object's table columns = [prop.key for prop in class_mapper(type(obj)).iterate_properties if isinstance(prop, sqlalchemy.orm.ColumnProperty)] # Create a filter for each column filters = {column: getattr(obj, column) for column in columns} # Check if a record with the same values already exists existing_obj = db.session.query(type(obj)).filter_by(**filters).first() return existing_obj is not None def commit_objects(objects_list, commit_message, debug=False): if isinstance(objects_list, dict): objects_list = objects_list.values() for obj in objects_list: if record_exists(obj): print("A record with the same values already exists. Skipping insertion.") continue try: # Save a copy of the object's attributes before making it transient inspector = inspect(type(obj)) # inspect the class, not the instance obj_attributes = {attr.key: getattr(obj, attr.key) for attr in inspector.attrs if not attr.key.startswith('_')} make_transient(obj) # Reassign the object's attributes after making it transient for key, value in obj_attributes.items(): setattr(obj, key, value) existing_obj = None if obj_attributes: query = db.session.query(type(obj)) for key, value in obj_attributes.items(): attr = getattr(type(obj), key) if key in [column.key for column in inspector.columns]: query = query.filter(attr == value) elif key in [relation.key for relation in inspector.relationships]: if is_collection(attr): related_obj = value if related_obj is not None and isinstance(related_obj, list) and len(related_obj) > 0: # check if object is a list before calling len() primary_key = inspector.primary_key[0].key query = query.filter(attr.any(**{primary_key: getattr(related_obj[0], primary_key)})) else: related_obj = value primary_key = inspector.primary_key[0].key query = query.filter(attr == value) else: print(f"Attribute '{key}' not found for {type(obj)}") continue existing_obj = query.first() if existing_obj: for key, value in obj_attributes.items(): setattr(existing_obj, key, value) if debug: print(f"Updated {type(obj).__name__} {getattr(existing_obj, inspector.primary_key[0].key)}") else: db.session.add(obj) if debug: print(f"Added {type(obj).__name__} {getattr(obj, inspector.primary_key[0].key)}") except Exception as e: print(f"Failed to commit {type(obj).__name__}: {str(e)}") db.session.rollback() continue try: db.session.commit() print(commit_message) except Exception as e: print(f"Failed to commit transaction: {str(e)}") db.session.rollback()
核心问题分析
你的代码逻辑过于复杂,尤其是make_transient的使用完全没必要,反而会破坏对象的session关联状态,导致无法正确添加到session。另外,record_exists函数用所有列作为判断条件并不合理——通常应该用业务唯一键(比如name+abbreviation)来判断记录是否存在,而不是所有字段(比如自增ID这类字段不应该参与判断)。
简化后的实现方案
我们可以重新设计两个核心函数:get_or_create(检查记录是否存在,不存在则创建)和commit_objects(批量处理对象),逻辑更清晰且避免不必要的操作。
1. 实现get_or_create函数
这个函数负责根据指定的唯一键查询记录,不存在则创建并返回实例:
from sqlalchemy.exc import IntegrityError def get_or_create(model, defaults=None, **kwargs): defaults = defaults or {} # 先根据唯一键查询 instance = db.session.query(model).filter_by(**kwargs).first() if instance: # 如果需要更新现有记录的其他字段 for key, value in defaults.items(): setattr(instance, key, value) return instance, False else: # 创建新实例 instance = model(**kwargs, **defaults) db.session.add(instance) try: db.session.commit() except IntegrityError: db.session.rollback() # 处理并发创建的情况 instance = db.session.query(model).filter_by(**kwargs).first() return instance, False return instance, True
2. 重构commit_objects函数
简化批量处理逻辑,支持字典和列表输入,同时处理依赖关系:
def commit_objects(objects_map, commit_message, debug=False): # 统一处理字典和列表 if isinstance(objects_map, dict): items = objects_map.items() else: items = enumerate(objects_map) for key, obj in items: # 提取业务唯一键(需要你根据每个模型定义,这里示例用常见字段) model = type(obj) unique_kwargs = {} defaults = {} # 针对不同模型定义唯一键,你可以扩展这个逻辑 if model == AssessmentCategory: unique_kwargs = {"name": obj.name, "abbreviation": obj.abbreviation} defaults = {"name": obj.name, "abbreviation": obj.abbreviation} elif model == AssessmentScore: unique_kwargs = {"score": obj.score} defaults = {"score": obj.score} elif model == AssessmentType: unique_kwargs = {"name": obj.name} defaults = {"name": obj.name, "sort_order": obj.sort_order} elif model == AssessmentOption: unique_kwargs = { "name": obj.name, "assessment_score_id": obj.assessment_score.id, "assessment_type_id": obj.assessment_type.id } defaults = { "name": obj.name, "assessment_score": obj.assessment_score, "assessment_type": obj.assessment_type } elif model == ActivityType: unique_kwargs = {"name": obj.name, "abbreviation": obj.abbreviation} defaults = {"name": obj.name, "abbreviation": obj.abbreviation} # 获取或创建实例 instance, created = get_or_create(model, defaults=defaults, **unique_kwargs) # 如果是字典,更新引用(确保后续步骤用的是数据库中存在的实例) if isinstance(objects_map, dict): objects_map[key] = instance if debug: if created: print(f"Added {model.__name__}: {instance}") else: print(f"Skipped {model.__name__} (already exists): {instance}") try: db.session.commit() print(commit_message) except Exception as e: db.session.rollback() print(f"Commit failed: {str(e)}")
3. 修改initial_seed函数
确保后续步骤使用的是数据库中已存在的实例:
def initial_seed(): # Step 1 - Objects with no dependencies ref_objects = {} ref_objects["tmro_cat_t"] = AssessmentCategory(name="Technology", abbreviation="T") ref_objects["rl1"] = AssessmentScore(score=1) ref_objects["rl_type_trl"] = AssessmentType(name="TRL", sort_order=2) commit_objects(ref_objects, "Importing Reference Objects - Step 1", True) # Step 2 - objects with dependencies from step 1 non_ref_objects_2 = [ AssessmentOption( assessment_score=ref_objects["rl1"], assessment_type=ref_objects["rl_type_trl"], name="Basic principles" ), ActivityType(name="Program", abbreviation="PGM"), ] commit_objects(non_ref_objects_2, "Importing Non-Reference Objects - Step 2", True)
关键改进点
- 去掉了不必要的
make_transient操作,避免破坏对象状态; - 使用业务唯一键而非全字段判断记录是否存在,逻辑更合理;
get_or_create函数处理并发创建的边界情况;- 更新字典中的实例引用,确保后续依赖步骤使用的是数据库中持久化的对象;
- 代码结构更清晰,易于扩展后续步骤。
内容的提问来源于stack exchange,提问作者Russ W.

