diff --git a/mobly/base_test.py b/mobly/base_test.py index 78651ba8..07c3877a 100644 --- a/mobly/base_test.py +++ b/mobly/base_test.py @@ -371,12 +371,27 @@ def _pre_run(self): return True except Exception as e: logging.exception('%s failed for %s.', stage_name, self.TAG) - record.test_error(e) + self._record_class_error(record, e) + return False + + def _record_class_error(self, record, e=None, exec_on_fail=False): + """Finalizes a class-level error record and writes it to results/summary. + + Args: + record: A TestResultRecord object for the class stage. + e: Optional exception that caused the failure. + exec_on_fail: bool, whether to execute `_on_fail` procedure. + """ + record.test_error(e) + record.update_record() + try: + if exec_on_fail: + self._exec_procedure_func(self._on_fail, record) + finally: self.results.add_class_error(record) self.summary_writer.dump( record.to_dict(), records.TestSummaryEntryType.RECORD ) - return False def pre_run(self): """Preprocesses that need to be done before setup_class. @@ -414,23 +429,11 @@ def _setup_class(self): # Setup class failed for unknown reasons. # Fail the class and skip all tests. logging.exception('Error in %s#setup_class.', self.TAG) - class_record.test_error(e) - self.results.add_class_error(class_record) - self._exec_procedure_func(self._on_fail, class_record) - class_record.update_record() - self.summary_writer.dump( - class_record.to_dict(), records.TestSummaryEntryType.RECORD - ) + self._record_class_error(class_record, e, exec_on_fail=True) self._skip_remaining_tests(e) return self.results if expects.recorder.has_error: - self._exec_procedure_func(self._on_fail, class_record) - class_record.test_error() - class_record.update_record() - self.summary_writer.dump( - class_record.to_dict(), records.TestSummaryEntryType.RECORD - ) - self.results.add_class_error(class_record) + self._record_class_error(class_record, exec_on_fail=True) self._skip_remaining_tests(class_record.termination_signal.exception) return self.results @@ -464,20 +467,10 @@ def _teardown_class(self): raise except Exception as e: logging.exception('Error encountered in %s.', stage_name) - record.test_error(e) - record.update_record() - self.results.add_class_error(record) - self.summary_writer.dump( - record.to_dict(), records.TestSummaryEntryType.RECORD - ) + self._record_class_error(record, e) else: if expects.recorder.has_error: - record.test_error() - record.update_record() - self.results.add_class_error(record) - self.summary_writer.dump( - record.to_dict(), records.TestSummaryEntryType.RECORD - ) + self._record_class_error(record) finally: self._clean_up() @@ -1192,9 +1185,4 @@ def _clean_up(self): self._record_controller_info() self._controller_manager.unregister_controllers() if expects.recorder.has_error: - record.test_error() - record.update_record() - self.results.add_class_error(record) - self.summary_writer.dump( - record.to_dict(), records.TestSummaryEntryType.RECORD - ) + self._record_class_error(record) diff --git a/tests/mobly/base_test_test.py b/tests/mobly/base_test_test.py index 1b43dcaf..8073ab32 100755 --- a/tests/mobly/base_test_test.py +++ b/tests/mobly/base_test_test.py @@ -1344,7 +1344,10 @@ def on_fail(self, record): 'Error 0, Executed 1, Failed 1, Passed 0, Requested 3, Skipped 2', ) - def test_abort_all_in_on_fail_from_setup_class(self): + @mock.patch('mobly.records.TestSummaryWriter.dump') + def test_abort_all_in_on_fail_from_setup_class(self, mock_dump): + on_fail_record_state = {} + class MockBaseTest(base_test.BaseTestClass): def setup_class(self): @@ -1360,6 +1363,8 @@ def test_3(self): never_call() def on_fail(self, record): + on_fail_record_state['end_time'] = record.end_time + on_fail_record_state['details'] = record.details asserts.abort_all(MSG_EXPECTED_EXCEPTION) bt_cls = MockBaseTest(self.mock_test_cls_configs) @@ -1367,6 +1372,8 @@ def on_fail(self, record): signals.TestAbortAll, MSG_EXPECTED_EXCEPTION ) as context: bt_cls.run(test_names=['test_1', 'test_2', 'test_3']) + self.assertIsNotNone(on_fail_record_state['end_time']) + self.assertEqual(on_fail_record_state['details'], MSG_UNEXPECTED_EXCEPTION) setup_class_record = bt_cls.results.error[0] self.assertEqual(setup_class_record.test_name, 'setup_class') self.assertTrue(hasattr(context.exception, 'results')) @@ -1374,6 +1381,50 @@ def on_fail(self, record): bt_cls.results.summary_str(), 'Error 1, Executed 0, Failed 0, Passed 0, Requested 3, Skipped 3', ) + setup_class_dict = mock_dump.call_args_list[1][0][0] + self.assertEqual(setup_class_dict['Test Name'], 'setup_class') + self.assertEqual(setup_class_dict['Result'], 'ERROR') + + @mock.patch('mobly.records.TestSummaryWriter.dump') + def test_abort_all_in_on_fail_from_setup_class_with_expects(self, mock_dump): + on_fail_record_state = {} + + class MockBaseTest(base_test.BaseTestClass): + + def setup_class(self): + expects.expect_true(False, MSG_UNEXPECTED_EXCEPTION) + + def test_1(self): + never_call() + + def test_2(self): + never_call() + + def test_3(self): + never_call() + + def on_fail(self, record): + on_fail_record_state['end_time'] = record.end_time + on_fail_record_state['details'] = record.details + asserts.abort_all(MSG_EXPECTED_EXCEPTION) + + bt_cls = MockBaseTest(self.mock_test_cls_configs) + with self.assertRaisesRegex( + signals.TestAbortAll, MSG_EXPECTED_EXCEPTION + ) as context: + bt_cls.run(test_names=['test_1', 'test_2', 'test_3']) + self.assertIsNotNone(on_fail_record_state['end_time']) + self.assertEqual(on_fail_record_state['details'], MSG_UNEXPECTED_EXCEPTION) + setup_class_record = bt_cls.results.error[0] + self.assertEqual(setup_class_record.test_name, 'setup_class') + self.assertTrue(hasattr(context.exception, 'results')) + self.assertEqual( + bt_cls.results.summary_str(), + 'Error 1, Executed 0, Failed 0, Passed 0, Requested 3, Skipped 3', + ) + setup_class_dict = mock_dump.call_args_list[1][0][0] + self.assertEqual(setup_class_dict['Test Name'], 'setup_class') + self.assertEqual(setup_class_dict['Result'], 'ERROR') def test_abort_all_in_test(self): class MockBaseTest(base_test.BaseTestClass): @@ -1777,6 +1828,7 @@ def test_2(self): def test_expect_in_setup_class(self, mock_dump): must_call = mock.Mock() must_call2 = mock.Mock() + on_fail_record_state = {} class MockBaseTest(base_test.BaseTestClass): @@ -1788,12 +1840,16 @@ def test_func(self): pass def on_fail(self, record): + on_fail_record_state['end_time'] = record.end_time + on_fail_record_state['details'] = record.details must_call2('on_fail') bt_cls = MockBaseTest(self.mock_test_cls_configs) bt_cls.run() must_call.assert_called_once_with('ha') must_call2.assert_called_once_with('on_fail') + self.assertIsNotNone(on_fail_record_state['end_time']) + self.assertEqual(on_fail_record_state['details'], MSG_EXPECTED_EXCEPTION) actual_record = bt_cls.results.error[0] self.assertEqual(actual_record.test_name, 'setup_class') self.assertEqual(actual_record.details, MSG_EXPECTED_EXCEPTION)