如何在SQLAlchemy同级添加列并适配FastAPI的Pydantic模型?
问题
我正在使用FastAPI和SQLAlchemy,希望返回包含User及其关联部门名称的查询结果。当前通过session.query(User).join(Department).add_column(Department.name.label('department_name'))查询后,返回的是嵌套结构:
[{ "User": { "name": "John", "age": 20 }, "department_name": "Apple" }]
我需要改成字段同级的结构:
[{ "name": "John", "age": 20, "department_name": "Apple" }]
直接指定UserDetail作为响应模型会触发ValidationError,而且我不想显式查询User的每个字段(因为模型可能变更),希望能自动加载所有User字段。
相关代码如下:
SQLAlchemy 模型
class User(Base): id = Column(Integer, primary_key=True) name = Column(String) age = Column(Integer) department_id = Column(Integer, ForeignKey("department.id"), nullable=False) class Department(Base): id = Column(Integer, primary_key=True) name = Column(String)
Pydantic 模型
class UserDetail(BaseModel): name: str age: int department_name: str class Config: orm_mode = True
FastAPI 路径操作函数
@app.get("/users") def read_users(): q = session.query(User).join(Department).set_label_style(LABEL_STYLE_DISAMBIGUATE_ONLY) # 设置标签样式:使用不带表名的列名 q.add_column(Department.name.label('department_name')) # 我想要添加的列 return q.all()
解决方案
方法一:将查询结果转换为扁平字典
修改路径操作函数,自动获取User的所有字段并与部门名称合并成扁平结构:
from sqlalchemy.orm import class_mapper @app.get("/users", response_model=list[UserDetail]) def read_users(): q = session.query(User, Department.name.label('department_name')).join(Department) results = q.all() flat_results = [] for user, dept_name in results: # 自动获取User所有列并转成字典 user_dict = {c.key: getattr(user, c.key) for c in class_mapper(User).columns} # 合并部门名称字段 user_dict['department_name'] = dept_name flat_results.append(user_dict) return flat_results
通过class_mapper(User).columns动态获取User的所有字段,后续模型变更无需修改这段代码。
方法二:重写Pydantic的from_orm方法适配元组结果
修改Pydantic模型,让它能自动解析查询返回的(User, department_name)元组:
class UserDetail(BaseModel): id: int name: str age: int department_name: str class Config: orm_mode = True @classmethod def from_orm(cls, obj): if isinstance(obj, tuple): user, dept_name = obj return cls(**user.__dict__, department_name=dept_name) return super().from_orm(obj)
然后调整路径操作函数即可直接返回查询结果:
@app.get("/users", response_model=list[UserDetail]) def read_users(): q = session.query(User, Department.name.label('department_name')).join(Department) return q.all()
方法三:直接查询所有列返回扁平映射
通过SQLAlchemy的select语句查询User的所有列加部门名称,返回字典格式的结果:
from sqlalchemy import select @app.get("/users", response_model=list[UserDetail]) def read_users(): # 自动获取User的所有列 user_columns = [User.__table__.c[col] for col in User.__table__.columns.keys()] q = select(*user_columns, Department.name.label('department_name')).join(Department) # 返回字典格式的查询结果 return session.execute(q).mappings().all()
内容的提问来源于stack exchange,提问作者Baekjun Kim
相关产品推荐
相关产品推荐

