From adfbcf2e095230d290a8f72a4b8cf714a87c675d Mon Sep 17 00:00:00 2001 From: Ang Li Date: Sat, 8 Aug 2026 07:43:09 +0000 Subject: [PATCH] Support using expects.expect_no_raises as a decorator --- mobly/expects.py | 69 +++++++++++++---- tests/mobly/base_test_test.py | 32 ++++++++ tests/mobly/expects_test.py | 137 ++++++++++++++++++++++++++++++++++ 3 files changed, 222 insertions(+), 16 deletions(-) create mode 100644 tests/mobly/expects_test.py diff --git a/mobly/expects.py b/mobly/expects.py index ad7cb8af5..c2b0dd727 100644 --- a/mobly/expects.py +++ b/mobly/expects.py @@ -13,6 +13,7 @@ # limitations under the License. import contextlib +import functools import logging import time @@ -135,31 +136,67 @@ def expect_equal(first, second, msg=None, extras=None): recorder.add_error(e) -@contextlib.contextmanager -def expect_no_raises(message=None, extras=None): - """Expects no exception is raised in a context. +class expect_no_raises(contextlib.ContextDecorator): + """Expects no exception is raised in a context or a decorated function. If the expectation is not met, the test is marked as fail after its execution finishes. - A default message is added to the exception `details`. + Can be used as a context manager: + with expects.expect_no_raises(message='Custom message'): + do_something() + + Or as a function decorator: + @expects.expect_no_raises(message='Custom message') + def helper_function(arg): + do_something(arg) + + @expects.expect_no_raises + def bare_decorated_function(): + do_something() Args: - message: string, custom message to add to exception's `details`. + message: string or callable, custom message to add to exception's `details`. + When used as a bare decorator (@expects.expect_no_raises), this argument + is the decorated function. extras: An optional field for extra information to be included in test result. """ - try: - yield - except Exception as e: - e_record = records.ExceptionRecord(e) - if extras: - e_record.extras = extras - msg = message or 'Got an unexpected exception' - details = '%s: %s' % (msg, e_record.details) - logging.exception(details) - e_record.details = details - recorder.add_error(e_record) + + def __new__(cls, *args, **kwargs): + if len(args) == 1 and callable(args[0]) and not kwargs: + # Used as bare decorator: @expects.expect_no_raises + func = args[0] + instance = super().__new__(cls) + instance.__init__() + + @functools.wraps(func) + def wrapped(*f_args, **f_kwargs): + with instance: + return func(*f_args, **f_kwargs) + + return wrapped + return super().__new__(cls) + + def __init__(self, message=None, extras=None): + self._message = message + self._extras = extras + + def __enter__(self): + return self + + def __exit__(self, exc_type, exc_val, exc_tb): + if exc_type is not None and issubclass(exc_type, Exception): + e_record = records.ExceptionRecord(exc_val) + if self._extras: + e_record.extras = self._extras + msg = self._message or 'Got an unexpected exception' + details = '%s: %s' % (msg, e_record.details) + logging.exception(details) + e_record.details = details + recorder.add_error(e_record) + return True + return False recorder = _ExpectErrorRecorder(DEFAULT_TEST_RESULT_RECORD) diff --git a/tests/mobly/base_test_test.py b/tests/mobly/base_test_test.py index d74ed0540..b226d081a 100755 --- a/tests/mobly/base_test_test.py +++ b/tests/mobly/base_test_test.py @@ -2565,6 +2565,38 @@ def test_something(self): for i, record in enumerate(bt_cls.results.passed): self.assertEqual(record.test_name, f'test_something_{i}') + def test_expect_no_raises_decorator_on_test_method(self): + class MockBaseTest(base_test.BaseTestClass): + + @expects.expect_no_raises(message='Step 1 failed') + def test_step(self): + raise ValueError('Expected failure in step') + + bt_cls = MockBaseTest(self.mock_test_cls_configs) + bt_cls.run() + self.assertEqual(1, len(bt_cls.results.failed)) + self.assertEqual(1, len(bt_cls.results.executed)) + record = bt_cls.results.failed[0] + self.assertIn('Step 1 failed', record.details) + self.assertIn('Expected failure in step', record.details) + + def test_expect_no_raises_decorator_on_helper_method(self): + class MockBaseTest(base_test.BaseTestClass): + + @expects.expect_no_raises + def helper_func(self): + raise RuntimeError('Helper failed') + + def test_method(self): + self.helper_func() + + bt_cls = MockBaseTest(self.mock_test_cls_configs) + bt_cls.run() + self.assertEqual(1, len(bt_cls.results.failed)) + self.assertEqual(1, len(bt_cls.results.executed)) + record = bt_cls.results.failed[0] + self.assertIn('Helper failed', record.details) + def test_repeat_with_consec_error_does_not_abort_repeat(self): repeat_count = 5 max_consec_error = 2 diff --git a/tests/mobly/expects_test.py b/tests/mobly/expects_test.py new file mode 100644 index 000000000..ba79fad94 --- /dev/null +++ b/tests/mobly/expects_test.py @@ -0,0 +1,137 @@ +# Copyright 2017 Google Inc. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +import unittest + +from mobly import expects +from mobly import records + + +class ExpectsTest(unittest.TestCase): + """Unit tests for the mobly.expects module.""" + + def setUp(self): + self.record = records.TestResultRecord('test_foo', 'TestClass') + expects.recorder.reset_internal_states(self.record) + + def test_expect_true_pass(self): + expects.expect_true(True, 'Should pass') + self.assertFalse(expects.recorder.has_error) + self.assertEqual(expects.recorder.error_count, 0) + + def test_expect_true_fail(self): + expects.expect_true(False, 'Expected true', extras='extra_info') + self.assertTrue(expects.recorder.has_error) + self.assertEqual(expects.recorder.error_count, 1) + err = list(self.record.extra_errors.values())[0] + self.assertEqual(err.extras, 'extra_info') + self.assertIn('Expected true', err.details) + + def test_expect_false_pass(self): + expects.expect_false(False, 'Should pass') + self.assertFalse(expects.recorder.has_error) + self.assertEqual(expects.recorder.error_count, 0) + + def test_expect_false_fail(self): + expects.expect_false(True, 'Expected false', extras='extra_info') + self.assertTrue(expects.recorder.has_error) + self.assertEqual(expects.recorder.error_count, 1) + err = list(self.record.extra_errors.values())[0] + self.assertEqual(err.extras, 'extra_info') + self.assertIn('Expected false', err.details) + + def test_expect_equal_pass(self): + expects.expect_equal(1, 1, 'Should pass') + self.assertFalse(expects.recorder.has_error) + self.assertEqual(expects.recorder.error_count, 0) + + def test_expect_equal_fail(self): + expects.expect_equal(1, 2, 'Values not equal', extras='extra_info') + self.assertTrue(expects.recorder.has_error) + self.assertEqual(expects.recorder.error_count, 1) + err = list(self.record.extra_errors.values())[0] + self.assertEqual(err.extras, 'extra_info') + self.assertIn('Values not equal', err.details) + + def test_expect_no_raises_context_manager_pass(self): + with expects.expect_no_raises(): + _ = 1 + 1 + self.assertFalse(expects.recorder.has_error) + self.assertEqual(expects.recorder.error_count, 0) + + def test_expect_no_raises_context_manager_fail(self): + with expects.expect_no_raises(message='Context error', extras='extra_info'): + raise ValueError('something went wrong') + self.assertTrue(expects.recorder.has_error) + self.assertEqual(expects.recorder.error_count, 1) + err = list(self.record.extra_errors.values())[0] + self.assertEqual(err.extras, 'extra_info') + self.assertIn('Context error', err.details) + self.assertIn('something went wrong', err.details) + + def test_expect_no_raises_bare_decorator_on_arbitrary_function(self): + """Verifies @expects.expect_no_raises on any arbitrary function.""" + + @expects.expect_no_raises + def arbitrary_helper(a, b, fail=False): + if fail: + raise RuntimeError('Helper failed') + return a + b + + # Success case on arbitrary function + result = arbitrary_helper(3, 4, fail=False) + self.assertEqual(result, 7) + self.assertFalse(expects.recorder.has_error) + + # Failure case on arbitrary function + result = arbitrary_helper(3, 4, fail=True) + self.assertIsNone(result) + self.assertTrue(expects.recorder.has_error) + self.assertEqual(expects.recorder.error_count, 1) + + def test_expect_no_raises_parameterized_decorator_on_arbitrary_function(self): + """Verifies @expects.expect_no_raises(...) on any arbitrary function.""" + + @expects.expect_no_raises(message='Custom step error', extras={'step': 1}) + def custom_calculation(x, y): + if y == 0: + raise ZeroDivisionError('divide by zero') + return x / y + + # Success case + self.assertEqual(custom_calculation(10, 2), 5) + self.assertFalse(expects.recorder.has_error) + + # Failure case + self.assertIsNone(custom_calculation(10, 0)) + self.assertTrue(expects.recorder.has_error) + self.assertEqual(expects.recorder.error_count, 1) + err = list(self.record.extra_errors.values())[0] + self.assertEqual(err.extras, {'step': 1}) + self.assertIn('Custom step error', err.details) + + def test_expect_no_raises_decorator_preserves_function_metadata(self): + """Verifies that docstrings, __name__, and metadata are preserved.""" + + @expects.expect_no_raises + def sample_func(x): + """Sample documentation.""" + return x + + self.assertEqual(sample_func.__name__, 'sample_func') + self.assertEqual(sample_func.__doc__, 'Sample documentation.') + + +if __name__ == '__main__': + unittest.main()