如何在Python pytest单元测试中Mock S3 Parquet数据加载?
用pytest Mock S3读取Parquet的单元测试实现思路
你要测试的load_parquet函数核心依赖wr.s3.read_parquet与S3交互,要避免实际调用S3,只要把这个方法mock掉,让它返回预设的测试数据就行。用pytest配合unittest.mock就能搞定,具体步骤如下:
1. 准备模拟测试数据
先构造一个和实际返回结构一致的DataFrame,包含重复数据,用来验证后续的去重逻辑是否生效:
import pandas as pd from datetime import date def get_mock_df(): # 构造带重复项的测试数据,包含函数用到的date列 data = { "id": [1, 2, 2, 3], "name": ["Alice", "Bob", "Bob", "Charlie"], "date": [date(2024, 1, 1), date(2024, 1, 2), date(2024, 1, 2), date(2024, 1, 3)] } return pd.DataFrame(data)
2. 用patch替换真实的S3读取方法
在测试函数里,用unittest.mock.patch把wr.s3.read_parquet替换成自定义mock函数,既可以返回预设数据,还能验证传入的参数是否正确(记得把your_module换成你实际的模块名):
from unittest.mock import patch from datetime import date from your_module import load_parquet def test_load_parquet_dedup_logic(): # 定义测试用参数 test_bucket = "test-bucket" test_prefix = "test-prefix" test_folder = "test-folder" required_cols = ["id", "name"] start_date = date(2024, 1, 1) end_date = date(2024, 1, 3) # 自定义mock函数,返回预设数据并验证参数 def mock_read_parquet(path, columns, partition_filter, dataset, use_threads): # 检查生成的S3路径是否正确 assert path == f"s3://{test_bucket}/{test_prefix}/{test_folder}/" # 检查传入的列是否符合预期 assert columns == required_cols # 返回预设的模拟DataFrame return get_mock_df() # 用上下文管理器替换原函数 with patch("your_module.wr.s3.read_parquet", side_effect=mock_read_parquet): result_df = load_parquet(test_bucket, test_prefix, test_folder, required_cols, start_date, end_date) # 验证去重结果:原数据4行,去重后应剩3行 assert len(result_df) == 3 # 验证去重后的id列表正确 assert result_df["id"].tolist() == [1, 2, 3]
3. 用fixture复用mock逻辑(可选)
如果多个测试都需要mock这个S3读取方法,可以把mock逻辑封装成pytest fixture,减少重复代码:
import pytest from unittest.mock import patch from your_module import load_parquet from datetime import date @pytest.fixture def mock_s3_read(): # 用patch包裹mock逻辑,yield给测试函数使用 with patch("your_module.wr.s3.read_parquet") as mock_read: mock_read.return_value = get_mock_df() yield mock_read def test_load_parquet_with_fixture(mock_s3_read): test_bucket = "test-bucket" test_prefix = "test-prefix" test_folder = "test-folder" required_cols = ["id", "name"] start_date = date(2024, 1, 1) end_date = date(2024, 1, 3) result_df = load_parquet(test_bucket, test_prefix, test_folder, required_cols, start_date, end_date) # 验证mock方法被调用了一次 mock_s3_read.assert_called_once() # 检查调用时传入的参数是否正确 call_args = mock_s3_read.call_args assert call_args[1]["path"] == f"s3://{test_bucket}/{test_prefix}/{test_folder}/" assert call_args[1]["columns"] == required_cols # 验证去重结果 assert len(result_df) == 3
4. 验证partition_filter逻辑(可选)
如果想确保partition_filter的日期筛选逻辑正确,可以在mock函数里测试这个lambda:
def mock_read_parquet(path, columns, partition_filter, dataset, use_threads): # 构造不同日期的分区,测试筛选逻辑 test_partitions = [ {"date": date(2023, 12, 31)}, {"date": date(2024, 1, 2)}, {"date": date(2024, 1, 4)} ] # 应用partition_filter筛选 filtered_partitions = [p for p in test_partitions if partition_filter(p)] # 预期只保留2024-1-2这个在日期范围内的分区 assert len(filtered_partitions) == 1 assert filtered_partitions[0]["date"] == date(2024, 1, 2) return get_mock_df()
内容的提问来源于stack exchange,提问作者pmanDS
相关产品推荐
相关产品推荐

