Skip to content
Open
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
29 changes: 27 additions & 2 deletions Lib/test/test_unittest/test_discovery.py
Original file line number Diff line number Diff line change
Expand Up @@ -386,7 +386,7 @@ def restore_path():
with self.assertRaises(ImportError):
loader.discover('/foo/bar', top_level_dir='/foo')

self.assertEqual(loader._top_level_dir, full_path)
self.assertIsNone(loader._top_level_dir)
self.assertIn(full_path, sys.path)

os.path.isfile = lambda path: True
Expand All @@ -408,7 +408,7 @@ def _find_tests(start_dir, pattern, namespace=None):
top_level_dir = os.path.abspath('/foo/bar')
start_dir = os.path.abspath('/foo/bar/baz')
self.assertEqual(suite, "['tests']")
self.assertEqual(loader._top_level_dir, os.path.abspath('/foo'))
self.assertIsNone(loader._top_level_dir)
self.assertEqual(_find_tests_args, [(start_dir, 'pattern')])
self.assertIn(top_level_dir, sys.path)

Expand Down Expand Up @@ -436,6 +436,31 @@ def restore():
loader.discover(dir, top_level_dir=top_level_dir)
self.assertEqual(loader._top_level_dir, dir2)

def test_discover_should_not_persist_top_level_dir_on_error(self):
original_isfile = os.path.isfile
original_isdir = os.path.isdir
original_sys_path = sys.path[:]
def restore():
os.path.isfile = original_isfile
os.path.isdir = original_isdir
sys.path[:] = original_sys_path
self.addCleanup(restore)

os.path.isfile = lambda path: False
os.path.isdir = lambda path: True
loader = unittest.TestLoader()
dir = '/foo/bar'
top_level_dir = '/foo'

with self.assertRaises(ImportError):
loader.discover(dir, top_level_dir=top_level_dir)
self.assertIsNone(loader._top_level_dir)

loader._top_level_dir = dir2 = '/previous/dir'
with self.assertRaises(ImportError):
loader.discover(dir, top_level_dir=top_level_dir)
self.assertEqual(loader._top_level_dir, dir2)

def test_discover_start_dir_is_package_calls_package_load_tests(self):
# This test verifies that the package load_tests in a package is indeed
# invoked when the start_dir is a package (and not the top level).
Expand Down
121 changes: 61 additions & 60 deletions Lib/unittest/loader.py
Original file line number Diff line number Diff line change
Expand Up @@ -280,70 +280,71 @@ def discover(self, start_dir, pattern='test*.py', top_level_dir=None):
sys.path.insert(0, top_level_dir)
self._top_level_dir = top_level_dir

is_not_importable = False
is_namespace = False
tests = []
if os.path.isdir(os.path.abspath(start_dir)):
start_dir = os.path.abspath(start_dir)
if start_dir != top_level_dir:
is_not_importable = not os.path.isfile(os.path.join(start_dir, '__init__.py'))
else:
# support for discovery from dotted module names
try:
__import__(start_dir)
except ImportError:
is_not_importable = True
try:
is_not_importable = False
is_namespace = False
tests = []
if os.path.isdir(os.path.abspath(start_dir)):
start_dir = os.path.abspath(start_dir)
if start_dir != top_level_dir:
is_not_importable = not os.path.isfile(os.path.join(start_dir, '__init__.py'))
else:
the_module = sys.modules[start_dir]
if not hasattr(the_module, "__file__") or the_module.__file__ is None:
# look for namespace packages
try:
spec = the_module.__spec__
except AttributeError:
spec = None

if spec and spec.submodule_search_locations is not None:
is_namespace = True

for path in the_module.__path__:
if (not set_implicit_top and
not path.startswith(top_level_dir)):
continue
self._top_level_dir = \
(path.split(the_module.__name__
.replace(".", os.path.sep))[0])
tests.extend(self._find_tests(path, pattern, namespace=True))
elif the_module.__name__ in sys.builtin_module_names:
# builtin module
raise TypeError('Can not use builtin modules '
'as dotted module names') from None
else:
raise TypeError(
f"don't know how to discover from {the_module!r}"
) from None

# support for discovery from dotted module names
try:
__import__(start_dir)
except ImportError:
is_not_importable = True
else:
top_part = start_dir.split('.')[0]
start_dir = os.path.abspath(os.path.dirname((the_module.__file__)))

if set_implicit_top:
if not is_namespace:
if sys.modules[top_part].__file__ is None:
self._top_level_dir = os.path.dirname(the_module.__file__)
if self._top_level_dir not in sys.path:
sys.path.insert(0, self._top_level_dir)
the_module = sys.modules[start_dir]
if not hasattr(the_module, "__file__") or the_module.__file__ is None:
# look for namespace packages
try:
spec = the_module.__spec__
except AttributeError:
spec = None

if spec and spec.submodule_search_locations is not None:
is_namespace = True

for path in the_module.__path__:
if (not set_implicit_top and
not path.startswith(top_level_dir)):
continue
self._top_level_dir = \
(path.split(the_module.__name__
.replace(".", os.path.sep))[0])
tests.extend(self._find_tests(path, pattern, namespace=True))
elif the_module.__name__ in sys.builtin_module_names:
# builtin module
raise TypeError('Can not use builtin modules '
'as dotted module names') from None
else:
self._top_level_dir = \
self._get_directory_containing_module(top_part)
sys.path.remove(top_level_dir)

if is_not_importable:
raise ImportError('Start directory is not importable: %r' % start_dir)
raise TypeError(
f"don't know how to discover from {the_module!r}"
) from None

if not is_namespace:
tests = list(self._find_tests(start_dir, pattern))

self._top_level_dir = original_top_level_dir
else:
top_part = start_dir.split('.')[0]
start_dir = os.path.abspath(os.path.dirname((the_module.__file__)))

if set_implicit_top:
if not is_namespace:
if sys.modules[top_part].__file__ is None:
self._top_level_dir = os.path.dirname(the_module.__file__)
if self._top_level_dir not in sys.path:
sys.path.insert(0, self._top_level_dir)
else:
self._top_level_dir = \
self._get_directory_containing_module(top_part)
sys.path.remove(top_level_dir)

if is_not_importable:
raise ImportError('Start directory is not importable: %r' % start_dir)

if not is_namespace:
tests = list(self._find_tests(start_dir, pattern))
finally:
self._top_level_dir = original_top_level_dir
return self.suiteClass(tests)

def _get_directory_containing_module(self, module_name):
Expand Down
Original file line number Diff line number Diff line change
@@ -0,0 +1,3 @@
Fix :meth:`unittest.TestLoader.discover` so that a failed call no longer
changes the top-level directory used by later calls on the same loader.
Patch by Christian Aurich Zanettini Martins.
Loading