From 7760d420b5b187fd04f2f339e9e42b8f7444a762 Mon Sep 17 00:00:00 2001 From: Ang Li Date: Tue, 29 Sep 2026 17:17:48 +0000 Subject: [PATCH 1/3] Preserve test failure state during teardown_test Separate test execution and teardown in BaseTestClass.exec_one_test so that: 1. The test execution outcome is recorded on tr_record before teardown_test runs, ensuring self.current_test_info.record accurately reflects test result and details in teardown. 2. Exceptions or errors occurring during teardown_test do not overwrite an existing FAIL or ERROR result from the test. 3. TestAbortSignal raised during teardown_test after an existing failure is recorded in extra_errors['teardown_test'] and re-raised to properly abort subsequent tests without erasing the primary failure. Fixes #1036 --- mobly/base_test.py | 127 +++++++++-------- tests/mobly/base_test_test.py | 250 ++++++++++++++++++++++++++++++++++ 2 files changed, 320 insertions(+), 57 deletions(-) diff --git a/mobly/base_test.py b/mobly/base_test.py index 07c3877a..ae43c17d 100644 --- a/mobly/base_test.py +++ b/mobly/base_test.py @@ -512,6 +512,14 @@ def _setup_test(self, test_name): with self._log_test_stage(STAGE_NAME_SETUP_TEST): self.setup_test() + def _exec_setup_test(self, test_name): + """Executes _setup_test and converts TestFailure into TestError.""" + try: + self._setup_test(test_name) + except signals.TestFailure as e: + _, _, traceback = sys.exc_info() + raise signals.TestError(e.details, e.extras).with_traceback(traceback) + def setup_test(self): """Setup function that will be called every time before executing each test method in the test class. @@ -528,6 +536,34 @@ def _teardown_test(self, test_name): with self._log_test_stage(STAGE_NAME_TEARDOWN_TEST): self.teardown_test() + def _exec_teardown_test(self, test_name, tr_record): + """Executes _teardown_test and records any errors on the test record.""" + tr_record.update_record() + _, active_exc, _ = sys.exc_info() + test_failed_or_errored = tr_record.result in ( + records.TestResultEnums.TEST_RESULT_FAIL, + records.TestResultEnums.TEST_RESULT_ERROR, + ) + try: + self._teardown_test(test_name) + except signals.TestAbortSignal as e: + if test_failed_or_errored: + tr_record.add_error(STAGE_NAME_TEARDOWN_TEST, e) + else: + tr_record.test_fail(e) + if not isinstance(active_exc, signals.TestAbortSignal) or ( + isinstance(e, signals.TestAbortAll) + and not isinstance(active_exc, signals.TestAbortAll) + ): + raise + except Exception as e: + logging.exception( + 'Exception occurred in %s of %s.', + STAGE_NAME_TEARDOWN_TEST, + self.current_test_info.name, + ) + tr_record.add_error(STAGE_NAME_TEARDOWN_TEST, e) + def teardown_test(self): """Teardown function that will be called every time a test method has been executed. @@ -772,43 +808,9 @@ def exec_one_test(self, test_name, test_method, record=None): ) expects.recorder.reset_internal_states(tr_record) logging.info('%s %s', TEST_CASE_TOKEN, test_name) - # Did teardown_test throw an error. - teardown_test_failed = False try: - try: - try: - self._setup_test(test_name) - except signals.TestFailure as e: - _, _, traceback = sys.exc_info() - raise signals.TestError(e.details, e.extras).with_traceback(traceback) - test_method() - except (signals.TestPass, signals.TestAbortSignal, signals.TestSkip): - raise - except Exception: - logging.exception( - 'Exception occurred in %s.', self.current_test_info.name - ) - raise - finally: - before_count = expects.recorder.error_count - try: - self._teardown_test(test_name) - except signals.TestAbortSignal: - raise - except Exception as e: - logging.exception( - 'Exception occurred in %s of %s.', - STAGE_NAME_TEARDOWN_TEST, - self.current_test_info.name, - ) - tr_record.test_error() - tr_record.add_error(STAGE_NAME_TEARDOWN_TEST, e) - teardown_test_failed = True - else: - # Check if anything failed by `expects`. - if before_count < expects.recorder.error_count: - tr_record.test_error() - teardown_test_failed = True + self._exec_setup_test(test_name) + test_method() except (signals.TestFailure, AssertionError) as e: tr_record.test_fail(e) except signals.TestSkip as e: @@ -823,38 +825,49 @@ def exec_one_test(self, test_name, test_method, record=None): tr_record.test_pass(e) except Exception as e: # Exception happened during test. + logging.exception( + 'Exception occurred in %s.', self.current_test_info.name + ) tr_record.test_error(e) else: - # No exception is thrown from test and teardown, if `expects` has + # No exception is thrown from test, if `expects` has # error, the test should fail with the first error in `expects`. - if expects.recorder.has_error and not teardown_test_failed: + if expects.recorder.has_error: tr_record.test_fail() # Otherwise the test passed. - elif not teardown_test_failed: + else: tr_record.test_pass() finally: - tr_record.update_record() try: - if tr_record.result in ( - records.TestResultEnums.TEST_RESULT_ERROR, - records.TestResultEnums.TEST_RESULT_FAIL, - ): - self._exec_procedure_func(self._on_fail, tr_record) - elif tr_record.result == records.TestResultEnums.TEST_RESULT_PASS: - self._exec_procedure_func(self._on_pass, tr_record) - elif tr_record.result == records.TestResultEnums.TEST_RESULT_SKIP: - self._exec_procedure_func(self._on_skip, tr_record) + self._exec_teardown_test(test_name, tr_record) finally: - logging.info( - RESULT_LINE_TEMPLATE, tr_record.test_name, tr_record.result - ) - self.results.add_record(tr_record) - self.summary_writer.dump( - tr_record.to_dict(), records.TestSummaryEntryType.RECORD - ) - self.current_test_info = None + self._exec_procedure_and_finalize_record(tr_record) return tr_record + def _exec_procedure_and_finalize_record(self, tr_record): + """Executes the on_* procedure for a test and finalizes its record.""" + tr_record.end_time = utils.get_current_epoch_time() + tr_record.update_record() + try: + if tr_record.result in ( + records.TestResultEnums.TEST_RESULT_ERROR, + records.TestResultEnums.TEST_RESULT_FAIL, + ): + self._exec_procedure_func(self._on_fail, tr_record) + elif tr_record.result == records.TestResultEnums.TEST_RESULT_PASS: + self._exec_procedure_func(self._on_pass, tr_record) + elif tr_record.result == records.TestResultEnums.TEST_RESULT_SKIP: + self._exec_procedure_func(self._on_skip, tr_record) + finally: + logging.info( + RESULT_LINE_TEMPLATE, tr_record.test_name, tr_record.result + ) + self.results.add_record(tr_record) + self.summary_writer.dump( + tr_record.to_dict(), records.TestSummaryEntryType.RECORD + ) + self.current_test_info = None + def _assert_function_names_in_stack(self, expected_func_names): """Asserts that the current stack contains any of the given function names.""" current_frame = inspect.currentframe() diff --git a/tests/mobly/base_test_test.py b/tests/mobly/base_test_test.py index 8073ab32..afae6674 100755 --- a/tests/mobly/base_test_test.py +++ b/tests/mobly/base_test_test.py @@ -3152,6 +3152,256 @@ class RecoverableError(Exception): base_test.TEST_STAGE_END_LOG_TEMPLATE, 'TestClass', 'stage' ) + def test_current_test_info_record_populated_in_teardown_test(self): + observed_records = {} + + class MockBaseTest(base_test.BaseTestClass): + + def teardown_test(self): + record = self.current_test_info.record + observed_records[self.current_test_info.name] = ( + record.result, + record.details, + ) + + def test_pass_case(self): + pass + + def test_assert_fail_case(self): + asserts.fail('assert fail') + + def test_expect_fail_case(self): + expects.expect_true(False, 'expect fail') + + def test_error_case(self): + raise RuntimeError('uncaught error') + + def test_skip_case(self): + asserts.skip('skipped') + + bt_cls = MockBaseTest(self.mock_test_cls_configs) + bt_cls.run( + test_names=[ + 'test_pass_case', + 'test_assert_fail_case', + 'test_expect_fail_case', + 'test_error_case', + 'test_skip_case', + ] + ) + self.assertEqual( + observed_records, + { + 'test_pass_case': (records.TestResultEnums.TEST_RESULT_PASS, None), + 'test_assert_fail_case': ( + records.TestResultEnums.TEST_RESULT_FAIL, + 'assert fail', + ), + 'test_expect_fail_case': ( + records.TestResultEnums.TEST_RESULT_FAIL, + 'expect fail', + ), + 'test_error_case': ( + records.TestResultEnums.TEST_RESULT_ERROR, + 'uncaught error', + ), + 'test_skip_case': ( + records.TestResultEnums.TEST_RESULT_SKIP, + 'skipped', + ), + }, + ) + + def test_expect_in_test_and_exception_in_teardown_test(self): + class MockBaseTest(base_test.BaseTestClass): + + def test_func(self): + expects.expect_true(False, MSG_EXPECTED_EXCEPTION, extras=MOCK_EXTRA) + + def teardown_test(self): + raise RuntimeError('teardown error') + + bt_cls = MockBaseTest(self.mock_test_cls_configs) + bt_cls.run(test_names=['test_func']) + self.assertEqual(len(bt_cls.results.failed), 1) + self.assertEqual(len(bt_cls.results.error), 0) + actual_record = bt_cls.results.failed[0] + self.assertEqual(actual_record.test_name, 'test_func') + self.assertEqual(actual_record.details, MSG_EXPECTED_EXCEPTION) + self.assertEqual(actual_record.extras, MOCK_EXTRA) + self.assertEqual( + actual_record.extra_errors['teardown_test'].details, 'teardown error' + ) + + def test_expect_in_test_and_expect_in_teardown_test(self): + class MockBaseTest(base_test.BaseTestClass): + + def test_func(self): + expects.expect_true(False, MSG_EXPECTED_EXCEPTION, extras=MOCK_EXTRA) + + def teardown_test(self): + expects.expect_true(False, 'teardown expect error') + + bt_cls = MockBaseTest(self.mock_test_cls_configs) + bt_cls.run(test_names=['test_func']) + self.assertEqual(len(bt_cls.results.failed), 1) + self.assertEqual(len(bt_cls.results.error), 0) + actual_record = bt_cls.results.failed[0] + self.assertEqual(actual_record.test_name, 'test_func') + self.assertEqual(actual_record.details, MSG_EXPECTED_EXCEPTION) + self.assertEqual(actual_record.extras, MOCK_EXTRA) + self.assertEqual(len(actual_record.extra_errors), 1) + extra_error = next(iter(actual_record.extra_errors.values())) + self.assertEqual(extra_error.details, 'teardown expect error') + + def test_abort_class_in_teardown_test_preserves_test_failure(self): + class MockBaseTest(base_test.BaseTestClass): + + def test_1(self): + asserts.fail(MSG_EXPECTED_EXCEPTION, extras=MOCK_EXTRA) + + def test_2(self): + never_call() + + def teardown_test(self): + asserts.abort_class('abort class in teardown') + + bt_cls = MockBaseTest(self.mock_test_cls_configs) + bt_cls.run(test_names=['test_1', 'test_2']) + self.assertEqual(len(bt_cls.results.failed), 1) + self.assertEqual(len(bt_cls.results.skipped), 1) + actual_record = bt_cls.results.failed[0] + self.assertEqual(actual_record.test_name, 'test_1') + self.assertEqual(actual_record.details, MSG_EXPECTED_EXCEPTION) + self.assertEqual(actual_record.extras, MOCK_EXTRA) + self.assertEqual( + actual_record.extra_errors['teardown_test'].details, + 'abort class in teardown', + ) + + def test_abort_all_in_teardown_test_preserves_test_failure(self): + class MockBaseTest(base_test.BaseTestClass): + + def test_1(self): + asserts.fail(MSG_EXPECTED_EXCEPTION, extras=MOCK_EXTRA) + + def test_2(self): + never_call() + + def teardown_test(self): + asserts.abort_all('abort all in teardown') + + bt_cls = MockBaseTest(self.mock_test_cls_configs) + with self.assertRaisesRegex(signals.TestAbortAll, 'abort all in teardown'): + bt_cls.run(test_names=['test_1', 'test_2']) + self.assertEqual(len(bt_cls.results.failed), 1) + self.assertEqual(len(bt_cls.results.skipped), 1) + actual_record = bt_cls.results.failed[0] + self.assertEqual(actual_record.test_name, 'test_1') + self.assertEqual(actual_record.details, MSG_EXPECTED_EXCEPTION) + self.assertEqual(actual_record.extras, MOCK_EXTRA) + self.assertEqual( + actual_record.extra_errors['teardown_test'].details, + 'abort all in teardown', + ) + + def test_abort_all_in_teardown_test_preserves_test_error(self): + class MockBaseTest(base_test.BaseTestClass): + + def test_1(self): + raise RuntimeError(MSG_EXPECTED_EXCEPTION) + + def test_2(self): + never_call() + + def teardown_test(self): + asserts.abort_all('abort all in teardown') + + bt_cls = MockBaseTest(self.mock_test_cls_configs) + with self.assertRaisesRegex(signals.TestAbortAll, 'abort all in teardown'): + bt_cls.run(test_names=['test_1', 'test_2']) + self.assertEqual(len(bt_cls.results.error), 1) + self.assertEqual(len(bt_cls.results.skipped), 1) + actual_record = bt_cls.results.error[0] + self.assertEqual(actual_record.test_name, 'test_1') + self.assertEqual(actual_record.details, MSG_EXPECTED_EXCEPTION) + self.assertEqual( + actual_record.extra_errors['teardown_test'].details, + 'abort all in teardown', + ) + + def test_abort_class_in_teardown_test_when_test_passes(self): + class MockBaseTest(base_test.BaseTestClass): + + def test_1(self): + pass + + def test_2(self): + never_call() + + def teardown_test(self): + asserts.abort_class('abort class in teardown') + + bt_cls = MockBaseTest(self.mock_test_cls_configs) + bt_cls.run(test_names=['test_1', 'test_2']) + self.assertEqual(len(bt_cls.results.failed), 1) + self.assertEqual(len(bt_cls.results.skipped), 1) + actual_record = bt_cls.results.failed[0] + self.assertEqual(actual_record.test_name, 'test_1') + self.assertEqual(actual_record.details, 'abort class in teardown') + self.assertFalse(actual_record.extra_errors) + + def test_abort_all_in_test_and_nested_abort_class_in_teardown_test(self): + class MockBaseTest(base_test.BaseTestClass): + + def test_1(self): + asserts.abort_all('abort all in test') + + def test_2(self): + never_call() + + def teardown_test(self): + try: + raise RuntimeError('intermediate teardown error') + except RuntimeError: + asserts.abort_class('abort class in teardown') + + bt_cls = MockBaseTest(self.mock_test_cls_configs) + with self.assertRaisesRegex(signals.TestAbortAll, 'abort all in test'): + bt_cls.run(test_names=['test_1', 'test_2']) + self.assertEqual(len(bt_cls.results.failed), 1) + self.assertEqual(len(bt_cls.results.skipped), 1) + actual_record = bt_cls.results.failed[0] + self.assertEqual(actual_record.details, 'abort all in test') + self.assertEqual( + actual_record.extra_errors['teardown_test'].details, + 'abort class in teardown', + ) + + def test_expect_and_abort_class_in_teardown_test_when_test_passes(self): + class MockBaseTest(base_test.BaseTestClass): + + def test_1(self): + pass + + def test_2(self): + never_call() + + def teardown_test(self): + expects.expect_true(False, 'expect error in teardown') + asserts.abort_class('abort class in teardown') + + bt_cls = MockBaseTest(self.mock_test_cls_configs) + bt_cls.run(test_names=['test_1', 'test_2']) + self.assertEqual(len(bt_cls.results.failed), 1) + self.assertEqual(len(bt_cls.results.skipped), 1) + actual_record = bt_cls.results.failed[0] + self.assertEqual(actual_record.details, 'abort class in teardown') + self.assertEqual(len(actual_record.extra_errors), 1) + extra_error = next(iter(actual_record.extra_errors.values())) + self.assertEqual(extra_error.details, 'expect error in teardown') + if __name__ == '__main__': unittest.main() + From 2049e2dea385c5052a14c5a38b03a8f8baa0d9fe Mon Sep 17 00:00:00 2001 From: Ang Li Date: Wed, 30 Sep 2026 05:04:20 +0000 Subject: [PATCH 2/3] Fix formatting --- mobly/base_test.py | 4 +--- tests/mobly/base_test_test.py | 1 - 2 files changed, 1 insertion(+), 4 deletions(-) diff --git a/mobly/base_test.py b/mobly/base_test.py index ae43c17d..b7a220f4 100644 --- a/mobly/base_test.py +++ b/mobly/base_test.py @@ -859,9 +859,7 @@ def _exec_procedure_and_finalize_record(self, tr_record): elif tr_record.result == records.TestResultEnums.TEST_RESULT_SKIP: self._exec_procedure_func(self._on_skip, tr_record) finally: - logging.info( - RESULT_LINE_TEMPLATE, tr_record.test_name, tr_record.result - ) + logging.info(RESULT_LINE_TEMPLATE, tr_record.test_name, tr_record.result) self.results.add_record(tr_record) self.summary_writer.dump( tr_record.to_dict(), records.TestSummaryEntryType.RECORD diff --git a/tests/mobly/base_test_test.py b/tests/mobly/base_test_test.py index afae6674..8fb65d67 100755 --- a/tests/mobly/base_test_test.py +++ b/tests/mobly/base_test_test.py @@ -3404,4 +3404,3 @@ def teardown_test(self): if __name__ == '__main__': unittest.main() - From e9f5133845d8e496fe89405447469e99866a66ba Mon Sep 17 00:00:00 2001 From: Ang Li Date: Wed, 30 Sep 2026 18:24:16 +0000 Subject: [PATCH 3/3] Add unit tests for setup_test override and teardown abort edge cases --- tests/mobly/base_test_test.py | 70 +++++++++++++++++++++++++++++++++++ 1 file changed, 70 insertions(+) diff --git a/tests/mobly/base_test_test.py b/tests/mobly/base_test_test.py index 8fb65d67..9fca837f 100755 --- a/tests/mobly/base_test_test.py +++ b/tests/mobly/base_test_test.py @@ -3401,6 +3401,76 @@ def teardown_test(self): extra_error = next(iter(actual_record.extra_errors.values())) self.assertEqual(extra_error.details, 'expect error in teardown') + def test_private_setup_test_override_fail_by_test_signal(self): + class MockBaseTest(base_test.BaseTestClass): + + def _setup_test(self, test_name): + asserts.fail(MSG_EXPECTED_EXCEPTION) + super()._setup_test(test_name) + + def test_something(self): + never_call() + + bt_cls = MockBaseTest(self.mock_test_cls_configs) + bt_cls.run(test_names=['test_something']) + self.assertEqual(len(bt_cls.results.error), 1) + self.assertEqual(len(bt_cls.results.failed), 0) + actual_record = bt_cls.results.error[0] + self.assertEqual(actual_record.test_name, self.mock_test_name) + self.assertEqual(actual_record.details, MSG_EXPECTED_EXCEPTION) + + def test_abort_class_in_test_and_abort_class_in_teardown_test(self): + class MockBaseTest(base_test.BaseTestClass): + + def test_1(self): + asserts.abort_class('abort class in test') + + def test_2(self): + never_call() + + def teardown_test(self): + asserts.abort_class('abort class in teardown') + + bt_cls = MockBaseTest(self.mock_test_cls_configs) + bt_cls.run(test_names=['test_1', 'test_2']) + self.assertEqual(len(bt_cls.results.failed), 1) + self.assertEqual(len(bt_cls.results.skipped), 1) + failed_record = bt_cls.results.failed[0] + self.assertEqual(failed_record.details, 'abort class in test') + self.assertEqual( + failed_record.extra_errors['teardown_test'].details, + 'abort class in teardown', + ) + skipped_record = bt_cls.results.skipped[0] + self.assertEqual( + skipped_record.details, + 'Test class aborted due to: abort class in test', + ) + + def test_abort_class_in_test_escalated_to_abort_all_in_teardown_test(self): + class MockBaseTest(base_test.BaseTestClass): + + def test_1(self): + asserts.abort_class('abort class in test') + + def test_2(self): + never_call() + + def teardown_test(self): + asserts.abort_all('abort all in teardown') + + bt_cls = MockBaseTest(self.mock_test_cls_configs) + with self.assertRaisesRegex(signals.TestAbortAll, 'abort all in teardown'): + bt_cls.run(test_names=['test_1', 'test_2']) + self.assertEqual(len(bt_cls.results.failed), 1) + self.assertEqual(len(bt_cls.results.skipped), 1) + failed_record = bt_cls.results.failed[0] + self.assertEqual(failed_record.details, 'abort class in test') + self.assertEqual( + failed_record.extra_errors['teardown_test'].details, + 'abort all in teardown', + ) + if __name__ == '__main__': unittest.main()