You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

SQLAlchemy单元测试:如何将内存会话传入模拟会话

问题分析与解决方案

你的测试始终返回生产数据库数据,核心原因有两点:

  1. 生产环境的engine在模块加载阶段就已初始化,User模型默认绑定到该生产引擎,即使替换Session,查询仍会指向生产库。
  2. 测试中对Session的Mock逻辑错误,未正确将查询引导至内存数据库,且测试数据未写入正确的会话。

最佳测试方案(分两种场景)

场景1:重构代码以优化测试友好性(推荐)

先改造app/data_structures/base.py,将数据库初始化逻辑封装为可配置函数,避免模块加载时直接绑定生产引擎:

from contextlib import contextmanager
from os import environ
from os.path import join, realpath

from sqlalchemy import Column, ForeignKey, Table, create_engine
from sqlalchemy.orm import declarative_base, scoped_session, sessionmaker

Base = declarative_base()

def init_db(db_url=None):
    """可配置的数据库初始化函数,测试时可传入内存库URL"""
    if db_url is None:
        db_name = environ.get("DB_NAME")
        ROOT_DIR = environ.get("ROOT_DIR")
        db_path = realpath(join(ROOT_DIR, "data", db_name))
        db_url = f"sqlite:///{db_path}"
        
    engine = create_engine(db_url, connect_args={"timeout": 120})
    session_factory = sessionmaker(bind=engine)
    sql_session = scoped_session(session_factory)
    
    @contextmanager
    def Session():
        session = sql_session()
        try:
            yield session
            session.commit()
        except Exception:
            session.rollback()
            raise
            
    return engine, Session

# 生产环境默认初始化
engine, Session = init_db()

对应的测试代码:

from app.data_structures.base import Base, init_db
from app.helpers.db_helper.get_users_to_query import get_users_to_query
from unittest.mock import patch

class TestSearchUser(UserHelperTestCase):
    def setUp(self):
        super().setUp()
        # 初始化内存数据库
        self.test_engine, self.test_Session = init_db("sqlite:///:memory:")
        Base.metadata.create_all(self.test_engine)
        
    def tearDown(self) -> None:
        # 测试后销毁表,保证用例隔离
        Base.metadata.drop_all(self.test_engine)
        return super().tearDown()
        
    def test_get_users_to_query(self) -> None:
        # 替换被测试函数中的Session为测试用Session
        with patch("app.helpers.db_helper.Session", self.test_Session):
            # 往测试数据库写入测试数据
            with self.test_Session() as session:
                fake_user1 = User(user_name="jdb1", flag=False)
                fake_user2 = User(user_name="jdb2", flag=True)
                session.add(fake_user1)
                session.add(fake_user2)
                session.commit()
            
            # 调用目标函数并验证结果
            result = get_users_to_query()
            assert result == [{"user_name": "jdb1"}]
            
            result_true = get_users_to_query(flag=True)
            assert result_true == [{"user_name": "jdb2"}]

场景2:不重构现有代码,直接修正测试逻辑

如果暂时无法修改生产代码,可通过重新绑定模型引擎+修正Mock逻辑解决:

from app.data_structures.base import Base, Session as prod_Session
from app.helpers.db_helper.get_users_to_query import get_users_to_query
from sqlalchemy import create_engine
from sqlalchemy.orm import sessionmaker
from unittest.mock import patch

class TestSearchUser(UserHelperTestCase):
    def setUp(self):
        super().setUp()
        # 创建内存引擎并绑定模型
        self.test_engine = create_engine("sqlite:///:memory:", connect_args={"timeout": 120})
        Base.metadata.create_all(self.test_engine)
        self.test_session_factory = sessionmaker(bind=self.test_engine)
        
    def tearDown(self) -> None:
        Base.metadata.drop_all(self.test_engine)
        return super().tearDown()
        
    @patch("app.data_structures.base.Session")
    def test_get_users_to_query(self, mock_session) -> None:
        # 让Mock的Session返回测试会话上下文
        def mock_session_context():
            session = self.test_session_factory()
            try:
                yield session
                session.commit()
            except Exception:
                session.rollback()
                raise
                
        mock_session.side_effect = mock_session_context
        
        # 写入测试数据到内存库
        with self.test_session_factory() as session:
            fake_user1 = User(user_name="jdb1", flag=False)
            fake_user2 = User(user_name="jdb2", flag=True)
            session.add(fake_user1)
            session.add(fake_user2)
            session.commit()
        
        # 调用函数并验证
        result = get_users_to_query()
        assert result == [{"user_name": "jdb1"}]

额外需要修正的代码问题

  1. User模型中flag: bool = false()是错误写法,Python中布尔值应写为False(大写)。
  2. 建议统一使用SQLAlchemy 2.0的Mapped+mapped_column语法,将flag = Column(Boolean)改为flag: Mapped[bool] = mapped_column(Boolean)。

内容的提问来源于stack exchange,提问作者abinitio

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.12 23:44:53