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()
关键注意点
- Patch路径正确性:如果
main.py中是from sqlalchemy import create_engine,则patch路径为main.create_engine;如果是import sqlalchemy后用sqlalchemy.create_engine,则patch路径应为main.sqlalchemy.create_engine - Mock层级匹配:必须依次Mock
create_engine的返回值(引擎对象)、引擎对象的connect()返回值(连接对象),确保原方法的每一步调用都能找到对应的Mock实例,避免出现解包错误 - 调用验证:通过
assert_called_once()和assert_called_with()确认方法调用的正确性,保证测试的有效性
内容的提问来源于stack exchange,提问作者gwc
相关产品推荐
相关产品推荐

