From f662ae4304ce3321a85632731444a2b122489af4 Mon Sep 17 00:00:00 2001 From: Cyril Roelandt Date: Tue, 8 Sep 2026 20:47:47 +0200 Subject: [PATCH] Add subTest support for testtools.TestCase Commit 267c0be80d6de80242ff427c4b3e5d15086d64f7 added support for unittest.TestCase.subTest. This meant the following code: import unittest class TestSubTests(unittest.TestCase): def test_even_numbers(self): for i in range(5): with self.subTest(i=i): self.assertEqual(i % 2, 0) would list the exact values of "i" for which the test failed: $ stestr run test_example_subtest 2>&1 | grep ^Captured Captured traceback (i=1): Captured traceback (i=3): Captured traceback (i=1): Captured traceback (i=3): But the following code: import testtools class TestSubTests(testtools.TestCase): def test_even_numbers(self): for i in range(5): with self.subTest(i=i): self.assertEqual(i % 2, 0) would not produce a similar output: $ stestr run test_example_subtest_testtools 2>&1 | grep ^Captured Captured traceback: Captured traceback: This commit fixes this so that we get the same output whether we use unittest.TestCase.subTest or testtools.TestCase.subTest. Closes: #317 Assisted-by: Claude Opus 4.6 (1M context) Signed-off-by: Cyril Roelandt --- tests/test_testcase.py | 144 +++++++++++++++++++++++++++++++++++++++++ testtools/runtest.py | 9 +++ testtools/testcase.py | 62 ++++++++++++++++++ 3 files changed, 215 insertions(+) diff --git a/tests/test_testcase.py b/tests/test_testcase.py index e7f0bcdc..e03be132 100644 --- a/tests/test_testcase.py +++ b/tests/test_testcase.py @@ -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("()", 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 diff --git a/testtools/runtest.py b/testtools/runtest.py index f4ef9648..1d62a010 100644 --- a/testtools/runtest.py +++ b/testtools/runtest.py @@ -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() diff --git a/testtools/testcase.py b/testtools/testcase.py index bdd640a9..c1603e9d 100644 --- a/testtools/testcase.py +++ b/testtools/testcase.py @@ -15,6 +15,7 @@ "unique_text_generator", ] +import contextlib import copy import datetime import functools @@ -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 "()" + + 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") @@ -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) @@ -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: + # 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