Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
17 changes: 4 additions & 13 deletions Lib/unittest/mock.py
Original file line number Diff line number Diff line change
Expand Up @@ -1042,8 +1042,7 @@ class _AnyComparer(list):
the left."""
def __contains__(self, item):
for _call in self:
if len(item) != len(_call):
continue
assert len(item) == len(_call)
if all([
expected == actual
for expected, actual in zip(item, _call)
Expand Down Expand Up @@ -1856,7 +1855,8 @@ def _unpatch_dict(self):

def __exit__(self, *args):
"""Unpatch the dict."""
self._unpatch_dict()
if self._original is not None:
self._unpatch_dict()
return False


Expand Down Expand Up @@ -2168,7 +2168,7 @@ def __init__(self, /, *args, **kwargs):
self.__dict__['__code__'] = code_mock

async def _execute_mock_call(self, /, *args, **kwargs):
# This is nearly just like super(), except for sepcial handling
# This is nearly just like super(), except for special handling
# of coroutines

_call = self.call_args
Expand Down Expand Up @@ -2541,12 +2541,6 @@ def __getattribute__(self, attr):
return tuple.__getattribute__(self, attr)


def count(self, /, *args, **kwargs):
return self.__getattr__('count')(*args, **kwargs)

def index(self, /, *args, **kwargs):
return self.__getattr__('index')(*args, **kwargs)

def _get_call_arguments(self):
if len(self) == 2:
args, kwargs = self
Expand Down Expand Up @@ -2921,9 +2915,6 @@ def __init__(self, iterator):
code_mock.co_flags = inspect.CO_ITERABLE_COROUTINE
self.__dict__['__code__'] = code_mock

def __aiter__(self):
return self

async def __anext__(self):
try:
return next(self.iterator)
Expand Down
59 changes: 18 additions & 41 deletions Lib/unittest/test/testmock/testasync.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,38 +16,28 @@ def tearDownModule():


class AsyncClass:
def __init__(self):
pass
async def async_method(self):
pass
def normal_method(self):
pass
def __init__(self): pass
async def async_method(self): pass
def normal_method(self): pass

@classmethod
async def async_class_method(cls):
pass
async def async_class_method(cls): pass

@staticmethod
async def async_static_method():
pass
async def async_static_method(): pass


class AwaitableClass:
def __await__(self):
yield
def __await__(self): yield

async def async_func():
pass
async def async_func(): pass

async def async_func_args(a, b, *, c):
pass
async def async_func_args(a, b, *, c): pass

def normal_func():
pass
def normal_func(): pass

class NormalClass(object):
def a(self):
pass
def a(self): pass


async_foo_name = f'{__name__}.AsyncClass'
Expand Down Expand Up @@ -402,17 +392,15 @@ def test_magicmock_lambda_spec(self):

class AsyncArguments(IsolatedAsyncioTestCase):
async def test_add_return_value(self):
async def addition(self, var):
return var + 1
async def addition(self, var): pass

mock = AsyncMock(addition, return_value=10)
output = await mock(5)

self.assertEqual(output, 10)

async def test_add_side_effect_exception(self):
async def addition(var):
return var + 1
async def addition(var): pass
mock = AsyncMock(addition, side_effect=Exception('err'))
with self.assertRaises(Exception):
await mock(5)
Expand Down Expand Up @@ -553,18 +541,14 @@ def test_magic_methods_are_async_functions(self):
class AsyncContextManagerTest(unittest.TestCase):

class WithAsyncContextManager:
async def __aenter__(self, *args, **kwargs):
return self
async def __aenter__(self, *args, **kwargs): pass

async def __aexit__(self, *args, **kwargs):
pass
async def __aexit__(self, *args, **kwargs): pass

class WithSyncContextManager:
def __enter__(self, *args, **kwargs):
return self
def __enter__(self, *args, **kwargs): pass

def __exit__(self, *args, **kwargs):
pass
def __exit__(self, *args, **kwargs): pass

class ProductionCode:
# Example real-world(ish) code
Expand Down Expand Up @@ -673,16 +657,9 @@ class WithAsyncIterator(object):
def __init__(self):
self.items = ["foo", "NormalFoo", "baz"]

def __aiter__(self):
return self

async def __anext__(self):
try:
return self.items.pop()
except IndexError:
pass
def __aiter__(self): pass

raise StopAsyncIteration
async def __anext__(self): pass

def test_aiter_set_return_value(self):
mock_iter = AsyncMock(name="tester")
Expand Down
5 changes: 5 additions & 0 deletions Lib/unittest/test/testmock/testmock.py
Original file line number Diff line number Diff line change
Expand Up @@ -1868,6 +1868,11 @@ def test_mock_open_using_next(self):
with self.assertRaises(StopIteration):
next(f1)

def test_mock_open_next_with_readline_with_return_value(self):
mopen = mock.mock_open(read_data='foo\nbarn')
mopen.return_value.readline.return_value = 'abc'
self.assertEqual('abc', next(mopen()))

def test_mock_open_write(self):
# Test exception in file writing write()
mock_namedtemp = mock.mock_open(mock.MagicMock(name='JLV'))
Expand Down
8 changes: 8 additions & 0 deletions Lib/unittest/test/testmock/testpatch.py
Original file line number Diff line number Diff line change
Expand Up @@ -770,6 +770,14 @@ def test_patch_dict_start_stop(self):
self.assertEqual(d, original)


def test_patch_dict_stop_without_start(self):
d = {'foo': 'bar'}
original = d.copy()
patcher = patch.dict(d, [('spam', 'eggs')], clear=True)
self.assertEqual(patcher.stop(), False)
self.assertEqual(d, original)


def test_patch_dict_class_decorator(self):
this = self
d = {'spam': 'eggs'}
Expand Down