数据库操作是后端开发最核心的部分之一。在Python中,直接写原生SQL虽然灵活,但在项目变得复杂后,ORM(对象关系映射)能帮我们节省大量时间。SQLAlchemy是Python生态中最强大的ORM框架,这篇文章不讲太深的理论,直接分享实际项目中最常用的操作和技巧。
一、为什么选择SQLAlchemy
Python的ORM有好几个选择,但SQLAlchemy是公认的"工业级"解决方案。
SQLAlchemy的核心优势:
-
支持多种数据库(PostgreSQL、MySQL、SQLite、Oracle等)
-
提供了两种使用方式:Core(SQL表达式)和ORM(对象映射)
-
性能优秀,SQL生成高效
-
连接池管理完善
-
支持异步(1.4版本+)
对比其他ORM:
| 独立使用 | ✅ | ❌(需要Django) | ✅ |
| 多数据库支持 | 完整 | 基本支持 | 基本支持 |
| SQL生成能力 | 极强 | 中等 | 中等 |
| 学习曲线 | 陡峭 | 平缓 | 平缓 |
| 灵活性 | 极高 | 中等 | 中等 |
二、基础:模型定义和连接配置
首先安装:
bash
pip install sqlalchemy
# 根据使用的数据库安装对应驱动
pip install psycopg2-binary # PostgreSQL
pip install pymysql # MySQL
pip install aiomysql # MySQL异步驱动
基础配置
python
from sqlalchemy import create_engine, Column, Integer, String, DateTime, Boolean, Text
from sqlalchemy.ext.declarative import declarative_base
from sqlalchemy.orm import sessionmaker, relationship
from datetime import datetime
# 创建数据库引擎
# 格式: 数据库类型://用户名:密码@主机:端口/数据库名
engine = create_engine(
'postgresql://user:password@localhost:5432/mydb',
echo=True, # 打印生成的SQL,调试时很有用
pool_size=10, # 连接池大小
max_overflow=20, # 连接池最大溢出数
pool_recycle=3600, # 连接回收时间(秒)
pool_pre_ping=True, # 使用前检测连接是否有效
)
# 创建基类
Base = declarative_base()
# 创建会话工厂
SessionLocal = sessionmaker(autocommit=False, autoflush=False, bind=engine)
定义模型
python
from sqlalchemy import Column, Integer, String, DateTime, ForeignKey, Float, Text, Boolean
from sqlalchemy.orm import relationship
from datetime import datetime
class User(Base):
__tablename__ = 'users'
id = Column(Integer, primary_key=True, index=True)
username = Column(String(50), unique=True, nullable=False, index=True)
email = Column(String(100), unique=True, nullable=False)
password_hash = Column(String(200), nullable=False)
full_name = Column(String(100))
age = Column(Integer)
is_active = Column(Boolean, default=True)
created_at = Column(DateTime, default=datetime.utcnow)
updated_at = Column(DateTime, default=datetime.utcnow, onupdate=datetime.utcnow)
# 关系
orders = relationship("Order", back_populates="user", cascade="all, delete-orphan")
def __repr__(self):
return f"<User(id={self.id}, username={self.username})>"
class Product(Base):
__tablename__ = 'products'
id = Column(Integer, primary_key=True, index=True)
name = Column(String(200), nullable=False)
description = Column(Text)
price = Column(Float, nullable=False)
stock = Column(Integer, default=0)
category = Column(String(50))
created_at = Column(DateTime, default=datetime.utcnow)
class Order(Base):
__tablename__ = 'orders'
id = Column(Integer, primary_key=True, index=True)
user_id = Column(Integer, ForeignKey('users.id'), nullable=False)
product_id = Column(Integer, ForeignKey('products.id'), nullable=False)
quantity = Column(Integer, nullable=False)
total_amount = Column(Float, nullable=False)
status = Column(String(20), default='pending') # pending, paid, shipped, completed, cancelled
created_at = Column(DateTime, default=datetime.utcnow)
# 关系
user = relationship("User", back_populates="orders")
product = relationship("Product")
# 创建表
Base.metadata.create_all(engine)
三、CRUD操作:最常用的增删改查
创建记录
python
from sqlalchemy.orm import Session
def create_user(db: Session, username: str, email: str, password_hash: str):
"""创建用户"""
user = User(
username=username,
email=email,
password_hash=password_hash,
created_at=datetime.utcnow()
)
db.add(user)
db.commit()
db.refresh(user) # 刷新获取生成的id
return user
# 批量创建
def create_users_bulk(db: Session, users_data: list):
"""批量创建用户"""
users = [User(**data) for data in users_data]
db.add_all(users)
db.commit()
return users
# 使用
with SessionLocal() as db:
user = create_user(db, "zhangsan", "zhangsan@example.com", "hashed_password")
print(user.id)
查询记录
python
from sqlalchemy import and_, or_, not_, desc, func
def get_user_by_id(db: Session, user_id: int):
"""根据ID获取用户"""
return db.query(User).filter(User.id == user_id).first()
def get_user_by_username(db: Session, username: str):
"""根据用户名获取用户"""
return db.query(User).filter(User.username == username).first()
def get_active_users(db: Session, skip: int = 0, limit: int = 100):
"""获取活跃用户列表"""
return db.query(User).filter(User.is_active == True)\\
.offset(skip).limit(limit).all()
def search_users(db: Session, keyword: str):
"""搜索用户(模糊匹配)"""
return db.query(User).filter(
or_(
User.username.like(f'%{keyword}%'),
User.email.like(f'%{keyword}%'),
User.full_name.like(f'%{keyword}%')
)
).all()
def get_users_with_orders(db: Session):
"""获取有订单的用户(使用join)"""
return db.query(User).join(Order).distinct().all()
def get_user_statistics(db: Session):
"""获取用户统计信息(聚合查询)"""
total_users = db.query(func.count(User.id)).scalar()
active_users = db.query(func.count(User.id)).filter(User.is_active == True).scalar()
avg_age = db.query(func.avg(User.age)).scalar()
return {
'total': total_users,
'active': active_users,
'avg_age': avg_age or 0
}
更新记录
python
def update_user(db: Session, user_id: int, **kwargs):
"""更新用户信息"""
user = db.query(User).filter(User.id == user_id).first()
if user:
for key, value in kwargs.items():
if hasattr(user, key):
setattr(user, key, value)
user.updated_at = datetime.utcnow()
db.commit()
db.refresh(user)
return user
def update_users_bulk(db: Session, user_ids: list, updates: dict):
"""批量更新用户"""
db.query(User).filter(User.id.in_(user_ids)).update(
updates,
synchronize_session=False
)
db.commit()
# 使用
with SessionLocal() as db:
# 单个更新
user = update_user(db, 1, full_name="张三丰", age=30)
# 批量更新
update_users_bulk(db, [1, 2, 3], {'is_active': False})
删除记录
python
def delete_user(db: Session, user_id: int):
"""删除用户(软删除:标记为不活跃)"""
user = db.query(User).filter(User.id == user_id).first()
if user:
user.is_active = False
db.commit()
return True
return False
def delete_user_permanent(db: Session, user_id: int):
"""永久删除用户"""
user = db.query(User).filter(User.id == user_id).first()
if user:
db.delete(user)
db.commit()
return True
return False
def delete_inactive_users(db: Session):
"""批量删除不活跃用户"""
deleted_count = db.query(User).filter(User.is_active == False).delete()
db.commit()
return deleted_count
四、高级查询技巧
复杂条件查询
python
from sqlalchemy import and_, or_, not_, between, in_
def complex_query(db: Session, filters: dict):
"""复杂条件查询"""
query = db.query(User)
# 组合条件
conditions = []
if filters.get('min_age'):
conditions.append(User.age >= filters['min_age'])
if filters.get('max_age'):
conditions.append(User.age <= filters['max_age'])
if filters.get('username_like'):
conditions.append(User.username.like(f"%{filters['username_like']}%"))
if filters.get('active_only'):
conditions.append(User.is_active == True)
if filters.get('exclude_ids'):
conditions.append(not_(User.id.in_(filters['exclude_ids'])))
if conditions:
query = query.filter(and_(*conditions))
return query.all()
排序和分页
python
from sqlalchemy import desc, asc
def paginated_query(db: Session, page: int = 1, page_size: int = 20, sort_by: str = 'id', sort_order: str = 'desc'):
"""分页查询"""
# 计算偏移量
offset = (page – 1) * page_size
# 排序方向
order_func = desc if sort_order == 'desc' else asc
# 获取排序字段
sort_field = getattr(User, sort_by, User.id)
# 查询
query = db.query(User).order_by(order_func(sort_field))
# 获取总数
total = query.count()
# 获取当前页数据
items = query.offset(offset).limit(page_size).all()
return {
'items': items,
'total': total,
'page': page,
'page_size': page_size,
'total_pages': (total + page_size – 1) // page_size
}
子查询和复杂查询
python
from sqlalchemy import func, select
def get_users_with_high_value_orders(db: Session, min_amount: float = 1000):
"""获取有高额订单的用户(子查询)"""
subquery = db.query(Order.user_id).filter(Order.total_amount > min_amount).subquery()
return db.query(User).filter(User.id.in_(subquery)).all()
def get_order_statistics(db: Session):
"""订单统计(分组查询)"""
result = db.query(
User.username,
func.count(Order.id).label('order_count'),
func.sum(Order.total_amount).label('total_amount'),
func.avg(Order.total_amount).label('avg_amount')
).join(Order, User.id == Order.user_id)\\
.group_by(User.id, User.username)\\
.order_by(desc('total_amount'))\\
.all()
return result
五、事务管理
python
from sqlalchemy.exc import IntegrityError, SQLAlchemyError
def transfer_order(db: Session, from_user_id: int, to_user_id: int, order_id: int):
"""转移订单(事务示例)"""
try:
# 开始事务(with块自动管理)
with db.begin():
# 获取订单
order = db.query(Order).filter(Order.id == order_id).first()
if not order:
raise ValueError("订单不存在")
# 更新订单所属用户
order.user_id = to_user_id
# 记录操作日志(假设有日志表)
# log = OperationLog(…)
# db.add(log)
# 所有操作都成功才提交
# with块结束时会自动commit
return True
except IntegrityError as e:
db.rollback()
print(f"数据完整性错误: {e}")
return False
except SQLAlchemyError as e:
db.rollback()
print(f"数据库错误: {e}")
return False
# 更复杂的嵌套事务
def complex_transaction(db: Session):
"""复杂事务示例"""
try:
# 方式1:显式管理
db.begin_nested() # 保存点
try:
# 执行一些操作
user = db.query(User).filter(User.id == 1).first()
user.is_active = False
# 如果这里出错,只会回滚到这个保存点
db.begin_nested()
try:
# 更细粒度的操作
order = db.query(Order).filter(Order.id == 1).first()
order.status = 'cancelled'
except Exception:
db.rollback() # 回滚到第二个保存点
except Exception:
db.rollback() # 回滚到第一个保存点
else:
db.commit()
except Exception as e:
db.rollback()
print(f"事务失败: {e}")
六、性能优化技巧
1. 使用selectinload避免N+1查询
python
from sqlalchemy.orm import selectinload, joinedload
# ❌ 错误方式:N+1查询
def get_orders_naive(db: Session):
orders = db.query(Order).all()
for order in orders:
# 每次访问order.user都会触发一次查询
print(order.user.username) # N次额外查询
# ✅ 正确方式:预加载关联数据
def get_orders_optimized(db: Session):
orders = db.query(Order).options(
selectinload(Order.user), # 一次查询加载所有用户
joinedload(Order.product) # 使用JOIN加载商品
).all()
for order in orders:
print(order.user.username) # 不会触发额外查询
2. 只查询需要的字段
python
# ❌ 查询所有字段
def get_all_users_fields(db: Session):
return db.query(User).all()
# ✅ 只查询需要的字段
def get_user_names_only(db: Session):
return db.query(User.id, User.username, User.email).all()
# 返回字典而不是对象
def get_user_dicts(db: Session):
return db.query(User.id, User.username).all() # 返回命名元组列表
3. 使用批量操作
python
# ❌ 逐条插入
def insert_users_one_by_one(db: Session, users_data: list):
for data in users_data:
user = User(**data)
db.add(user)
db.commit()
# ✅ 批量插入
def insert_users_bulk(db: Session, users_data: list):
db.bulk_insert_mappings(User, users_data)
db.commit()
4. 合理使用索引
python
from sqlalchemy import Index
# 在模型上定义索引
class User(Base):
__tablename__ = 'users'
id = Column(Integer, primary_key=True)
username = Column(String(50))
email = Column(String(100))
age = Column(Integer)
city = Column(String(50))
# 复合索引
__table_args__ = (
Index('idx_username_age', 'username', 'age'),
Index('idx_city_email', 'city', 'email'),
)
七、异步SQLAlchemy
SQLAlchemy 1.4+ 支持异步操作,配合FastAPI非常实用:
python
from sqlalchemy.ext.asyncio import create_async_engine, AsyncSession
from sqlalchemy.orm import sessionmaker
from sqlalchemy import select
# 创建异步引擎
async_engine = create_async_engine(
'postgresql+asyncpg://user:password@localhost:5432/mydb',
echo=True,
pool_size=10,
)
# 异步会话工厂
AsyncSessionLocal = sessionmaker(
async_engine,
class_=AsyncSession,
expire_on_commit=False,
)
# 异步CRUD
async def get_user_async(user_id: int):
async with AsyncSessionLocal() as db:
result = await db.execute(
select(User).where(User.id == user_id)
)
return result.scalar_one_or_none()
async def create_user_async(user_data: dict):
async with AsyncSessionLocal() as db:
user = User(**user_data)
db.add(user)
await db.commit()
await db.refresh(user)
return user
# 在FastAPI中使用
from fastapi import FastAPI
app = FastAPI()
@app.get("/users/{user_id}")
async def get_user(user_id: int):
user = await get_user_async(user_id)
if not user:
return {"error": "User not found"}
return {"id": user.id, "username": user.username}
八、常见问题与解决方案
1. 连接池耗尽
python
# 增加连接池大小
engine = create_engine(
'postgresql://…',
pool_size=20, # 基础连接数
max_overflow=30, # 最大额外连接数
pool_pre_ping=True, # 自动重连
)
2. 数据库连接超时
python
# 设置连接超时和回收
engine = create_engine(
'postgresql://…',
connect_args={'connect_timeout': 10},
pool_recycle=3600, # 一小时后回收连接
)
3. 查询性能慢
使用explain()查看执行计划:
python
from sqlalchemy import text
def explain_query(db: Session):
query = db.query(User).filter(User.age > 18)
# 打印SQL
print(str(query))
# 执行EXPLAIN
result = db.execute(
text(f"EXPLAIN ANALYZE {str(query)}")
)
for row in result:
print(row)
总结
SQLAlchemy功能强大,但学习曲线确实陡峭。刚入门的朋友可以从核心功能开始,逐步深入。
学习建议:
先熟悉CRUD基础操作
掌握查询过滤和排序
学会处理关系映射
了解性能优化技巧
再研究高级特性
记住:SQLAlchemy的目标不是让你忘记SQL,而是让SQL操作更安全、更高效。在复杂查询时,有时候直接写原生SQL反而更清晰。