【问题标题】:Python - Unable to mock call for Inherited classPython - 无法模拟对继承类的调用
【发布时间】:2018-09-18 15:01:17
【问题描述】:

我有这个主课

def main(args):
    if type == train_pipeline_type:
        strategy = TrainPipelineStrategy()
    else:
        strategy = TestPipelineStrategy()
    for table in fetch_table_information_by_region(region):
        split_required = DataUtils.load_from_dict(table, "split_required")
        if split_required:
            strategy.split(spark=spark, table_name=table_name,
                           data_loc=filtered_data_location, partition_column=partition_column,
                           split_output_dir= split_output_dir)
            logger.info("Data Split for table : {} completed".format(table_name))

我的 TrainPipelineStrategy 和 TestPipelineStrategy 看起来像这样 -

class PipelineTypeStrategy(object):

    def partition_data(self, x):
        # Something

    def prepare_split_data(self, y):
        # Something

    def write_split_data(self, z):
        # Something

    def split(self, p):
        # Something


class TrainPipelineStrategy(PipelineTypeStrategy):
    """"""


class TestPipelineStrategy(PipelineTypeStrategy):

    def write_split_data(self, y):
        # Something else

我的测试用例 - 我需要通过在 main 方法中模拟 split 功能来测试 split 被调用了多少次。

这是我尝试过的 -

@patch('module.PipelineTypeStrategy.TrainPipelineStrategy')
    def test_split_data_main_split_data_call_count(self, fake_train):
        fake_train_functions = mock.Mock()
        fake_train_functions.split.return_value = None
        fake_train.return_value = fake_train_functions
        test_args = ["", "--x=6"]
        SplitData.main(args=test_args)
        assert fake_train_functions.split.call_count == 10

当我尝试运行我的测试时,它会创建模拟,但最终会调用实际的拆分函数。我做错了什么?

【问题讨论】:

  • 我无法理解您的代码,但猴子修补比看起来更难。如果您将要模拟的内容作为参数传递给 SUT,会更容易。

标签: python python-3.x python-unittest


【解决方案1】:

这段代码的主要问题是,如果TrainPipelineStrategyPipelineTypeStrategy 的嵌套类,但TrainPipelineStrategyPipelineTypeStrategy 的子类,那么您设置patch 的方式。

由于TrainPipelineStrategy 继承自PipelineTypeStrategy,它可以直接访问split,因此您可以在不引用PipelineTypeStrategy 的情况下修补split(除非您特别想修补定义在PipelineTypeStrategy)。

但是,如果您只想模拟 PipelineTypeStrategy 类的 split 方法,您应该使用 patch.object 装饰器来模拟 split 而不是模拟整个类,因为它更干净一点。这是一个例子:

class TestClass(unittest.TestCase):
    @patch.object(TrainPipelineStrategy, 'split', return_value=None)
    def test_split_data_main_split_data_call_count(self, mock_split):
        test_args = ["", "--x=6"]
        SplitData.main(args=test_args)
        self.assertEqual(mock_split.call_count, 10)

【讨论】:

    猜你喜欢
    • 1970-01-01
    • 2011-08-08
    • 2020-11-11
    • 2016-03-04
    • 1970-01-01
    • 1970-01-01
    • 2016-06-06
    • 2016-06-19
    • 1970-01-01
    相关资源
    最近更新 更多