diff --git a/mobly/base_test.py b/mobly/base_test.py index 78651ba8..47face95 100644 --- a/mobly/base_test.py +++ b/mobly/base_test.py @@ -416,21 +416,25 @@ def _setup_class(self): 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 - ) + try: + self._exec_procedure_func(self._on_fail, class_record) + finally: + class_record.update_record() + self.summary_writer.dump( + class_record.to_dict(), records.TestSummaryEntryType.RECORD + ) 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) + try: + self._exec_procedure_func(self._on_fail, class_record) + finally: + 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._skip_remaining_tests(class_record.termination_signal.exception) return self.results diff --git a/tests/mobly/base_test_test.py b/tests/mobly/base_test_test.py index 1b43dcaf..8f0de3a8 100755 --- a/tests/mobly/base_test_test.py +++ b/tests/mobly/base_test_test.py @@ -1375,6 +1375,45 @@ def on_fail(self, record): 'Error 1, Executed 0, Failed 0, Passed 0, Requested 3, Skipped 3', ) + with open(self.summary_file) as summary: + setup_records = [ + entry + for entry in yaml.safe_load_all(summary) + if entry.get('Test Name') == 'setup_class' + ] + self.assertEqual(len(setup_records), 1) + self.assertEqual(setup_records[0]['Result'], 'ERROR') + self.assertEqual(setup_records[0]['Details'], MSG_UNEXPECTED_EXCEPTION) + self.assertIsNotNone(setup_records[0]['End Time']) + + def test_abort_all_in_on_fail_from_setup_class_expect(self): + class MockBaseTest(base_test.BaseTestClass): + + def setup_class(self): + expects.expect_true(False, MSG_UNEXPECTED_EXCEPTION) + + def test_1(self): + never_call() + + def on_fail(self, record): + asserts.abort_all(MSG_EXPECTED_EXCEPTION) + + bt_cls = MockBaseTest(self.mock_test_cls_configs) + with self.assertRaisesRegex(signals.TestAbortAll, MSG_EXPECTED_EXCEPTION): + bt_cls.run(test_names=['test_1']) + self.assertEqual(len(bt_cls.results.error), 1) + self.assertEqual(len(bt_cls.results.skipped), 1) + with open(self.summary_file) as summary: + setup_records = [ + entry + for entry in yaml.safe_load_all(summary) + if entry.get('Test Name') == 'setup_class' + ] + self.assertEqual(len(setup_records), 1) + self.assertEqual(setup_records[0]['Result'], 'ERROR') + self.assertEqual(setup_records[0]['Details'], MSG_UNEXPECTED_EXCEPTION) + self.assertIsNotNone(setup_records[0]['End Time']) + def test_abort_all_in_test(self): class MockBaseTest(base_test.BaseTestClass):