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
144 changes: 144 additions & 0 deletions tests/test_testcase.py
Original file line number Diff line number Diff line change
Expand Up @@ -2343,6 +2343,150 @@ def test_other_attribute(self):
self.assertRaises(AttributeError, getattr, orig, "thing")


class TestSubTest(TestCase):
"""Tests for testtools.TestCase.subTest support."""

run_test_with = FullStackRunTest

def _run_case(self, case):
result = ExtendedTestResult()
case.run(result)
return result

def test_passing_subtests(self):
class Case(TestCase):
def test_it(self):
for i in (0, 2, 4):
with self.subTest(i=i):
self.assertEqual(i % 2, 0)

result = self._run_case(Case("test_it"))
self.assertIn("addSuccess", [e[0] for e in result._events])
self.assertEqual([], result.failures)

def test_single_failure(self):
class Case(TestCase):
def test_it(self):
for i in (0, 1, 2):
with self.subTest(i=i):
self.assertEqual(i % 2, 0)

result = self._run_case(Case("test_it"))
self.assertNotIn("addSuccess", [e[0] for e in result._events])
self.assertEqual(1, len(result.failures))
subtest = result.failures[0][0]
self.assertIn("(i=1)", str(subtest))

def test_multiple_failures(self):
class Case(TestCase):
def test_it(self):
for i in range(4):
with self.subTest(i=i):
self.assertEqual(i % 2, 0)

result = self._run_case(Case("test_it"))
self.assertEqual(2, len(result.failures))
descriptions = [str(f[0]) for f in result.failures]
self.assertTrue(
any("(i=1)" in d for d in descriptions),
f"Expected a failure for (i=1), got {descriptions}",
)
self.assertTrue(
any("(i=3)" in d for d in descriptions),
f"Expected a failure for (i=3), got {descriptions}",
)

def test_failure_continues_loop(self):
class Case(TestCase):
def __init__(self, *args, **kwargs):
super().__init__(*args, **kwargs)
self.iterations = []

def test_it(self):
for i in range(4):
with self.subTest(i=i):
self.iterations.append(i)
self.assertEqual(i % 2, 0)

case = Case("test_it")
self._run_case(case)
self.assertEqual([0, 1, 2, 3], case.iterations)

def test_with_msg(self):
class Case(TestCase):
def test_it(self):
with self.subTest(msg="my label", x=42):
self.fail("boom")

result = self._run_case(Case("test_it"))
self.assertEqual(1, len(result.failures))
description = str(result.failures[0][0])
self.assertIn("[my label]", description)
self.assertIn("(x=42)", description)

def test_multiple_params(self):
class Case(TestCase):
def test_it(self):
with self.subTest(a=1, b="two"):
self.fail("boom")

result = self._run_case(Case("test_it"))
self.assertEqual(1, len(result.failures))
description = str(result.failures[0][0])
self.assertIn("a=1", description)
self.assertIn("b='two'", description)

def test_no_params(self):
class Case(TestCase):
def test_it(self):
with self.subTest():
self.fail("boom")

result = self._run_case(Case("test_it"))
self.assertEqual(1, len(result.failures))
description = str(result.failures[0][0])
self.assertIn("(<subtest>)", description)

def test_nested_subtests(self):
class Case(TestCase):
def test_it(self):
for a in (1, 2):
with self.subTest(a=a):
for b in (3, 4):
with self.subTest(b=b):
self.assertEqual(a, b)

result = self._run_case(Case("test_it"))
self.assertEqual(4, len(result.failures))
descriptions = [str(f[0]) for f in result.failures]
self.assertTrue(
any("a=1" in d and "b=3" in d for d in descriptions),
f"Expected a failure with both a=1 and b=3, got {descriptions}",
)
self.assertTrue(
any("a=2" in d and "b=4" in d for d in descriptions),
f"Expected a failure with both a=2 and b=4, got {descriptions}",
)

def test_skip_inside_subtest(self):
class Case(TestCase):
def __init__(self, *args, **kwargs):
super().__init__(*args, **kwargs)
self.reached = []

def test_it(self):
for i in range(3):
with self.subTest(i=i):
if i == 1:
self.skipTest("skip this one")
self.reached.append(i)

case = Case("test_it")
self._run_case(case)
self.assertEqual([0, 2], case.reached)
self.assertEqual(1, len(case._subtest_skips))


def test_suite():
from unittest import TestLoader

Expand Down
9 changes: 9 additions & 0 deletions testtools/runtest.py
Original file line number Diff line number Diff line change
Expand Up @@ -184,6 +184,15 @@ def _run_core(self) -> None:
if getattr(self.case, "force_failure", None):
self._run_user(_raise_force_fail_error)
failed = True
for subtest, err in getattr(self.case, "_subtest_failures", ()):
add_subtest = getattr(self.result, "addSubTest", None)
if add_subtest is not None:
add_subtest(self.case, subtest, err)
else:
self.result.addFailure(self.case, err)
failed = True
for subtest, reason in getattr(self.case, "_subtest_skips", ()):
self.result.addSkip(subtest, reason=reason)
if not failed:
self.result.addSuccess(
self.case, details=self.case.getDetails()
Expand Down
62 changes: 62 additions & 0 deletions testtools/testcase.py
Original file line number Diff line number Diff line change
Expand Up @@ -15,6 +15,7 @@
"unique_text_generator",
]

import contextlib
import copy
import datetime
import functools
Expand Down Expand Up @@ -99,6 +100,44 @@ class _ExpectedFailure(Exception):
"""


_subtest_msg_sentinel = object()


class SubTest(unittest.TestCase):
"""Describes a single subTest iteration for failure reporting.

Carries the message and parameters passed to ``subTest`` so that
result objects can label each failure with its subTest context.
"""

def __init__(
self, test_case: "TestCase", msg: object, params: dict[str, Any]
) -> None:
super().__init__()
self.test_case = test_case
self.failureException = test_case.failureException
self._msg = msg
self._params = params

def _subDescription(self) -> str:
parts: list[str] = []
if self._msg is not _subtest_msg_sentinel:
parts.append(f"[{self._msg}]")
if self._params:
params_desc = ", ".join(f"{k}={v!r}" for k, v in self._params.items())
parts.append(f"({params_desc})")
return " ".join(parts) or "(<subtest>)"

def id(self) -> str:
return f"{self.test_case.id()} {self._subDescription()}"

def shortDescription(self) -> str | None:
return self.test_case.shortDescription()

def __str__(self) -> str:
return self.id()


# TypeVar for decorators
_P = ParamSpec("_P")
_R = TypeVar("_R")
Expand Down Expand Up @@ -378,6 +417,9 @@ def _reset(self) -> None:
# force_failure is set by expectThat() on mismatch; must be
# cleared so re-runs of the same test can succeed.
self.force_failure: bool | None = None
self._subtest_failures: list[tuple[SubTest, ExcInfo]] = []
self._subtest_skips: list[tuple[SubTest, str]] = []
self._subtest_params: dict[str, Any] = {}

def __eq__(self, other: object) -> bool:
eq = getattr(unittest.TestCase, "__eq__", None)
Expand Down Expand Up @@ -887,6 +929,26 @@ def _report_traceback(
),
)

@contextlib.contextmanager
def subTest(
self, msg: object = _subtest_msg_sentinel, **params: Any
) -> Iterator[None]:
"""Return a context manager for a subTest."""
merged_params = {**self._subtest_params, **params}
subtest = SubTest(self, msg, merged_params)
old_params, self._subtest_params = self._subtest_params, merged_params
try:
yield
except SkipTest as e:
reason = str(e)
self._subtest_skips.append((subtest, reason))
except Exception:

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

This catches SkipTest, so self.skipTest(...) inside a subTest is recorded as an error rather than a skip. Against stdlib, for a case that skips one subtest and fails another.

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

OK I caught SkipTest and reraised the exception.

# Inside except block, exc_info() is guaranteed to have non-None values
exc_info = sys.exc_info()
self._subtest_failures.append((subtest, exc_info)) # type: ignore[arg-type]
finally:
self._subtest_params = old_params

@staticmethod
def _report_unexpected_success(
self: "TestCase", result: TestResult, err: BaseException
Expand Down