diff --git a/cxplat/cxplat_test/cxplat_fault_injection_test.cpp b/cxplat/cxplat_test/cxplat_fault_injection_test.cpp index 8ed8ec5..8bd7afe 100644 --- a/cxplat/cxplat_test/cxplat_fault_injection_test.cpp +++ b/cxplat/cxplat_test/cxplat_fault_injection_test.cpp @@ -102,6 +102,21 @@ TEST_CASE("fault_injection", "[fault_injection]") // Clear the fault injection state. cxplat_fault_injection_reset(); + auto inject_fault_from_same_callsite = []() { return cxplat_fault_injection_inject_fault(); }; + + cxplat_fault_injection_suspend(); + REQUIRE(cxplat_fault_injection_is_enabled() == true); + REQUIRE(inject_fault_from_same_callsite() == false); + + cxplat_fault_injection_suspend(); + cxplat_fault_injection_resume(); + REQUIRE(inject_fault_from_same_callsite() == false); + + cxplat_fault_injection_resume(); + REQUIRE(inject_fault_from_same_callsite() == true); + + cxplat_fault_injection_reset(); + for (_fault_injection_expected_outcome state = _fault_injection_expected_outcome::ExpectFault; state <= _fault_injection_expected_outcome::ExpectFaultDifferentCallsite; state = (_fault_injection_expected_outcome)((int)state + 1)) { diff --git a/cxplat/inc/winuser/cxplat_fault_injection.h b/cxplat/inc/winuser/cxplat_fault_injection.h index 5de0257..2bad899 100644 --- a/cxplat/inc/winuser/cxplat_fault_injection.h +++ b/cxplat/inc/winuser/cxplat_fault_injection.h @@ -6,6 +6,8 @@ #ifndef CXPLAT_DEBUGGING_FEATURES_ENABLED #define cxplat_fault_injection_is_enabled() false #define cxplat_fault_injection_inject_fault() false +#define cxplat_fault_injection_suspend() ((void)0) +#define cxplat_fault_injection_resume() ((void)0) #else #include @@ -50,6 +52,20 @@ cxplat_fault_injection_inject_fault() CXPLAT_NOEXCEPT; bool cxplat_fault_injection_is_enabled() CXPLAT_NOEXCEPT; +/** + * @brief Suspend fault injection. Calls may be nested and must be matched by calls to + * cxplat_fault_injection_resume(). This function is thread safe. + */ +void +cxplat_fault_injection_suspend() CXPLAT_NOEXCEPT; + +/** + * @brief Resume fault injection after a matching call to cxplat_fault_injection_suspend(). + * This function is thread safe. + */ +void +cxplat_fault_injection_resume() CXPLAT_NOEXCEPT; + /** * @brief Reset fault injection. This function is thread safe. */ diff --git a/cxplat/src/cxplat_winuser/fault_injection.cpp b/cxplat/src/cxplat_winuser/fault_injection.cpp index 04fd548..5b91989 100644 --- a/cxplat/src/cxplat_winuser/fault_injection.cpp +++ b/cxplat/src/cxplat_winuser/fault_injection.cpp @@ -11,6 +11,7 @@ #include #include #include +#include #include #include #include @@ -40,6 +41,12 @@ typedef class _cxplat_fault_injection bool inject_fault(); + void + suspend(); + + void + resume(); + /** * @brief Reset the fault injection state, both in memory and on disk. */ @@ -139,6 +146,9 @@ typedef class _cxplat_fault_injection */ std::mutex _mutex; + std::shared_mutex _suspension_mutex; + _Guarded_by_(_suspension_mutex) size_t _suspension_count = 0; + size_t _stack_depth; _Guarded_by_(_mutex) std::vector _last_fault_stack; @@ -209,9 +219,30 @@ _cxplat_fault_injection::~_cxplat_fault_injection() bool _cxplat_fault_injection::inject_fault() { + std::shared_lock lock(_suspension_mutex); + if (_suspension_count > 0) { + return false; + } return is_new_stack(); } +void +_cxplat_fault_injection::suspend() +{ + std::unique_lock lock(_suspension_mutex); + _suspension_count++; +} + +void +_cxplat_fault_injection::resume() +{ + std::unique_lock lock(_suspension_mutex); + CXPLAT_RUNTIME_ASSERT(_suspension_count > 0); + if (_suspension_count > 0) { + _suspension_count--; + } +} + void _cxplat_fault_injection::reset() { @@ -419,6 +450,22 @@ cxplat_fault_injection_is_enabled() noexcept return _cxplat_fault_injection_singleton != nullptr; } +void +cxplat_fault_injection_suspend() noexcept +{ + if (_cxplat_fault_injection_singleton) { + _cxplat_fault_injection_singleton->suspend(); + } +} + +void +cxplat_fault_injection_resume() noexcept +{ + if (_cxplat_fault_injection_singleton) { + _cxplat_fault_injection_singleton->resume(); + } +} + void cxplat_fault_injection_reset() noexcept {