sqlalchemy ORM异步


注意事项
sqlalchemy的查询结果都要使用 result.scalar()获取模型实列:
result.scalar() 单个对象
result.scalars().all() 列表形式
result.fetchone() 元组形式
result.to_dict() 字典
定义的模型都需要定义一个转化为字典的方法来传递给前端:
def to_dict(self):
return {
"id": self.id,
"name": self.name,
}

python
from sqlalchemy import Boolean, Column, DateTime, String, Integer, Text, Float, func, select
from sqlalchemy.ext.asyncio import create_async_engine, AsyncSession
from sqlalchemy.orm import sessionmaker, declarative_base
from datetime import datetime
import asyncio
from collections.abc import AsyncGenerator

# 创建异步engine和Base
DATABASE_URL = "sqlite+aiosqlite:///example_async.db"  # 使用aiosqlite作为异步驱动
# 对于PostgreSQL,可以使用:
# DATABASE_URL = "postgresql+asyncpg://user:password@localhost/dbname"

# 创建异步引擎
engine = create_async_engine(
    DATABASE_URL, 
    echo=True,  # 显示SQL语句,便于调试
    future=True  # 使用SQLAlchemy 2.0风格API
)

# 创建Base
Base = declarative_base()

# 创建异步会话工厂
async_session = sessionmaker(
    engine, 
    class_=AsyncSession, 
    expire_on_commit=False
)

# 定义模型
class User(Base):
    """
    用户模型
    """
    __tablename__ = "users"
    
    id = Column(Integer, primary_key=True)
    username = Column(String(50), nullable=False, unique=True)
    email = Column(String(100), nullable=False, unique=True)
    password_hash = Column(String(128), nullable=False)
    full_name = Column(String(100))
    bio = Column(Text)
    avatar_url = Column(String(255))
    is_active = Column(Boolean, default=True)
    is_admin = Column(Boolean, default=False)
    created_at = Column(DateTime, default=datetime.utcnow)
    updated_at = Column(DateTime, default=datetime.utcnow, onupdate=datetime.utcnow)
    login_count = Column(Integer, default=0)
    last_login = Column(DateTime)
    
    def __repr__(self):
        """
        对象的字符串表示
        """
        return f"<User(id={self.id}, username='{self.username}', email='{self.email}')>"
	# 自定义序列化方法,sqlalchemy的查询结果是不能直接传递给前端的需要在查询结果中调用此方法返回需要的字典	
	def to_dict(self): 
        return {
            "id": self.id,
            "name": self.name,
            "email": self.email
        }
# 异步上下文管理器,用于获取会话
async def get_session() -> AsyncGenerator[AsyncSession, None]:
    """
    创建异步数据库会话的上下文管理器
    
    @yields {AsyncSession} 异步数据库会话
    """
    async with async_session() as session:
		yield sesession #更建议下面这样
     #    try:
	    #     yield sesession
	    #     await session.commit()
	    # except Exception:
	    #     await session.rollback()
	    #     raise
	    # finally:
	    #     await session.close()

# 数据库操作函数
async def create_tables():
    """
    异步创建所有表
    """
    async with engine.begin() as conn:
        await conn.run_sync(Base.metadata.create_all)

async def drop_tables():
    """
    异步删除所有表
    """
    async with engine.begin() as conn:
        await conn.run_sync(Base.metadata.drop_all)

async def create_user(
    username: str, 
    email: str, 
    password_hash: str, 
    full_name: str | None = None, 
    bio: str | None = None, 
    avatar_url: str | None = None, 
    is_admin: bool = False
) -> User:
    """
    异步创建新用户
    
    @param {str} username - 用户名
    @param {str} email - 电子邮件
    @param {str} password_hash - 密码哈希值
    @param {str|None} full_name - 全名(可选)
    @param {str|None} bio - 简介(可选)
    @param {str|None} avatar_url - 头像URL(可选)
    @param {bool} is_admin - 是否为管理员(默认False)
    @returns {User} 创建的用户对象
    """
    async with async_session() as session:
        async with session.begin():
            # 创建用户对象
            new_user = User(
                username=username,
                email=email,
                password_hash=password_hash,
                full_name=full_name,
                bio=bio,
                avatar_url=avatar_url,
                is_admin=is_admin
            )
            
            # 添加到session
            session.add(new_user)
            
            # 提交会自动进行
            await session.flush()
            
            # 返回新创建的用户
            return new_user

async def get_user_by_id(user_id: int) -> User | None:
    """
    异步通过ID获取单个用户
    
    @param {int} user_id - 用户ID
    @returns {User|None} 用户对象,如果未找到则返回None
    """
    async with async_session() as session:
        # 查询单个用户
        result = await session.execute(
            select(User).where(User.id == user_id)
        )
        return result.scalars().first()

async def update_user(user_id: int, **kwargs) -> bool:
    """
    异步更新用户信息
    
    @param {int} user_id - 要更新的用户ID
    @param {dict} kwargs - 要更新的字段和值
    @returns {bool} 是否成功更新
    """
    async with async_session() as session:
        async with session.begin():
            # 查询用户
            result = await session.execute(
                select(User).where(User.id == user_id)
            )
            user = result.scalars().first()
            
            if not user:
                return False
                
            # 更新字段
            for key, value in kwargs.items():
                if hasattr(user, key):
                    setattr(user, key, value)
                    
            # 提交会自动进行
            return True

async def delete_user(user_id: int) -> bool:
    """
    异步删除用户
    
    @param {int} user_id - 要删除的用户ID
    @returns {bool} 是否成功删除
    """
    async with async_session() as session:
        async with session.begin():
            # 查询用户
            result = await session.execute(
                select(User).where(User.id == user_id)
            )
            user = result.scalars().first()
            
            if not user:
                return False
                
            # 删除用户
            await session.delete(user)
            return True

async def get_all_users() -> list[User]:
    """
    异步获取所有用户
    
    @returns {list} 所有用户对象的列表
    """
    async with async_session() as session:
        # 查询所有用户
        result = await session.execute(select(User))
        return result.scalars().all()

async def get_users_paginated(page: int = 1, per_page: int = 20) -> list[User]:
    """
    异步分页获取用户
    
    @param {int} page - 页码(从1开始)
    @param {int} per_page - 每页记录数
    @returns {list} 用户对象的列表
    """
    async with async_session() as session:
        # 计算偏移量
        offset = (page - 1) * per_page
        
        # 查询指定页的用户
        result = await session.execute(
            select(User).order_by(User.id).offset(offset).limit(per_page)
        )
        return result.scalars().all()

async def count_users() -> int:
    """
    异步统计用户总数
    
    @returns {int} 用户总数
    """
    async with async_session() as session:
        result = await session.execute(select(func.count(User.id)))
        return result.scalar()

async def get_active_users() -> list[User]:
    """
    异步获取所有激活状态的用户
    
    @returns {list} 激活状态的用户列表
    """
    async with async_session() as session:
        result = await session.execute(
            select(User).where(User.is_active == True)
        )
		# 使用.scalars().all()返回实列列表(直接的查询结果是不能使用的)
        return result.scalars().all()

async def update_login_info(user_id: int) -> None:
    """
    异步更新用户登录信息
    
    @param {int} user_id - 用户ID
    """
    async with async_session() as session:
        async with session.begin():
            result = await session.execute(
                select(User).where(User.id == user_id)
            )
			# 取实列(直接的查询结果是不能使用的)
            user = result.scalars().first()
            
            if user:
                user.login_count += 1
                user.last_login = datetime.utcnow()


# 示例用法
async def main():
    """
    异步主函数,演示所有功能
    """
    # 重新创建表
    await drop_tables()
    await create_tables()
    
    print("=== 创建用户示例 ===")
    # 创建示例用户
    user1 = await create_user(
        username="admin",
        email="admin@example.com",
        password_hash="hashed_password_123",
        full_name="管理员",
        is_admin=True
    )
    print(f"创建的用户: {user1}")
    
    # 批量创建多个用户用于演示
    for i in range(1, 30):
        await create_user(
            username=f"user{i}",
            email=f"user{i}@example.com",
            password_hash=f"hashed_password_{i}",
            full_name=f"用户{i}"
        )
    
    print("\n=== 查询单个用户示例 ===")
    # 获取指定ID的用户
    user = await get_user_by_id(1)
    print(f"ID为1的用户: {user}")
    
    print("\n=== 更新用户示例 ===")
    # 更新用户
    update_success = await update_user(
        1,
        bio="这是一个管理员账号",
        avatar_url="https://example.com/avatars/admin.jpg"
    )
    print(f"更新用户结果: {'成功' if update_success else '失败'}")
    
    # 显示更新后的用户
    updated_user = await get_user_by_id(1)
    print(f"更新后的用户: {updated_user}")
    print(f"用户简介: {updated_user.bio}")
    print(f"头像URL: {updated_user.avatar_url}")
    
    print("\n=== 模拟用户登录 ===")
    # 更新登录信息
    await update_login_info(1)
    user_after_login = await get_user_by_id(1)
    print(f"登录次数: {user_after_login.login_count}")
    print(f"最后登录时间: {user_after_login.last_login}")
    
    print("\n=== 查询所有用户示例 ===")
    # 获取用户总数
    total_users = await count_users()
    print(f"总用户数: {total_users}")
    
    # 获取所有用户 (谨慎使用,数据量大时会有性能问题)
    all_users = await get_all_users()
    print(f"所有用户数量: {len(all_users)}")
    
    print("\n=== 分页查询用户示例 ===")
    # 分页获取用户
    page1_users = await get_users_paginated(page=1, per_page=20)
    print(f"第1页用户数量: {len(page1_users)}")
    for i, user in enumerate(page1_users, 1):
        print(f"  {i}. {user.username} ({user.email})")
    
    page2_users = await get_users_paginated(page=2, per_page=20)
    print(f"\n第2页用户数量: {len(page2_users)}")
    for i, user in enumerate(page2_users, 1):
        print(f"  {i}. {user.username} ({user.email})")
    
    print("\n=== 删除用户示例 ===")
    # 删除用户
    delete_success = await delete_user(5)
    print(f"删除用户结果: {'成功' if delete_success else '失败'}")
    
    # 确认删除
    deleted_user = await get_user_by_id(5)
    print(f"ID为5的用户现在是: {deleted_user}")  # 应为None
    
    print("\n=== 获取激活用户示例 ===")
    # 禁用某些用户
    await update_user(2, is_active=False)
    await update_user(3, is_active=False)
    
    # 获取激活用户
    active_users = await get_active_users()
    print(f"激活用户数量: {len(active_users)}")
    

if __name__ == "__main__":
    # 运行异步主函数
    asyncio.run(main()) 

聚合管道查询

python
# sqlalchemy 管道聚合查询,随机获取一条热点事件
def get_one_hot_event(self, interests: list[str], persona_id: int) -> HotEventModel:
    """
    根据兴趣标签随机获取一条热点事件
    
    Args:
        interests (list[str]): 兴趣标签数组  
        persona_id (int): 人设ID
        
    Returns:
        HotEventModel: 随机匹配的热点事件,如果没有匹配的返回None
    """
    session = self.Session()
    try:
        # 方法1:使用ORDER BY RANDOM() LIMIT 1直接获取随机一条记录
        # 这种方法对于小到中等规模的数据集效率较高
        result = session.query(HotEventModel).filter(
            HotEventModel.events.op('&&')(interests),
            ~HotEventModel.id.in_(
                session.query(AdoptedModel.hot_id).filter(
                    AdoptedModel.persona_id == persona_id
                )
            )
        ).order_by(func.random()).limit(1).first()
        
        if result:
            print(f"为persona_id {persona_id} 随机选择热点事件: {result.title}")
            return result
        else:
            print(f"没有找到符合兴趣 {interests} 且未被persona_id {persona_id} 采用的热点事件")
            return None
        
    except Exception as e:
        print(f"随机获取热点事件失败: {e}")
        return None
    finally:
        session.close()
    
# postgres 原始sql高效聚合管道操作,随机获取一条热点事件
def get_one_hot_event_v2(self, interests: list[str], persona_id: int) -> HotEventModel:
    """
    根据兴趣标签随机获取一条热点事件,排除已被该人设采用过的文章
    使用PostgreSQL的CTE和随机抽样优化查询性能
    
    Args:
        interests (list[str]): 兴趣标签数组  
        persona_id (int): 人设ID
        
    Returns:
        HotEventModel: 随机匹配的热点事件,如果没有匹配的返回None
    """
    from sqlalchemy import func, text
    
    session = self.Session()
    try:
        # 使用SQLAlchemy构建带有CTE的SQL查询
        # 添加显式类型转换:CAST(:interests AS VARCHAR[])
        query = text("""
            WITH eligible_events AS (
                SELECT h.* FROM hot h
                WHERE h.events && CAST(:interests AS VARCHAR[])
                AND h.id NOT IN (
                    SELECT hot_id FROM adopted
                    WHERE persona_id = :persona_id
                )
            )
            SELECT * FROM eligible_events
            ORDER BY RANDOM()
            LIMIT 1
        """)
        
        # 执行查询
        result_proxy = session.execute(
            query, 
            {"interests": interests, "persona_id": persona_id}
        )
        
        # 获取一行结果
        result = result_proxy.fetchone()
        
        if result:
            # 正确处理RowProxy对象
            hot_event = HotEventModel()
            
            # 获取结果列名
            column_names = result_proxy.keys()
            
            # 遍历列名和值
            for idx, column_name in enumerate(column_names):
                if hasattr(hot_event, column_name):
                    # 使用索引获取值
                    setattr(hot_event, column_name, result[idx])
            
            print(f"为persona_id {persona_id} 随机选择热点事件: {hot_event.title}")
            return hot_event
        else:
            print(f"没有找到符合兴趣 {interests} 且未被persona_id {persona_id} 采用的热点事件")
            return None
        
    except Exception as e:
        print(f"随机获取热点事件失败: {e}")
        return None
    finally:
        session.close()

联级删除

python
# 热点事件
class HotEventModel(Base):
    __tablename__ = "hot"
    id = Column(Integer, primary_key=True)
    title = Column(String(500))
    content = Column(Text, nullable=False)  # 长文本,必填
    created_at = Column(DateTime(timezone=True), default=datetime.now)
    adopted_records = relationship("AdoptedModel", cascade="all, delete-orphan", back_populates="hot_event") #概念加反向关联

# 创建一个已发帖的映射表,设置为HotEventModel的外键,实现关联删除,如果HotEventModel被删除,则AdoptedModel也删除
class AdoptedModel(Base):
    __tablename__ = "adopted"
    
    id = Column(Integer, primary_key=True)
    # 数据库级别的级联删除
    hot_id = Column(Integer, ForeignKey("hot.id", ondelete='CASCADE'), nullable=False)
    persona_id = Column(Integer, ForeignKey("personas.id"), nullable=False)
    created_at = Column(DateTime(timezone=True), default=datetime.now)
    # 反向关系
    hot_event = relationship("HotEventModel", back_populates="adopted_records")