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

如何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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.20 06:26:06