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

Python单元测试:如何Mock PostgreSQL连接字符串?

问题场景与报错

我在main.py中实现了connect_with_postgres方法,用于将DataFrame写入PostgreSQL,代码如下:

def connect_with_postgres(df_table,sql_conn):
    db = create_engine(sql_conn)
    conn=db.connect()
    #build a table and put the dataframe information into
    df_table.to_sql('table_name', con=conn, if_exists='replace', schema='schemaname', index=False)
    conn.close()
    print("The data has been stored into the database.")

单元测试中读取mocks/table.csv作为测试数据:

with open('mocks/table.csv', 'r') as file:
    mock_table_csv = file.read()
    file.close()

测试函数如下:

def test_connect_with_postgres(self):
    conn_result = Mock()
    mock_conn.connect.return_value = conn_result
    mock_table = pd.read_csv(StringIO(mock_table_csv))
    output = connect_with_postgres(mock_table,conn)

运行测试时抛出错误:

TypeError: cannot unpack non-iterable Mock object

使用真实连接字符串测试正常,但希望通过Mock虚假连接完成测试,该如何处理?


解决方法

原方法通过create_engine(sql_conn)创建引擎对象,再调用其connect()方法获取连接,直接Mock连接对象无法覆盖完整调用链路,导致报错。正确的做法是Mockcreate_engine函数,构建完整的Mock调用层级:

修改后的测试代码

from unittest.mock import Mock, patch
import pandas as pd
import unittest
from io import StringIO
from main import connect_with_postgres

class TestPostgresConnection(unittest.TestCase):
    def setUp(self):
        # 提前读取测试CSV数据
        with open('mocks/table.csv', 'r') as file:
            self.mock_table_csv = file.read()

    @patch('main.create_engine')  # 注意:patch路径要对应create_engine在main.py中的引用路径
    def test_connect_with_postgres(self, mock_create_engine):
        # 1. 构建Mock连接对象
        mock_conn = Mock()
        
        # 2. 构建Mock引擎对象,让其connect方法返回Mock连接
        mock_engine = Mock()
        mock_engine.connect.return_value = mock_conn
        mock_create_engine.return_value = mock_engine

        # 3. 加载测试用DataFrame
        mock_table = pd.read_csv(StringIO(self.mock_table_csv))

        # 4. 调用待测试方法,传入虚假连接字符串
        connect_with_postgres(mock_table, "fake_postgres_conn_string")

        # 5. 验证关键方法的调用是否符合预期
        mock_create_engine.assert_called_once_with("fake_postgres_conn_string")
        mock_engine.connect.assert_called_once()
        mock_table.to_sql.assert_called_once_with(
            'table_name', 
            con=mock_conn, 
            if_exists='replace', 
            schema='schemaname', 
            index=False
        )
        mock_conn.close.assert_called_once()

关键注意点

  1. Patch路径正确性:如果main.py中是from sqlalchemy import create_engine,则patch路径为main.create_engine;如果是import sqlalchemy后用sqlalchemy.create_engine,则patch路径应为main.sqlalchemy.create_engine
  2. Mock层级匹配:必须依次Mockcreate_engine的返回值(引擎对象)、引擎对象的connect()返回值(连接对象),确保原方法的每一步调用都能找到对应的Mock实例,避免出现解包错误
  3. 调用验证:通过assert_called_once()和assert_called_with()确认方法调用的正确性,保证测试的有效性

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.27 02:45:19