Unable to pass mocked fixture as an argument in parametrized tests in pytest

Viewed 22

See below, is it even possible to pass a mock fixture (mock_adam, "adam", "0.001", ...) in parametrized tests for reusability purposes?

import pytest

from contextlib import contextmanager
from unittest import mock
from my_module import get_optimizer


@contextmanager
def does_not_raise():
    yield


@pytest.fixture(autouse=True)
def mock_adam():
    with mock.patch("my_module.optimizers.Adam") as mocker:
        yield mocker


@pytest.fixture(autouse=True)
def mock_RMSprop():
    with mock.patch("my_module.optimizers.RMSprop") as mocker:
        yield mocker


class TestGetOptimizers:
    @pytest.mark.parametrize(
        "mock_optimizer, optimizer_name, learning_rate, clipnorm, expectation",
        [
            (mock_adam, "adam", "0.001", "1.1", does_not_raise()),
            (mock_RMSprop, "rmsprop", "0.001", "1.1", does_not_raise()),
        ],
    )
    def test_get_optimizer(self, mock_optimizer, optimizer_name, learning_rate, clipnorm, expectation):
        with expectation:
            get_optimizer(
                optimizer_name=optimizer_name,
                learning_rate=learning_rate,
                clipnorm=clipnorm,
            )
            mock_optimizer.assert_called_once_with(lr=learning_rate, clipnorm=clipnorm)
AttributeError: 'function' object has no attribute 'assert_called_once_with'
1 Answers

To achieve this, you can create a single fixture that would get the path to mock as an indirect parameter, like so:

import pytest

from contextlib import contextmanager
from unittest import mock
from my_module import get_optimizer


@contextmanager
def does_not_raise():
    yield


@pytest.fixture
def mock_with_param(request):
    with mock.patch(request.param) as mocker:
        yield mocker


class TestGetOptimizers:
    @pytest.mark.parametrize(
        "mock_with_param, optimizer_name, learning_rate, clipnorm, expectation",
        [
            ("my_module.optimizers.Adam", "adam", "0.001", "1.1", does_not_raise()),
            ("my_module.optimizers.RMSprop", "rmsprop", "0.001", "1.1", does_not_raise()),
        ], indirect=["mock_with_param"]
    )
    def test_get_optimizer(self, mock_with_param, optimizer_name, learning_rate, clipnorm, expectation):
        with expectation:
            get_optimizer(
                optimizer_name=optimizer_name,
                learning_rate=learning_rate,
                clipnorm=clipnorm,
            )
            mock_with_param.assert_called_once_with(lr=learning_rate, clipnorm=clipnorm)
Related