一、SQLAlchemy引入
1、安装
python
pip install sqlalchemy
2、在项目的配置类
python
import os
from dotenv import load_dotenv
from sqlalchemy.ext.asyncio import async_sessionmaker, AsyncSession, create_async_engine
load_dotenv()
# 数据库URL
ASYNC_DATABASE_URL = os.getenv('ASYNC_DATABASE_URL')
POOL_SIZE = int(os.getenv('POOL_SIZE'))
MAX_OVERFLOW = int(os.getenv('MAX_OVERFLOW'))
POOL_RECYCLE = int(os.getenv('POOL_RECYCLE'))
POOL_TIMEOUT = int(os.getenv('POOL_TIMEOUT'))
SQL_ECHO = bool(os.getenv('SQL_ECHO'))
# 创建异步引擎
async_engine = create_async_engine(
ASYNC_DATABASE_URL = ASYNC_DATABASE_URL,
echo = SQL_ECHO, # 输出SQL日志
pool_size = POOL_SIZE, # 设置连接池中保持的持久连接数
max_overflow = MAX_OVERFLOW, # 设置连接池允许创建的额外连接数
pool_timeout = POOL_TIMEOUT,
pool_recycle = POOL_RECYCLE
)
# 创建异步会话工厂
AsyncSessionLocal = async_sessionmaker(
bind=async_engine,
class_=AsyncSession,
expire_on_commit=False
)
# 依赖项,用于获取数据库会话
async def get_db():
async with AsyncSessionLocal() as session:
try:
yield session
await session.commit()
except Exception:
await session.rollback()
raise
finally:
await session.close()
二、使用SQLAlchemy
1、实体类定义
例如:一个User实体类
- Optionalstr表示字段可以为字符串也可以为None
python
from datetime import datetime
from typing import Optional
from sqlalchemy import Index, Integer, String, Enum, DateTime, ForeignKey
from sqlalchemy.orm import DeclarativeBase, Mapped, mapped_column
class Base(DeclarativeBase):
pass
class User(Base):
"""
用户信息表ORM模型
"""
__tablename__ = 'user'
# 创建索引
__table_args__ = (
Index('username_UNIQUE', 'username'),
Index('phone_UNIQUE', 'phone'),
)
id: Mapped[int] = mapped_column(Integer, primary_key=True, autoincrement=True, comment="用户ID")
username: Mapped[str] = mapped_column(String(50), unique=True, nullable=False, comment="用户名")
password: Mapped[str] = mapped_column(String(255), nullable=False, comment="密码(加密存储)")
nickname: Mapped[Optional[str]] = mapped_column(String(50), comment="昵称")
avatar: Mapped[Optional[str]] = mapped_column(String(255), comment="头像URL", default='https://fastly.jsdelivr.net/npm/@vant/assets/cat.jpeg')
gender: Mapped[Optional[str]] = mapped_column(Enum('male', 'female', 'unknown'), comment="性别", default='unknown')
bio: Mapped[Optional[str]] = mapped_column(String(500), comment="个人简介", default='这个人很懒,什么都没留下!')
phone: Mapped[Optional[str]] = mapped_column(String(20), unique=True, comment="手机号")
created_at: Mapped[datetime] = mapped_column(DateTime, default=datetime.now(), comment="创建时间")
updated_at: Mapped[datetime] = mapped_column(DateTime, default=datetime.now(), onupdate=datetime.now(), comment="更新时间")
#在开发/测试时,可以将对象转为字符串输出,方便调试
def __repr__(self):
return f"<User(id={self.id}, username='{self.username}', nickname='{self.nickname}')>"
三、使用SQLAlchemy增删改查
- 使用commit()将已经保存的对象提交到数据库
- 使用rollback()撤销需要提交的对象
- 使用refresh()用数据库中最新的数据,覆盖掉当前对象在内存中的值
1、增
两种方式:
- 使用add()/add_all()将需要添加的对象进行保存,直接commit即可
- 使用insert后,执行execute,最终commit
python
async def create_user(db: AsyncSession, username: str, password: str):
#先密码加密处理
hashed_password = get_hash_password(password)
#创建用户
user = User(username=username, password=hashed_password)
db.add(user)
await db.commit()
await db.refresh(user)
return user
2、删
- 使用delete()
python
async def delete_user(db: AsyncSession, user_id: int):
stmt = delete(User).where(User.id == user_id)
result = await db.execute(stmt)
deleted_rows = result.rowcount
await db.commit()
return deleted_rows
3、改
- update
python
#更新用户信息的模型类
class UserUpdateRequest(BaseModel):
nickname: str = None # type: ignore
avatar: str = None # type: ignore
gender: str = None # type: ignore
bio: str = None # type: ignore
phone: str = None # type: ignore
#更新用户方法
async def update_user(db: AsyncSession, user_name:str, user_data:UserUpdateRequest):
#update(User).where(User.username == user_name).values(字段=值,字段=值)
#user_data是一个pydantic模型类,不能直接用在update里,需要转换成字典->**解包
query = update(User).where(User.username == user_name).values(**user_data.model_dump(
exclude_unset=True,
exclude_none=True #没有设置值的不更新
))
result = await db.execute(query)
await db.commit()
return result.rowcount
4、查
- select,查询语句类似MySQL
python
# SQLAlchemy 方法链顺序
stmt = (
select(column1, column2) # 1. SELECT
.select_from(table1) # 2. FROM(可选,通常自动推断)
.join(table2, condition) # 3. JOIN
.where(condition) # 4. WHERE
.group_by(column) # 5. GROUP BY
.having(group_condition) # 6. HAVING
.order_by(column) # 7. ORDER BY
.limit(n) # 8. LIMIT
.offset(m) # 9. OFFSET
)
四、原子性
- 使用begin()来确保多个sql语句执行时的原子性
python
from sqlalchemy import select, update
from sqlalchemy.ext.asyncio import AsyncSession
async def transfer_points(db: AsyncSession, from_user_id: int, to_user_id: int, points: int):
"""原子性地将积分从用户A转移到用户B"""
# 使用 async with 开启事务
async with db.begin():
# 1. 扣减转出方的积分(先查询再更新,确保业务逻辑正确)
stmt_from = (
update(User)
.where(User.id == from_user_id, User.points >= points) # 加条件防止负积分
.values(points=User.points - points)
)
result_from = await db.execute(stmt_from)
# 如果转出方不存在或积分不足,affected_rows 为 0
if result_from.rowcount == 0:
raise ValueError("转出用户不存在或积分不足")
# 2. 增加转入方的积分
stmt_to = (
update(User)
.where(User.id == to_user_id)
.values(points=User.points + points)
)
result_to = await db.execute(stmt_to)
if result_to.rowcount == 0:
raise ValueError("转入用户不存在")
# 3. 记录积分变动日志(第三个表)
log_entry = PointsLog(
from_user_id=from_user_id,
to_user_id=to_user_id,
points=points,
created_at=func.now()
)
db.add(log_entry)
# 注意:这里不需要手动 commit()
# 退出 with 块时,自动 commit()
# 如果抛异常,自动 rollback()
# 调用示例
try:
await transfer_points(db, 1, 2, 100)
print("积分转移成功!")
except ValueError as e:
print(f"业务失败: {e}")
# 事务已经自动回滚,无需额外操作