如何Mock Spanner的StreamedResultSet修复Pytest执行SQL的列匹配报错?
解决Spanner StreamedResultSet Mock导致的DataFrame列数不匹配问题
问题背景
在为Google Spanner的Database类编写Pytest用例时,execute_query的测试用例触发报错:ValueError: 0 columns passed, passed data had 1 columns。
原Database类实现
class Database: _client = None def __init__(self, instance_id: str, database_id: str, pool=None): if not Database._client: Database._client = SpannerClient() instance = Database._client.instance(instance_id) self._database = instance.database(database_id, pool=pool) def execute_query( self, query: str, params: Dict | None = None, param_types: Dict | None = None ): try: with self._database.snapshot() as snapshot: results = snapshot.execute_sql(query, params, param_types) df = DataFrame( data=[row for row in results], columns=[col.name for col in results.fields], ) return df except GoogleAPICallError as e: print(f"Error code:{e.code},Error message: {e.message}") raise GoogleAPIError() from e def get_database(self): return self._database
已编写的测试代码
import pytest from unittest.mock import patch, MagicMock from database import Database from google.api_core.exceptions import GoogleAPICallError from pandas import DataFrame @pytest.fixture def mock_spanner_client(): with patch("database.Database._client") as MockClient: yield MockClient @pytest.fixture def mock_instance(mock_spanner_client): mock_instance = MagicMock() mock_spanner_client.instance.return_value = mock_instance yield mock_instance @pytest.fixture def mock_database(mock_instance): mock_database = MagicMock() mock_instance.database.return_value = mock_database yield mock_database def test_database_initialization(mock_spanner_client, mock_instance, mock_database): db = Database("test_instance", "test_database") assert db._database == mock_database def test_get_database(mock_database): db = Database("test_instance", "test_database") assert db.get_database() == mock_database def test_execute_query_success(mock_database): mock_snapshot = MagicMock() mock_snapshot.execute_sql.return_value = MagicMock( __iter__=lambda self: iter([["row1"], ["row2"]]), fields=[MagicMock(name="col1")], ) mock_database.snapshot.return_value.__enter__.return_value = mock_snapshot db = Database("test_instance", "test_database") query = "SELECT * FROM test_table" result = db.execute_query(query) assert isinstance(result, DataFrame) assert not result.empty assert list(result.columns) == ["col1"]
问题原因
原测试中,在MagicMock初始化时同时定义__iter__和fields属性,导致MagicMock的属性拦截机制干扰了fields的正常访问,使得[col.name for col in results.fields]返回空列表,最终触发DataFrame列数不匹配错误。
解决方案
方式一:使用自定义模拟类
创建简单类模拟StreamedResultSet的行为,明确定义fields和迭代逻辑:
def test_execute_query_success(mock_database): mock_snapshot = MagicMock() # 自定义模拟StreamedResultSet class MockStreamedResultSet: def __iter__(self): return iter([["row1"], ["row2"]]) @property def fields(self): return [MagicMock(name="col1")] mock_snapshot.execute_sql.return_value = MockStreamedResultSet() mock_database.snapshot.return_value.__enter__.return_value = mock_snapshot db = Database("test_instance", "test_database") query = "SELECT * FROM test_table" result = db.execute_query(query) assert isinstance(result, DataFrame) assert not result.empty assert list(result.columns) == ["col1"]
方式二:分步设置MagicMock属性
避免在MagicMock初始化时同时定义多个属性,改为分步赋值确保fields正常读取:
def test_execute_query_success(mock_database): mock_snapshot = MagicMock() # 创建模拟的StreamedResultSet mock_result = MagicMock() mock_result.__iter__ = lambda self: iter([["row1"], ["row2"]]) mock_result.fields = [MagicMock(name="col1")] mock_snapshot.execute_sql.return_value = mock_result mock_database.snapshot.return_value.__enter__.return_value = mock_snapshot db = Database("test_instance", "test_database") query = "SELECT * FROM test_table" result = db.execute_query(query) assert isinstance(result, DataFrame) assert not result.empty assert list(result.columns) == ["col1"]
内容的提问来源于stack exchange,提问作者Rudra
相关产品推荐
相关产品推荐

