diff --git a/.github/workflows/memory-leak-analysis.yml b/.github/workflows/memory-leak-analysis.yml index 682316aa0..982ca149e 100644 --- a/.github/workflows/memory-leak-analysis.yml +++ b/.github/workflows/memory-leak-analysis.yml @@ -95,6 +95,22 @@ jobs: -BaselinePath .github/memory-leak-baseline.csv -TargetArguments "--gtest_filter=-OfflineStorageTests_SQLite.StoreThousandEventsTakesLessThanASecond" + - name: Analyze modern and disabled network detection + shell: pwsh + run: | + ./.github/scripts/run-drmemory.ps1 ` + -DrMemoryPath "$env:RUNNER_TEMP/DrMemory-Windows-$env:DRMEMORY_VERSION/bin64/drmemory.exe" ` + -LogDirectory drmemory-results ` + -Scenario network-native ` + -TargetPath Solutions/out/Debug/x64/UnitTests/UnitTests.exe ` + -TargetArguments "--gtest_filter=NetworkDetectorTests.StartsReadsCostAndStopsWithoutNetworkListManager:NetworkDetectorTests.ConfigurationDisablesDetectionWithoutLoadingNetworkListManager" + $summary = @(Import-Csv drmemory-results/summary.csv | Where-Object Scenario -eq "network-native") + if ($summary.Count -ne 1 -or + [int64]$summary[0].TotalLeaks -ne 0 -or + [int64]$summary[0].TotalPossibleLeaks -ne 0) { + throw "Modern and disabled network detection must have zero actual and possible leaks." + } + - name: Analyze functional tests shell: pwsh run: >- @@ -116,11 +132,13 @@ jobs: -BaselinePath .github/memory-leak-baseline.csv -TargetPath Solutions/out/Debug/x64/SampleCppMini/SampleCppMini.exe - - name: Verify Network List Manager is not loaded + - name: Verify modern production paths do not load Network List Manager shell: pwsh run: | $moduleLogs = @() - foreach ($scenario in @("unit-tests", "functional-tests", "sample-cpp-mini")) { + # The full unit suite deliberately exercises the legacy COM fallback too. + # Keep its leak baseline unchanged; inspect the native backend separately. + foreach ($scenario in @("network-native", "functional-tests", "sample-cpp-mini")) { $scenarioLogs = @(Get-ChildItem "drmemory-results/$scenario" -Filter global.*.log -File -Recurse) if ($scenarioLogs.Count -eq 0) { throw "Dr. Memory did not produce a module log for $scenario." diff --git a/README.md b/README.md index 078c4c553..096dd5621 100644 --- a/README.md +++ b/README.md @@ -89,26 +89,28 @@ Other resources to learn how to setup the build system: ## Target Platforms - | Target Platform | Supported | Covered by CI | - | ------------------------------ | ------------------ | ------------------ | - | Android (API 23+) | :white_check_mark: | :white_check_mark: | - | iOS 12+ (simulator) | :white_check_mark: | :white_check_mark: | - | iOS 12+ (arm64, arm64e) | :white_check_mark: | | - | Linux (x86, x64, arm, aarch64) | :white_check_mark: | | - | macOS 10.15+ | :white_check_mark: | | - | macOS (latest) | :white_check_mark: | :white_check_mark: | - | Ubuntu 20.04.x LTS | :white_check_mark: | :white_check_mark: | - | Ubuntu 22.04.x LTS | :white_check_mark: | :white_check_mark: | - | Ubuntu (latest) | :white_check_mark: | :white_check_mark: | - | Windows 10.x | :white_check_mark: | | - | Windows 11 | :white_check_mark: | | - | Windows Server 2016 | :white_check_mark: | | - | Windows Server 2019 | :white_check_mark: | | - | Windows Server 2022 | :white_check_mark: | :white_check_mark: | + | Target Platform | Supported | Covered by CI | + | ------------------------------ | ------------------ | ------------- | + | Android (API 23+) | :white_check_mark: | [Native builds](.github/workflows/test-embedding.yml), [Gradle build and Java unit tests](.github/workflows/build-android.yml); no device runtime tests | + | iOS 12+ (simulator) | :white_check_mark: | [Debug/Release simulator tests](.github/workflows/build-ios-mac.yml) on current runtimes | + | iOS 12+ (arm64, arm64e) | :white_check_mark: | [arm64 cross-builds](.github/workflows/test-embedding.yml) with deployment target 13.0; no arm64e or device runtime tests | + | Linux (x86, x64, arm, aarch64) | :white_check_mark: | [x64 build and tests](.github/workflows/build-posix-latest.yml); no x86/arm/aarch64 jobs | + | macOS 10.15+ | :white_check_mark: | No macOS 10.15 runner | + | macOS (latest) | :white_check_mark: | [Debug/Release tests](.github/workflows/build-posix-latest.yml), [arm64/universal builds](.github/workflows/test-embedding.yml) | + | Ubuntu 20.04.x LTS | :white_check_mark: | No current runner | + | Ubuntu 22.04.x LTS | :white_check_mark: | [Debug/Release tests](.github/workflows/build-ubuntu-2204.yml), [Dr. Memory](.github/workflows/memory-leak-analysis.yml) | + | Ubuntu (latest) | :white_check_mark: | [Debug/Release tests](.github/workflows/build-posix-latest.yml), [embedding/package tests](.github/workflows/test-embedding.yml) | + | Windows 10.x | :white_check_mark: | [API-floor compile checks](.github/workflows/test-win-latest.yml) on Server 2022; no Windows 10 runner | + | Windows 11 | :white_check_mark: | No Windows 11 runner | + | Windows Server 2016 | :white_check_mark: | No Server 2016 runner | + | Windows Server 2019 | :white_check_mark: | No Server 2019 runner | + | Windows Server 2022 | :white_check_mark: | [Win32/x64 Debug/Release tests](.github/workflows/test-win-latest.yml), [Dr. Memory](.github/workflows/memory-leak-analysis.yml) | * **Supported** - these platforms are known to work well with the SDK in production. -* **Covered by CI** - these platforms are tested as part of CI. +* **Covered by CI** - current GitHub Actions coverage, distinguishing builds, + runtime tests, and specific runner OS versions. Cross-builds and API-floor + checks do not establish runtime coverage on every supported OS version. * Windows 7, Windows 8, and Windows 8.1 are not supported. Windows desktop builds target the Windows 10 API floor in CI. * For iOS simulator, CI covers representative supported simulator diff --git a/docs/building-custom-SKU.md b/docs/building-custom-SKU.md index b3a03d394..44006c833 100644 --- a/docs/building-custom-SKU.md +++ b/docs/building-custom-SKU.md @@ -28,12 +28,62 @@ Build recipe must contain the following preprocessor definitions: | HAVE_MAT_WIN_LOG | off | Will log statements to disk on windows if trace enabled and HAVE_MAT_LOGGING defined | | HAVE_MAT_EVT_TRACEID | off | Enable event tracking by adding trace-id to http request header on Windows. This is for debugging purpose, and not recommended to be enabled in production. The collector doesn't parse/read this header. As of now, this is meant to be used through the capi, where the http-send handler should remove this header from the event data before sending it to collector. | | HAVE_MAT_STORAGE | on | Enable SQLite persistent offline storage | -| HAVE_MAT_NETDETECT | on | _Win32 Desktop only_: Use Windows Runtime APIs for network cost detection on Windows 10+ | +| HAVE_MAT_NETDETECT | on | _Win32 Desktop only_: Use IP Helper connectivity hints on Windows 10 version 2004+; use native Network List Manager COM APIs on older supported Windows versions | | HAVE_MAT_SHORT_NS | off | Use short "MAT::" namespace instead of "Microsoft::Applications::Events::" to reduce the .DLL size | | HAVE_CS4 | off | Build with Common Schema 4.0 support. Current default is `off`, i.e. building with Common Schema 3.0 support | | HAVE_CS4_FULL | off | Enable additional Common Schema 4.0 protocol features needed by server / services SDK | | COMPACT_SDK | off | Built-in build recipe for smallest possible SDK. Turns most features off. Includes_mat/config-compact.h_ | +### Windows desktop network cost lifecycle + +Build with a Windows SDK that declares `NL_NETWORK_CONNECTIVITY_HINT` in `nldef.h`. + +The IP Helper APIs are resolved at runtime, preserving the existing Windows 10 +and Windows Server 2016 minimum rather than adding newer loader imports. The +modern backend reports aggregate connectivity hints, not just the WinRT Internet +connection profile. Roaming and approaching/exceeded data limits map to the +restrictive `NetworkCost_Roaming` category. Connectivity hints do not expose +WinRT's separate background-data restriction flag. + +The fallback activates `INetworkListManager` on a private SDK-owned STA and +preserves the three original event families: network-list connectivity, network +properties, and connection properties. `INetworkCostManager` is queried only +as an optional capability. An unsupported cost interface reports +`NetworkCost_Unknown` while connectivity/property monitoring remains active, +matching the behavior before the WinRT-only detector change. It is not a startup +failure and is not treated as an unmetered connection. No cost-specific event +interface is required. + +Base NLM is documented for Windows Vista/Server 2008 onward; cost querying is +documented for Windows 8 clients with no supported Server versions. The fallback +therefore does not require Windows 8 cost support or WinRT on Windows 7 SP1 or +Server 2008 R2. This describes detector API coverage, not a change to the SDK's +overall support policy or compiler/runtime requirements. + +This fallback loads `netprofm.dll`; the modern backend does not. Subscription teardown, interface +release, and balanced COM shutdown happen before joining the listener thread. +The host does not need to initialize COM or retain an MTA across SDK DLL reloads. +Callback dispatch is drained on external stop; a reentrant stop does not wait on +itself. Restart and subsequent external stops still drain the previous callback. +Explicit COM disconnection failures are logged, and the non-agile sink's owning +STA still completes `CoUninitialize`, which closes its RPC connections, before +the thread is joined. This does not terminate the host. Native notification +cancellation failure remains fatal because there is no COM apartment rundown +to provide that safety guarantee. + +Consumers embedding the SDK in an unloadable library can avoid both network +backends by setting `CFG_BOOL_ENABLE_NET_DETECT` to `false` in the +`ILogConfiguration` passed to SDK initialization. The desktop implementation +does not construct or start a detector in that configuration. This preserves +the existing disabled-detection behavior (unmetered cost), rather than the +unknown cost returned by an enabled detector without cost information. +The setting avoids detector-originated COM/NLM activity; it is not a guarantee +that every SDK component or host dependency is leak-free. + +For build-time exclusion, omit `HAVE_MAT_NETDETECT` from a custom SDK recipe. +Defining it as `0` does not disable the feature because it uses presence-based +preprocessor checks. + ## Building custom SDK SKU: MSBuild example Command: diff --git a/lib/callbacks/DebugSource.cpp b/lib/callbacks/DebugSource.cpp index cc70009d3..aec31393d 100644 --- a/lib/callbacks/DebugSource.cpp +++ b/lib/callbacks/DebugSource.cpp @@ -15,7 +15,7 @@ namespace MAT_NS_BEGIN { namespace { - thread_local std::vector pendingListeners; + thread_local std::vector* pendingListeners = nullptr; std::atomic pendingReleaseCallback{nullptr}; @@ -23,14 +23,23 @@ namespace MAT_NS_BEGIN { { public: explicit PendingListenersScope(const std::vector& listeners) : - remaining(listeners) + remaining(listeners), + previous(pendingListeners) { - pendingListeners.insert( - pendingListeners.end(), + if (previous != nullptr) + { + pending = *previous; + } + pending.insert( + pending.end(), listeners.begin(), listeners.end()); + pendingListeners = &pending; } + PendingListenersScope(const PendingListenersScope&) = delete; + PendingListenersScope& operator=(const PendingListenersScope&) = delete; + ~PendingListenersScope() { for (auto listener : remaining) @@ -42,6 +51,7 @@ namespace MAT_NS_BEGIN { callback(listener); } } + pendingListeners = previous; } void BeginCallback(DebugEventListener* listener) @@ -55,28 +65,31 @@ namespace MAT_NS_BEGIN { } private: - static void RemovePending(DebugEventListener* listener) + void RemovePending(DebugEventListener* listener) { - auto pending = std::find( - pendingListeners.rbegin(), - pendingListeners.rend(), + auto entry = std::find( + pending.rbegin(), + pending.rend(), listener); - if (pending != pendingListeners.rend()) + if (entry != pending.rend()) { - pendingListeners.erase(std::next(pending).base()); + pending.erase(std::next(entry).base()); } } std::vector remaining; + std::vector pending; + std::vector* previous; }; } bool IsDebugEventListenerPending(const DebugEventListener* listener) noexcept { - return std::find( - pendingListeners.begin(), - pendingListeners.end(), - listener) != pendingListeners.end(); + return pendingListeners != nullptr && + std::find( + pendingListeners->begin(), + pendingListeners->end(), + listener) != pendingListeners->end(); } void SetDebugEventListenerPendingReleaseCallback( diff --git a/lib/callbacks/DebugSourceInternal.hpp b/lib/callbacks/DebugSourceInternal.hpp index 5adfe75bd..35aa6c3bc 100644 --- a/lib/callbacks/DebugSourceInternal.hpp +++ b/lib/callbacks/DebugSourceInternal.hpp @@ -11,6 +11,7 @@ namespace MAT_NS_BEGIN using DebugEventListenerPendingReleaseCallback = void (*)(DebugEventListener*); + // Pending storage is dispatch-scope-owned; querying outside dispatch does not allocate. bool IsDebugEventListenerPending(const DebugEventListener* listener) noexcept; void SetDebugEventListenerPendingReleaseCallback( DebugEventListenerPendingReleaseCallback callback) noexcept; diff --git a/lib/pal/desktop/NetworkDetector.cpp b/lib/pal/desktop/NetworkDetector.cpp index 4694b5091..7f46596ae 100644 --- a/lib/pal/desktop/NetworkDetector.cpp +++ b/lib/pal/desktop/NetworkDetector.cpp @@ -6,9 +6,15 @@ #include "mat/config.h" #ifdef HAVE_MAT_NETDETECT -#pragma comment(lib, "runtimeobject.lib") +#pragma comment(lib, "iphlpapi.lib") +#pragma comment(lib, "ole32.lib") +#include +#include #include "NetworkDetector.hpp" +#include +#include +#include #include "DebugEvents.hpp" #include "ILogManager.hpp" @@ -25,18 +31,41 @@ namespace MAT_NS_BEGIN struct NetworkDetector::CallbackState { - std::atomic listenerThreadId{0}; + mutable std::mutex mutex; + DWORD listenerThreadId = 0; + + void SetListenerThreadId(DWORD threadId) + { + std::lock_guard lock(mutex); + listenerThreadId = threadId; + } bool QueueRefresh() const { - const auto threadId = listenerThreadId.load(std::memory_order_acquire); - return threadId != 0 && - PostThreadMessage(threadId, NETDETECTOR_REFRESH, 0, NULL) != FALSE; + std::lock_guard lock(mutex); + if (listenerThreadId == 0) + { + return false; + } + if (!PostThreadMessage(listenerThreadId, NETDETECTOR_REFRESH, 0, NULL)) + { + LOG_ERROR("Unable to post a network cost refresh: %lu.", GetLastError()); + return false; + } + return true; } }; struct NetworkDetector::EventDispatchState : std::enable_shared_from_this { + ~EventDispatchState() + { + if (work != nullptr) + { + CloseThreadpoolWork(work); + } + } + bool Queue(NetworkCost cost) { std::lock_guard lock(mutex); @@ -52,37 +81,63 @@ namespace MAT_NS_BEGIN return true; } - workerScheduled = true; - auto context = new (std::nothrow) std::shared_ptr(shared_from_this()); - if (context == nullptr || - !QueueUserWorkItem(DispatchPendingEvents, context, WT_EXECUTEDEFAULT)) + if (work == nullptr) { - delete context; - workerScheduled = false; - return false; + work = CreateThreadpoolWork(DispatchPendingEvents, this, nullptr); + if (work == nullptr) + { + LOG_ERROR("Unable to create network callback work: %lu.", GetLastError()); + return false; + } } + workerScheduled = true; + SubmitThreadpoolWork(work); return true; } - void StopAndWait() + void Stop() { - std::unique_lock lock(mutex); + std::lock_guard lock(mutex); acceptEvents = false; eventPending = false; + } + + void Wait() + { if (currentNetworkEventDispatch == this) { return; } - cv.wait(lock, [this]() - { return !workerScheduled; }); + if (work != nullptr) + { + WaitForThreadpoolWorkCallbacks(work, FALSE); + } } private: - static DWORD CALLBACK DispatchPendingEvents(void* context) + static void CALLBACK DispatchPendingEvents(PTP_CALLBACK_INSTANCE instance, void* context, PTP_WORK) { std::shared_ptr state = - *static_cast*>(context); - delete static_cast*>(context); + static_cast(context)->shared_from_this(); + HMODULE module = nullptr; + const auto address = reinterpret_cast(&DispatchPendingEvents); + if (!GetModuleHandleExW(GET_MODULE_HANDLE_EX_FLAG_FROM_ADDRESS | + GET_MODULE_HANDLE_EX_FLAG_UNCHANGED_REFCOUNT, + address, &module) || + (module != GetModuleHandleW(nullptr) && + !GetModuleHandleExW(GET_MODULE_HANDLE_EX_FLAG_FROM_ADDRESS, address, &module))) + { + LOG_ERROR("Unable to retain the network callback module: %lu.", GetLastError()); + std::lock_guard lock(state->mutex); + state->workerScheduled = false; + state->eventPending = false; + return; + } + if (module != GetModuleHandleW(nullptr)) + { + // Reentrant Stop can return on this callback; release the DLL reference after return. + FreeLibraryWhenCallbackReturns(instance, module); + } currentNetworkEventDispatch = state.get(); while (true) @@ -93,9 +148,8 @@ namespace MAT_NS_BEGIN if (!state->acceptEvents || !state->eventPending) { state->workerScheduled = false; - state->cv.notify_all(); currentNetworkEventDispatch = nullptr; - return 0; + return; } cost = state->latestCost; state->eventPending = false; @@ -110,7 +164,7 @@ namespace MAT_NS_BEGIN } std::mutex mutex; - std::condition_variable cv; + PTP_WORK work = nullptr; NetworkCost latestCost = NetworkCost_Unknown; bool acceptEvents = true; bool eventPending = false; @@ -118,113 +172,137 @@ namespace MAT_NS_BEGIN }; NetworkCost MapNetworkCost( - NetworkCostType costType, - boolean roaming, - boolean overDataLimit, - boolean approachingDataLimit, - boolean backgroundDataUsageRestricted) + NL_NETWORK_CONNECTIVITY_COST_HINT costType, + bool roaming, + bool overDataLimit, + bool approachingDataLimit) { - if (roaming || overDataLimit || approachingDataLimit || backgroundDataUsageRestricted) + if (roaming || overDataLimit || approachingDataLimit) { return NetworkCost_Roaming; } switch (costType) { - case NetworkCostType_Unrestricted: + case NetworkConnectivityCostHintUnrestricted: return NetworkCost_Unmetered; - case NetworkCostType_Fixed: - case NetworkCostType_Variable: + case NetworkConnectivityCostHintFixed: + case NetworkConnectivityCostHintVariable: return NetworkCost_Metered; - case NetworkCostType_Unknown: + case NetworkConnectivityCostHintUnknown: default: return NetworkCost_Unknown; } } - static NetworkCost QueryCurrentNetworkCost(INetworkInformationStatics* networkInfoStats) + NetworkCost MapLegacyNetworkCost(DWORD cost) { - NetworkCost result = NetworkCost_Unknown; - LOG_TRACE("get network cost...\n"); - - if (networkInfoStats == nullptr) + if ((cost & (NLM_CONNECTION_COST_ROAMING | NLM_CONNECTION_COST_OVERDATALIMIT | + NLM_CONNECTION_COST_APPROACHINGDATALIMIT | NLM_CONNECTION_COST_CONGESTED)) != 0) { - LOG_WARN("Windows network information is unavailable!"); - return result; + return NetworkCost_Roaming; } - - ComPtr connectionProfile; - HRESULT hr = networkInfoStats->GetInternetConnectionProfile(&connectionProfile); - if (FAILED(hr) || connectionProfile == nullptr) + if ((cost & NLM_CONNECTION_COST_UNRESTRICTED) != 0) { - return result; + return NetworkCost_Unmetered; } - - ComPtr connectionCost; - hr = connectionProfile->GetConnectionCost(&connectionCost); - if (FAILED(hr) || connectionCost == nullptr) + if ((cost & (NLM_CONNECTION_COST_FIXED | NLM_CONNECTION_COST_VARIABLE)) != 0) { - return result; + return NetworkCost_Metered; } + return NetworkCost_Unknown; + } - boolean roaming = false; - boolean overDataLimit = false; - boolean approachingDataLimit = false; - boolean backgroundDataUsageRestricted = false; - NetworkCostType costType = NetworkCostType_Unknown; - if (FAILED(connectionCost->get_Roaming(&roaming)) || - FAILED(connectionCost->get_OverDataLimit(&overDataLimit)) || - FAILED(connectionCost->get_ApproachingDataLimit(&approachingDataLimit)) || - FAILED(connectionCost->get_NetworkCostType(&costType))) + NetworkDetector::NetworkDetector() + { + const auto module = GetModuleHandleW(L"iphlpapi.dll"); + getConnectivityHint = reinterpret_cast( + GetProcAddress(module, "GetNetworkConnectivityHint")); + notifyConnectivityHint = reinterpret_cast( + GetProcAddress(module, "NotifyNetworkConnectivityHintChange")); + if (getConnectivityHint == nullptr || notifyConnectivityHint == nullptr) { - return result; + getConnectivityHint = nullptr; + notifyConnectivityHint = nullptr; } + } - ComPtr connectionCost2; - if (SUCCEEDED(connectionCost.As(&connectionCost2)) && - FAILED(connectionCost2->get_BackgroundDataUsageRestricted(&backgroundDataUsageRestricted))) + NetworkCost NetworkDetector::QueryNetworkCost() + { + if (getConnectivityHint != nullptr) { - return result; + NL_NETWORK_CONNECTIVITY_HINT hint {}; + const auto error = getConnectivityHint(&hint); + if (error != NO_ERROR) + { + LOG_ERROR("Unable to query network connectivity cost: %lu.", error); + return NetworkCost_Unknown; + } + return MapNetworkCost(hint.ConnectivityCost, hint.Roaming != FALSE, + hint.OverDataLimit != FALSE, hint.ApproachingDataLimit != FALSE); } - - return MapNetworkCost( - costType, - roaming, - overDataLimit, - approachingDataLimit, - backgroundDataUsageRestricted); + if (networkCostManager == nullptr) + { + return NetworkCost_Unknown; + } + DWORD cost = NLM_CONNECTION_COST_UNKNOWN; + const auto hr = networkCostManager->GetCost(&cost, nullptr); + if (FAILED(hr)) + { + LOG_ERROR("Unable to query legacy network cost: 0x%08lx.", hr); + return NetworkCost_Unknown; + } + return MapLegacyNetworkCost(cost); } - /// - /// Get current realtime network cost synchronously. - /// This function provides an SEH handler for Windows Runtime failures. - /// - static int RefreshNetworkCost( - INetworkInformationStatics* networkInfoStats, - std::atomic& currentNetworkCostState) + struct NetworkDetector::NetworkStatusChangedSink : + RuntimeClass, INetworkListManagerEvents, + INetworkEvents, INetworkConnectionEvents> { - NetworkCost currentNetworkCost = NetworkCost_Unknown; - __try + explicit NetworkStatusChangedSink(std::shared_ptr state) : state(std::move(state)) { - currentNetworkCost = QueryCurrentNetworkCost(networkInfoStats); } - //****************************************************************************************************************************** - // This code is required as a workaround for an issue in Visual Studio debug host mode: crash in W.N.C.dll - // - // onecoreuap\net\netprofiles\winrt\networkinformation\lib\handlemanager.cpp(132)\Windows.Networking.Connectivity.dll!0FBCFB9E: - // (caller: 0FBCEE2C) ReturnHr(1) tid(4584) 80070426 The service has not been started. - // - // Exception thrown at XXX (KernelBase.dll) in YYY : The binding handle is invalid. - // If there is a handler for this exception, the program may be safely continued. - //******************************************************************************************************************************* -#pragma warning(suppress : 6320) - __except (EXCEPTION_EXECUTE_HANDLER) + HRESULT STDMETHODCALLTYPE ConnectivityChanged(NLM_CONNECTIVITY) override { - LOG_ERROR("Unable to obtain network state!"); + state->QueueRefresh(); + return S_OK; } + HRESULT STDMETHODCALLTYPE NetworkAdded(GUID) override + { + state->QueueRefresh(); + return S_OK; + } + HRESULT STDMETHODCALLTYPE NetworkDeleted(GUID) override + { + state->QueueRefresh(); + return S_OK; + } + HRESULT STDMETHODCALLTYPE NetworkConnectivityChanged(GUID, NLM_CONNECTIVITY) override + { + state->QueueRefresh(); + return S_OK; + } + HRESULT STDMETHODCALLTYPE NetworkPropertyChanged(GUID, NLM_NETWORK_PROPERTY_CHANGE) override + { + state->QueueRefresh(); + return S_OK; + } + HRESULT STDMETHODCALLTYPE NetworkConnectionConnectivityChanged(GUID, NLM_CONNECTIVITY) override + { + state->QueueRefresh(); + return S_OK; + } + HRESULT STDMETHODCALLTYPE NetworkConnectionPropertyChanged(GUID, NLM_CONNECTION_PROPERTY_CHANGE) override + { + state->QueueRefresh(); + return S_OK; + } + std::shared_ptr state; + }; - currentNetworkCostState.store(currentNetworkCost, std::memory_order_relaxed); - return currentNetworkCost; + void WINAPI NetworkDetector::NetworkHintChanged(void* context, NL_NETWORK_CONNECTIVITY_HINT) + { + static_cast(context)->QueueRefresh(); } NetworkCost NetworkDetector::GetNetworkCost() { @@ -233,12 +311,37 @@ namespace MAT_NS_BEGIN int NetworkDetector::GetCurrentNetworkCost() { - const auto currentNetworkCost = - RefreshNetworkCost(networkInfoStats.Get(), *m_currentNetworkCost); + { + std::unique_lock lock(m_lock); + if (m_listener_tid != GetCurrentThreadId()) + { + if (startupState != StartupState::Ready || stopRequested) + { + return GetNetworkCost(); + } + const auto sequence = refreshSequence; + const auto callbackState = networkStatusCallbackState; + if (!callbackState->QueueRefresh()) + { + LOG_ERROR("Unable to request a network cost refresh."); + return GetNetworkCost(); + } + cv.wait(lock, [this, sequence, callbackState]() { + return refreshSequence != sequence || stopRequested || + startupState != StartupState::Ready || + networkStatusCallbackState != callbackState; + }); + return GetNetworkCost(); + } + } + const auto currentNetworkCost = QueryNetworkCost(); + m_currentNetworkCost->store(currentNetworkCost, std::memory_order_relaxed); std::shared_ptr dispatchState; { std::lock_guard lock(m_lock); dispatchState = eventDispatchState; + ++refreshSequence; + cv.notify_all(); } if (dispatchState != nullptr && !dispatchState->Queue(static_cast(currentNetworkCost))) { @@ -258,49 +361,99 @@ namespace MAT_NS_BEGIN } /// - /// Get activation factory and look-up network info statistics + /// Initialize the network cost backend on its owning thread /// /// - bool NetworkDetector::GetNetworkInfoStats() + bool NetworkDetector::InitializeNetworkCost() { - HRESULT hr = GetActivationFactory(HString::MakeReference(RuntimeClass_Windows_Networking_Connectivity_NetworkInformation).Get(), &networkInfoStats); - if (hr != S_OK) + if (getConnectivityHint != nullptr) + { + return true; + } + auto hr = CoCreateInstance(CLSID_NetworkListManager, nullptr, CLSCTX_ALL, + IID_PPV_ARGS(networkListManager.GetAddressOf())); + if (FAILED(hr)) { - LOG_ERROR("Unable to get Windows::Networking::Connectivity::NetworkInformation"); + LOG_ERROR("Unable to initialize the legacy network list manager: 0x%08lx.", hr); return false; } + hr = queryLegacyCost(networkListManager.Get(), networkCostManager.GetAddressOf()); + if (FAILED(hr)) + { + networkCostManager.Reset(); + LOG_WARN("Legacy network cost information is unavailable (0x%08lx); monitoring connectivity with Unknown cost.", hr); + } return true; } + HRESULT WINAPI NetworkDetector::QueryLegacyCostInterface( + INetworkListManager* manager, INetworkCostManager** cost) + { + return manager->QueryInterface(IID_PPV_ARGS(cost)); + } + + HRESULT WINAPI NetworkDetector::FindLegacyConnectionPoint( + IConnectionPointContainer* container, REFIID iid, IConnectionPoint** point) + { + return container->FindConnectionPoint(iid, point); + } + bool NetworkDetector::RegisterAndListen() noexcept { MSG msg; PeekMessage(&msg, nullptr, WM_USER, WM_USER, PM_NOREMOVE); const auto callbackState = networkStatusCallbackState; - callbackState->listenerThreadId.store(GetCurrentThreadId(), std::memory_order_release); - networkStatusChangedHandler = Callback( - [callbackState](IInspectable*) -> HRESULT - { - callbackState->QueueRefresh(); - return S_OK; - }); - if (networkStatusChangedHandler == nullptr) + callbackState->SetListenerThreadId(GetCurrentThreadId()); + if (notifyConnectivityHint != nullptr) { - LOG_ERROR("Unable to create network status handler."); - callbackState->listenerThreadId.store(0, std::memory_order_release); - return false; + const auto error = notifyConnectivityHint( + NetworkHintChanged, callbackState.get(), FALSE, &networkStatusNotification); + if (error != NO_ERROR) + { + LOG_ERROR("Unable to subscribe to network connectivity changes: %lu.", error); + return false; + } } - - HRESULT hr = networkInfoStats->add_NetworkStatusChanged( - networkStatusChangedHandler.Get(), - &networkStatusChangedToken); - if (FAILED(hr)) + else { - LOG_ERROR("Unable to subscribe to network status changes."); - callbackState->listenerThreadId.store(0, std::memory_order_release); - networkStatusChangedHandler.Reset(); - return false; + auto sink = Make(callbackState); + if (sink == nullptr) + { + LOG_ERROR("Unable to create a legacy network status handler."); + return false; + } + auto hr = sink.As(&networkStatusChangedHandler); + ComPtr container; + if (SUCCEEDED(hr)) + { + hr = networkListManager.As(&container); + } + if (FAILED(hr)) + { + LOG_ERROR("Unable to obtain the legacy network event container: 0x%08lx.", hr); + return false; + } + const IID interfaces[] = { + __uuidof(INetworkListManagerEvents), + __uuidof(INetworkEvents), + __uuidof(INetworkConnectionEvents) + }; + for (size_t index = 0; index < legacySubscriptions.size(); ++index) + { + auto& subscription = legacySubscriptions[index]; + hr = findLegacyPoint(container.Get(), interfaces[index], subscription.point.GetAddressOf()); + if (SUCCEEDED(hr)) + { + hr = subscription.point->Advise(networkStatusChangedHandler.Get(), &subscription.cookie); + subscription.subscribed = SUCCEEDED(hr); + } + if (FAILED(hr)) + { + LOG_ERROR("Unable to subscribe to legacy network event %zu: 0x%08lx.", index, hr); + return false; + } + } } { @@ -365,58 +518,81 @@ namespace MAT_NS_BEGIN { if (networkStatusCallbackState != nullptr) { - networkStatusCallbackState->listenerThreadId.store(0, std::memory_order_release); + networkStatusCallbackState->SetListenerThreadId(0); } - if (networkStatusChangedToken.value != 0 && networkInfoStats != nullptr) + if (networkStatusNotification != nullptr) { - const auto token = networkStatusChangedToken; - networkStatusChangedToken.value = 0; - networkInfoStats->remove_NetworkStatusChanged(token); + const auto error = CancelMibChangeNotify2(networkStatusNotification); + if (error != NO_ERROR) + { + LOG_ERROR("Unable to cancel network connectivity notifications: %lu.", error); + std::terminate(); + } + networkStatusNotification = nullptr; } - networkStatusChangedHandler.Reset(); - networkInfoStats.Reset(); - } - - /// - /// Register for Windows Runtime events and block-wait in RegisterAndListen - /// - void NetworkDetector::run() - { - bool isRoInitialized = false; - - __try + for (auto& subscription : legacySubscriptions) { - __try + if (subscription.subscribed) { - HRESULT hr = RoInitialize(RO_INIT_MULTITHREADED); + const auto hr = subscription.point->Unadvise(subscription.cookie); if (FAILED(hr)) { - LOG_ERROR("RoInitialize failed."); - return; - } - - isRoInitialized = true; - if (GetNetworkInfoStats()) - { - RefreshNetworkCost(networkInfoStats.Get(), *m_currentNetworkCost); - LOG_TRACE("start listening to events..."); - RegisterAndListen(); + LOG_ERROR("Unable to unsubscribe from legacy network changes: 0x%08lx.", hr); } + subscription.subscribed = false; } - __finally + } + if (networkStatusChangedHandler != nullptr) + { + const auto hr = disconnectLegacyHandler(networkStatusChangedHandler.Get(), 0); + if (FAILED(hr)) { - Reset(); + LOG_ERROR("Unable to disconnect the legacy network handler: 0x%08lx.", hr); } } -#pragma warning(suppress : 6320) - __except (EXCEPTION_EXECUTE_HANDLER) + // The non-agile sink also disconnects when the owning STA completes CoUninitialize. + networkStatusChangedHandler.Reset(); + for (auto& subscription : legacySubscriptions) { - LOG_ERROR("Handled exception in Windows Runtime network cost detection."); + subscription.point.Reset(); } + networkCostManager.Reset(); + networkListManager.Reset(); + } - if (isRoInitialized) + /// + /// Own the network backend and notifications for the listener thread + /// + void NetworkDetector::run() + { + const bool useCom = getConnectivityHint == nullptr; + if (useCom) + { + const auto hr = CoInitializeEx(nullptr, COINIT_APARTMENTTHREADED); + if (FAILED(hr)) + { + LOG_ERROR("Unable to initialize the legacy network COM apartment: 0x%08lx.", hr); + return; + } + } + struct Cleanup + { + NetworkDetector& detector; + bool useCom; + ~Cleanup() + { + detector.Reset(); + if (useCom) + { + CoUninitialize(); + } + } + } cleanup { *this, useCom }; + if (InitializeNetworkCost()) { - RoUninitialize(); + m_currentNetworkCost->store(QueryNetworkCost(), std::memory_order_relaxed); + LOG_TRACE("start listening to events..."); + RegisterAndListen(); } } /// @@ -425,7 +601,7 @@ namespace MAT_NS_BEGIN /// true - if start is successful, false - otherwise bool NetworkDetector::Start() { - std::lock_guard lifecycleLock(m_lifecycleLock); + std::unique_lock lifecycleLock(m_lifecycleLock); { std::unique_lock lock(m_lock); if (startupState == StartupState::Starting) @@ -439,6 +615,24 @@ namespace MAT_NS_BEGIN return true; } + while (eventDispatchState != nullptr) + { + const auto previousDispatch = eventDispatchState; + lock.unlock(); + lifecycleLock.unlock(); + previousDispatch->Wait(); + lifecycleLock.lock(); + lock.lock(); + if (startupState == StartupState::Ready) + { + return true; + } + if (eventDispatchState == previousDispatch) + { + break; + } + } + lock.unlock(); if (netDetectThread.joinable()) { @@ -506,12 +700,17 @@ namespace MAT_NS_BEGIN } if (!started) { - std::lock_guard lock(m_lock); - CloseHandle(stopEvent); - stopEvent = nullptr; - networkStatusCallbackState.reset(); - eventDispatchState->StopAndWait(); - eventDispatchState.reset(); + std::shared_ptr dispatchState; + { + std::lock_guard lock(m_lock); + CloseHandle(stopEvent); + stopEvent = nullptr; + networkStatusCallbackState.reset(); + dispatchState = std::move(eventDispatchState); + } + dispatchState->Stop(); + lifecycleLock.unlock(); + dispatchState->Wait(); } return started; } @@ -522,15 +721,17 @@ namespace MAT_NS_BEGIN /// void NetworkDetector::Stop() { - std::lock_guard lifecycleLock(m_lifecycleLock); + std::unique_lock lifecycleLock(m_lifecycleLock); if (netDetectThread.joinable()) { { std::lock_guard lock(m_lock); stopRequested = true; + cv.notify_all(); + eventDispatchState->Stop(); if (networkStatusCallbackState != nullptr) { - networkStatusCallbackState->listenerThreadId.store(0, std::memory_order_release); + networkStatusCallbackState->SetListenerThreadId(0); } if (!SetEvent(stopEvent)) { @@ -539,16 +740,21 @@ namespace MAT_NS_BEGIN } netDetectThread.join(); - eventDispatchState->StopAndWait(); - - std::lock_guard lock(m_lock); - CloseHandle(stopEvent); - stopEvent = nullptr; - startupState = StartupState::Stopped; - stopRequested = false; - networkStatusCallbackState.reset(); - eventDispatchState.reset(); - LOG_TRACE("NetworkDetector tid=%p has stopped.", m_listener_tid); + { + std::lock_guard lock(m_lock); + CloseHandle(stopEvent); + stopEvent = nullptr; + startupState = StartupState::Stopped; + stopRequested = false; + networkStatusCallbackState.reset(); + } + } + const auto dispatchState = eventDispatchState; + lifecycleLock.unlock(); + if (dispatchState != nullptr) + { + dispatchState->Stop(); + dispatchState->Wait(); } }; diff --git a/lib/pal/desktop/NetworkDetector.hpp b/lib/pal/desktop/NetworkDetector.hpp index 2b4080e04..1752781b1 100644 --- a/lib/pal/desktop/NetworkDetector.hpp +++ b/lib/pal/desktop/NetworkDetector.hpp @@ -17,10 +17,11 @@ #include #include -#include -#include +#include +#include #include +#include #include #include #include @@ -29,26 +30,26 @@ #include "Enums.hpp" using namespace Microsoft::WRL; -using namespace Microsoft::WRL::Wrappers; -using namespace ABI::Windows::Foundation; -using namespace ABI::Windows::Networking::Connectivity; namespace MAT_NS_BEGIN { namespace Windows { NetworkCost MapNetworkCost( - NetworkCostType costType, - boolean roaming, - boolean overDataLimit, - boolean approachingDataLimit, - boolean backgroundDataUsageRestricted); + NL_NETWORK_CONNECTIVITY_COST_HINT costType, + bool roaming, + bool overDataLimit, + bool approachingDataLimit); + + NetworkCost MapLegacyNetworkCost(DWORD cost); class NetworkDetector { private: struct CallbackState; struct EventDispatchState; + struct NetworkStatusChangedSink; + friend class NetworkDetectorTestAccess; enum class StartupState { Stopped, @@ -60,9 +61,29 @@ namespace MAT_NS_BEGIN /// /// Current network info stats /// - ComPtr networkInfoStats; - ComPtr networkStatusChangedHandler; - EventRegistrationToken networkStatusChangedToken{}; + using GetConnectivityHint = DWORD(WINAPI*)(NL_NETWORK_CONNECTIVITY_HINT*); + using HintChangedCallback = void(WINAPI*)(void*, NL_NETWORK_CONNECTIVITY_HINT); + using NotifyConnectivityHint = DWORD(WINAPI*)(HintChangedCallback, void*, BOOLEAN, HANDLE*); + GetConnectivityHint getConnectivityHint = nullptr; + NotifyConnectivityHint notifyConnectivityHint = nullptr; + HANDLE networkStatusNotification = nullptr; + ComPtr networkListManager; + ComPtr networkCostManager; + struct LegacySubscription + { + ComPtr point; + DWORD cookie = 0; + bool subscribed = false; + }; + std::array legacySubscriptions; + ComPtr networkStatusChangedHandler; + using QueryLegacyCost = HRESULT(WINAPI*)(INetworkListManager*, INetworkCostManager**); + using FindLegacyPoint = HRESULT(WINAPI*)(IConnectionPointContainer*, REFIID, IConnectionPoint**); + static HRESULT WINAPI QueryLegacyCostInterface(INetworkListManager*, INetworkCostManager**); + static HRESULT WINAPI FindLegacyConnectionPoint(IConnectionPointContainer*, REFIID, IConnectionPoint**); + QueryLegacyCost queryLegacyCost = QueryLegacyCostInterface; + FindLegacyPoint findLegacyPoint = FindLegacyConnectionPoint; + decltype(&CoDisconnectObject) disconnectLegacyHandler = CoDisconnectObject; std::shared_ptr networkStatusCallbackState; std::shared_ptr eventDispatchState; @@ -70,7 +91,9 @@ namespace MAT_NS_BEGIN /// Get instance of network info stats /// /// - bool GetNetworkInfoStats(); + bool InitializeNetworkCost(); + NetworkCost QueryNetworkCost(); + static void WINAPI NetworkHintChanged(void* context, NL_NETWORK_CONNECTIVITY_HINT hint); std::mutex m_lifecycleLock; std::mutex m_lock; @@ -79,6 +102,7 @@ namespace MAT_NS_BEGIN std::thread netDetectThread; StartupState startupState = StartupState::Stopped; bool stopRequested = false; + uint64_t refreshSequence = 0; HANDLE stopEvent = nullptr; /// @@ -113,7 +137,7 @@ namespace MAT_NS_BEGIN /// /// Createa network status listener /// - NetworkDetector() = default; + NetworkDetector(); /// /// @@ -144,7 +168,7 @@ namespace MAT_NS_BEGIN NetworkCost GetNetworkCost(); /// - /// Queue the same refresh performed by a WinRT network status callback. + /// Queue the same refresh performed by a network status callback. /// bool QueueNetworkCostRefresh(); }; diff --git a/tests/CMakeLists.txt b/tests/CMakeLists.txt index 216590ebd..ecc1ccb44 100644 --- a/tests/CMakeLists.txt +++ b/tests/CMakeLists.txt @@ -35,4 +35,7 @@ if(MATSDK_BUILD_UNIT_TESTS) target_include_directories(matsdk_test_config INTERFACE ${CMAKE_CURRENT_SOURCE_DIR}/unittests) add_subdirectory(unittests) + if(MSVC AND NOT BUILD_SHARED_LIBS) + add_subdirectory(dll-unload) + endif() endif() diff --git a/tests/common/network-detector-test-access.hpp b/tests/common/network-detector-test-access.hpp new file mode 100644 index 000000000..4b5d61dca --- /dev/null +++ b/tests/common/network-detector-test-access.hpp @@ -0,0 +1,124 @@ +// +// Copyright (c) Microsoft Corporation. All rights reserved. +// SPDX-License-Identifier: Apache-2.0 +// +#pragma once + +#include "pal/desktop/NetworkDetector.hpp" + +#ifdef HAVE_MAT_NETDETECT +namespace MAT_NS_BEGIN +{ + namespace Windows + { + class NetworkDetectorTestAccess + { + public: + static bool HasNativeBackend(const NetworkDetector& detector) + { + return detector.getConnectivityHint != nullptr && detector.notifyConnectivityHint != nullptr; + } + + static void UseLegacyBackend(NetworkDetector& detector) + { + detector.getConnectivityHint = nullptr; + detector.notifyConnectivityHint = nullptr; + detector.queryLegacyCost = NetworkDetector::QueryLegacyCostInterface; + detector.findLegacyPoint = NetworkDetector::FindLegacyConnectionPoint; + detector.disconnectLegacyHandler = CoDisconnectObject; + } + + static void UseLegacyBackendWithoutCost(NetworkDetector& detector) + { + UseLegacyBackend(detector); + detector.queryLegacyCost = UnsupportedCost; + } + + static bool HasLegacyCost(const NetworkDetector& detector) + { + return detector.networkCostManager != nullptr; + } + + static bool HasLegacyManager(const NetworkDetector& detector) + { + return detector.networkListManager != nullptr; + } + + static size_t LegacySubscriptionCount(const NetworkDetector& detector) + { + size_t count = 0; + for (const auto& subscription : detector.legacySubscriptions) + { + count += subscription.subscribed ? 1 : 0; + } + return count; + } + + static void FailLegacyCostQuery(NetworkDetector& detector) + { + UseLegacyBackend(detector); + detector.queryLegacyCost = RejectCostQuery; + } + + static void FailLegacySubscription(NetworkDetector& detector) + { + UseLegacyBackend(detector); + detector.findLegacyPoint = RejectConnectionEvents; + } + + static void FailLegacyDisconnect(NetworkDetector& detector) + { + UseLegacyBackend(detector); + detector.disconnectLegacyHandler = RejectDisconnect; + } + + static void FailNativeSubscription(NetworkDetector& detector) + { + detector.getConnectivityHint = GetUnknownHint; + detector.notifyConnectivityHint = RejectSubscription; + } + + private: + static HRESULT WINAPI UnsupportedCost(INetworkListManager*, INetworkCostManager** cost) + { + *cost = nullptr; + return E_NOINTERFACE; + } + + static HRESULT WINAPI RejectCostQuery(INetworkListManager*, INetworkCostManager** cost) + { + *cost = nullptr; + return E_ACCESSDENIED; + } + + static HRESULT WINAPI RejectConnectionEvents( + IConnectionPointContainer* container, REFIID iid, IConnectionPoint** point) + { + if (iid == __uuidof(INetworkConnectionEvents)) + { + *point = nullptr; + return E_ACCESSDENIED; + } + return container->FindConnectionPoint(iid, point); + } + + static HRESULT WINAPI RejectDisconnect(IUnknown*, DWORD) + { + return E_FAIL; + } + + static DWORD WINAPI GetUnknownHint(NL_NETWORK_CONNECTIVITY_HINT* hint) + { + *hint = {}; + return NO_ERROR; + } + + static DWORD WINAPI RejectSubscription(NetworkDetector::HintChangedCallback, void*, BOOLEAN, HANDLE*) + { + return ERROR_ACCESS_DENIED; + } + }; + } +} +MAT_NS_END +#endif diff --git a/tests/dll-unload/CMakeLists.txt b/tests/dll-unload/CMakeLists.txt new file mode 100644 index 000000000..1e3ce458c --- /dev/null +++ b/tests/dll-unload/CMakeLists.txt @@ -0,0 +1,34 @@ +add_library(debug-listener-unload-module MODULE debug-listener-unload-module.cpp) +target_link_libraries(debug-listener-unload-module PRIVATE mat matsdk_internal_config) +target_include_directories(debug-listener-unload-module PRIVATE + "${PROJECT_SOURCE_DIR}/lib" + "${PROJECT_SOURCE_DIR}/lib/include" + "${PROJECT_SOURCE_DIR}/lib/include/public" + "${PROJECT_SOURCE_DIR}/lib/include/mat" + "${PROJECT_SOURCE_DIR}/tests") + +# The host must not link the SDK: its only reference is the explicitly loaded DLL. +add_executable(debug-listener-unload-test debug-listener-unload-test.cpp) +target_compile_definitions(debug-listener-unload-test PRIVATE WIN32_LEAN_AND_MEAN NOMINMAX) +add_dependencies(debug-listener-unload-test debug-listener-unload-module) +foreach(mode IN ITEMS idle dispatch network-native network-disabled network-legacy network-legacy-no-cost network-legacy-failures) + add_test(NAME debug-listener-unload-${mode} + COMMAND debug-listener-unload-test + $ ${mode} + CONFIGURATIONS Debug) + set_tests_properties(debug-listener-unload-${mode} PROPERTIES TIMEOUT 30) +endforeach() +foreach(mode IN ITEMS network-native network-disabled network-legacy network-legacy-no-cost network-legacy-failures) + set_tests_properties(debug-listener-unload-${mode} PROPERTIES SKIP_RETURN_CODE 77) +endforeach() + +add_executable(network-detector-reload-test network-detector-reload-test.cpp) +target_compile_definitions(network-detector-reload-test PRIVATE WIN32_LEAN_AND_MEAN NOMINMAX) +target_link_libraries(network-detector-reload-test PRIVATE ole32) +add_dependencies(network-detector-reload-test debug-listener-unload-module) +foreach(mode IN ITEMS native disabled legacy legacy-no-cost legacy-failures) + add_test(NAME network-detector-reload-${mode} + COMMAND network-detector-reload-test $ ${mode} + CONFIGURATIONS Debug) + set_tests_properties(network-detector-reload-${mode} PROPERTIES TIMEOUT 30 SKIP_RETURN_CODE 77) +endforeach() diff --git a/tests/dll-unload/README.md b/tests/dll-unload/README.md new file mode 100644 index 000000000..9b4036a68 --- /dev/null +++ b/tests/dll-unload/README.md @@ -0,0 +1,47 @@ +# Debug SDK DLL-unload regressions + +On MSVC static-SDK unit-test builds, `debug-listener-unload-test` loads a DLL +that embeds the SDK, creates seven native threads while it is loaded, and unloads +it while all seven threads remain alive. The host does not link the SDK. +Both targets use the Debug CRT and the default Debug STL iterator checking. + +The `idle` case never calls the SDK, modeling disabled telemetry. The `dispatch` +case also queries pending state and dispatches/removes a listener on each thread. +These cases do not start SDK background services. The `network-native` and +`network-legacy` cases start/stop a detector on each thread before unload; the +legacy case forces the compatibility backend. `network-legacy-no-cost` forces +the cost interface to be unavailable, exercising Server/older-client behavior +while all three original NLM event subscriptions remain active. +`network-legacy-failures` exercises rejected cost queries, partially registered +subscriptions, and failed explicit disconnection followed by STA rundown. +The `network-disabled` case constructs network information with +`CFG_BOOL_ENABLE_NET_DETECT = false` and checks that NLM is not loaded. +All seven require the DLL to be +unloaded and zero outstanding normal/client CRT blocks and bytes before allowing +the seven threads to exit. + +`network-detector-reload-test` repeats load/start/stop/unload five times with a +COM-uninitialized host. Default, disabled, legacy, legacy-no-cost, and legacy-failures modes must unload the +DLL and leave zero outstanding normal/client CRT blocks and bytes on every +iteration. The host's COM apartment must remain uninitialized. These tests do +not require a host-owned MTA or keep the SDK DLL permanently loaded. + +Build `debug-listener-unload-test` and `network-detector-reload-test` in Debug, then run: + +```text +ctest --test-dir -C Debug -R "debug-listener-unload|network-detector-reload" --output-on-failure +``` + +Nested dispatch, duplicate registrations, removal, exception cleanup, release +callback reentrancy, and thread isolation are covered by `DebugEventSourceTests` +in `UnitTests`. + +The network cases report CTest skip code 77 only when the embedded SDK's custom +SKU disables `HAVE_MAT_NETDETECT`; missing exports and failed startup are errors. +The default backend uses dynamically resolved IP Helper APIs on Windows 10 +version 2004/build 19041 and later. Older supported Windows uses an SDK-owned +STA with balanced Network List Manager subscriptions and COM teardown. +Unsupported optional cost interfaces return `Unknown` without disabling the +original network-list/network/connection event families. Cost-specific events +are not required. Forced +legacy testing on a modern OS does not replace execution on an older OS. diff --git a/tests/dll-unload/debug-listener-unload-module.cpp b/tests/dll-unload/debug-listener-unload-module.cpp new file mode 100644 index 000000000..a98b717fd --- /dev/null +++ b/tests/dll-unload/debug-listener-unload-module.cpp @@ -0,0 +1,168 @@ +// +// Copyright (c) Microsoft Corporation. All rights reserved. +// SPDX-License-Identifier: Apache-2.0 +// +#include "callbacks/DebugSourceInternal.hpp" +#include "pal/desktop/NetworkDetector.hpp" +#include "common/network-detector-test-access.hpp" +#include "config/RuntimeConfig_Default.hpp" +#include "pal/NetworkInformationImpl.hpp" + +#ifdef _DEBUG +static_assert(_ITERATOR_DEBUG_LEVEL == 2, "The regression requires Debug STL proxies."); +#ifndef _DLL +#error The unload regression requires the shared Debug CRT. +#endif +#endif + +namespace +{ + class Listener : public MAT::DebugEventListener + { + public: + void OnDebugEvent(MAT::DebugEvent&) override + { + ++calls; + } + + unsigned calls = 0; + }; +} + +extern "C" __declspec(dllexport) bool ExerciseDebugListeners() +{ + Listener listener; + if (MAT::IsDebugEventListenerPending(&listener)) + { + return false; + } + MAT::DebugEventSource source; + source.AddEventListener(MAT::EVT_LOG_EVENT, listener); + const bool dispatched = source.DispatchEvent(MAT::DebugEvent { MAT::EVT_LOG_EVENT }); + source.RemoveEventListener(MAT::EVT_LOG_EVENT, listener); + return dispatched && listener.calls == 1 && + !MAT::IsDebugEventListenerPending(&listener); +} + +extern "C" __declspec(dllexport) bool HasNetworkDetector() +{ +#ifdef HAVE_MAT_NETDETECT + return true; +#else + return false; +#endif +} + +#ifdef HAVE_MAT_NETDETECT +enum class NetworkBackend +{ + Native, + Legacy, + LegacyWithoutCost +}; + +static bool IsStopped(MATW::NetworkDetector& detector) +{ + return !detector.isUp() && !detector.QueueNetworkCostRefresh() && + !MATW::NetworkDetectorTestAccess::HasLegacyManager(detector) && + !MATW::NetworkDetectorTestAccess::HasLegacyCost(detector) && + MATW::NetworkDetectorTestAccess::LegacySubscriptionCount(detector) == 0; +} + +static bool ExerciseNetworkDetectorBackend(NetworkBackend backend) +{ + MATW::NetworkDetector detector; + if (backend == NetworkBackend::Legacy) + { + MATW::NetworkDetectorTestAccess::UseLegacyBackend(detector); + } + else if (backend == NetworkBackend::LegacyWithoutCost) + { + MATW::NetworkDetectorTestAccess::UseLegacyBackendWithoutCost(detector); + } + if (!detector.Start()) + { + return false; + } + const bool running = detector.isUp(); + const auto cost = detector.GetCurrentNetworkCost(); + const bool readable = cost == MAT::NetworkCost_Unknown || cost == MAT::NetworkCost_Unmetered || + cost == MAT::NetworkCost_Metered || cost == MAT::NetworkCost_Roaming; + const bool subscriptions = backend == NetworkBackend::Native || + MATW::NetworkDetectorTestAccess::LegacySubscriptionCount(detector) == 3; + const bool optionalCost = backend != NetworkBackend::LegacyWithoutCost || + (cost == MAT::NetworkCost_Unknown && + !MATW::NetworkDetectorTestAccess::HasLegacyCost(detector)); + detector.Stop(); + return running && readable && subscriptions && optionalCost && IsStopped(detector); +} + +extern "C" __declspec(dllexport) bool ExerciseNetworkDetector() +{ + return ExerciseNetworkDetectorBackend(NetworkBackend::Native); +} + +extern "C" __declspec(dllexport) bool ExerciseDisabledNetworkDetection() +{ + const auto moduleBefore = GetModuleHandleW(L"netprofm.dll"); + MAT::ILogConfiguration configuration; + configuration[MAT::CFG_BOOL_ENABLE_NET_DETECT] = false; + MAT::RuntimeConfig_Default runtimeConfig(configuration); + for (unsigned iteration = 0; iteration < 10; ++iteration) + { + auto network = PAL::NetworkInformationImpl::Create(runtimeConfig); + if (network->GetNetworkCost() != MAT::NetworkCost_Unmetered || + GetModuleHandleW(L"netprofm.dll") != moduleBefore) + { + return false; + } + } + return true; +} + +extern "C" __declspec(dllexport) bool ExerciseLegacyNetworkDetector() +{ + return ExerciseNetworkDetectorBackend(NetworkBackend::Legacy); +} + +extern "C" __declspec(dllexport) bool ExerciseLegacyNetworkDetectorWithoutCost() +{ + return ExerciseNetworkDetectorBackend(NetworkBackend::LegacyWithoutCost); +} + +extern "C" __declspec(dllexport) bool ExerciseLegacyNetworkDetectorFailures() +{ + MATW::NetworkDetector detector; + for (unsigned iteration = 0; iteration < 3; ++iteration) + { + MATW::NetworkDetectorTestAccess::FailLegacyCostQuery(detector); + if (!detector.Start() || detector.GetCurrentNetworkCost() != MAT::NetworkCost_Unknown || + MATW::NetworkDetectorTestAccess::HasLegacyCost(detector)) + { + return false; + } + detector.Stop(); + if (!IsStopped(detector)) + { + return false; + } + MATW::NetworkDetectorTestAccess::FailLegacySubscription(detector); + if (detector.Start() || !IsStopped(detector)) + { + return false; + } + MATW::NetworkDetectorTestAccess::FailLegacyDisconnect(detector); + if (!detector.Start()) + { + return false; + } + detector.GetCurrentNetworkCost(); + detector.Stop(); + if (!IsStopped(detector)) + { + return false; + } + } + return true; +} +#endif diff --git a/tests/dll-unload/debug-listener-unload-test.cpp b/tests/dll-unload/debug-listener-unload-test.cpp new file mode 100644 index 000000000..c466b03bc --- /dev/null +++ b/tests/dll-unload/debug-listener-unload-test.cpp @@ -0,0 +1,188 @@ +// +// Copyright (c) Microsoft Corporation. All rights reserved. +// SPDX-License-Identifier: Apache-2.0 +// +#include +#include +#include +#include + +#if defined(_DEBUG) && !defined(_DLL) +#error The unload regression requires the shared Debug CRT. +#endif + +namespace +{ + using Exercise = bool (*)(); + constexpr unsigned threadCount = 7; + + struct Worker + { + HANDLE ready = nullptr; + HANDLE exit = nullptr; + Exercise exercise = nullptr; + bool succeeded = false; + }; + + DWORD WINAPI RunWorker(void* context) + { + auto& worker = *static_cast(context); + worker.succeeded = worker.exercise == nullptr || worker.exercise(); + SetEvent(worker.ready); + return WaitForSingleObject(worker.exit, INFINITE) == WAIT_OBJECT_0 ? 0 : 1; + } +} + +int main(int argc, char** argv) +{ + if (argc != 3 || (std::strcmp(argv[2], "idle") != 0 && + std::strcmp(argv[2], "dispatch") != 0 && + std::strcmp(argv[2], "network-native") != 0 && + std::strcmp(argv[2], "network-disabled") != 0 && + std::strcmp(argv[2], "network-legacy") != 0 && + std::strcmp(argv[2], "network-legacy-no-cost") != 0 && + std::strcmp(argv[2], "network-legacy-failures") != 0)) + { + std::fprintf(stderr, "Usage: debug-listener-unload-test " + "\n"); + return 1; + } +#ifndef _DEBUG + std::fprintf(stderr, "This regression requires the Debug CRT and Debug STL.\n"); + return 1; +#else + std::printf("Debug SDK unload: %s, %u live threads\n", argv[2], threadCount); + _CrtMemState before {}; + _CrtMemState after {}; + _CrtMemState difference {}; + _CrtMemCheckpoint(&before); + + HMODULE module = LoadLibraryA(argv[1]); + if (module == nullptr) + { + std::fprintf(stderr, "LoadLibrary failed: %lu\n", GetLastError()); + return 1; + } + const char* entry = "ExerciseDebugListeners"; + if (std::strcmp(argv[2], "network-native") == 0) + { + entry = "ExerciseNetworkDetector"; + } + else if (std::strcmp(argv[2], "network-legacy") == 0) + { + entry = "ExerciseLegacyNetworkDetector"; + } + else if (std::strcmp(argv[2], "network-disabled") == 0) + { + entry = "ExerciseDisabledNetworkDetection"; + } + else if (std::strcmp(argv[2], "network-legacy-no-cost") == 0) + { + entry = "ExerciseLegacyNetworkDetectorWithoutCost"; + } + else if (std::strcmp(argv[2], "network-legacy-failures") == 0) + { + entry = "ExerciseLegacyNetworkDetectorFailures"; + } + if (std::strncmp(argv[2], "network-", 8) == 0) + { + auto hasDetector = reinterpret_cast(GetProcAddress(module, "HasNetworkDetector")); + if (hasDetector != nullptr && !hasDetector()) + { + std::printf("Network detection is disabled in this SDK SKU.\n"); + return FreeLibrary(module) ? 77 : 1; + } + } + auto exercise = reinterpret_cast(GetProcAddress(module, entry)); + if (exercise == nullptr) + { + std::fprintf(stderr, "GetProcAddress failed: %lu\n", GetLastError()); + FreeLibrary(module); + return 1; + } + + Worker workers[threadCount]; + HANDLE threads[threadCount] {}; + HANDLE exit = CreateEventW(nullptr, TRUE, FALSE, nullptr); + bool succeeded = exit != nullptr; + for (unsigned i = 0; succeeded && i < threadCount; ++i) + { + workers[i].exit = exit; + workers[i].ready = CreateEventW(nullptr, TRUE, FALSE, nullptr); + workers[i].exercise = std::strcmp(argv[2], "idle") != 0 ? exercise : nullptr; + if (workers[i].ready == nullptr) + { + succeeded = false; + break; + } + threads[i] = CreateThread(nullptr, 0, RunWorker, &workers[i], 0, nullptr); + succeeded = threads[i] != nullptr && + WaitForSingleObject(workers[i].ready, 10000) == WAIT_OBJECT_0 && + workers[i].succeeded; + } + + if (succeeded) + { + succeeded = FreeLibrary(module) != FALSE; + if (succeeded) + { + module = nullptr; + } + else + { + std::fprintf(stderr, "FreeLibrary failed: %lu\n", GetLastError()); + } + if (GetModuleHandleA(argv[1]) != nullptr) + { + std::fprintf(stderr, "The SDK DLL is still loaded.\n"); + succeeded = false; + } + for (auto thread : threads) + { + DWORD code = 0; + succeeded = GetExitCodeThread(thread, &code) != FALSE && + code == STILL_ACTIVE && succeeded; + } + _CrtMemCheckpoint(&after); + _CrtMemDifference(&difference, &before, &after); + const size_t blocks = difference.lCounts[_NORMAL_BLOCK] + difference.lCounts[_CLIENT_BLOCK]; + const size_t bytes = difference.lSizes[_NORMAL_BLOCK] + difference.lSizes[_CLIENT_BLOCK]; + std::printf("After unload, seven threads still alive: %zu blocks, %zu bytes\n", blocks, bytes); + succeeded = succeeded && blocks == 0 && bytes == 0; + if (blocks != 0 || bytes != 0) + { + _CrtMemDumpAllObjectsSince(&before); + } + } + else + { + std::fprintf(stderr, "Worker setup/exercise failed: %lu\n", GetLastError()); + } + + if (exit != nullptr) + { + SetEvent(exit); + } + for (unsigned i = 0; i < threadCount; ++i) + { + if (threads[i] != nullptr) + { + succeeded = WaitForSingleObject(threads[i], INFINITE) == WAIT_OBJECT_0 && succeeded; + CloseHandle(threads[i]); + } + if (workers[i].ready != nullptr) + { + CloseHandle(workers[i].ready); + } + } + if (exit != nullptr) + { + CloseHandle(exit); + } + if (module != nullptr) + { + FreeLibrary(module); + } + return succeeded ? 0 : 1; +#endif +} diff --git a/tests/dll-unload/network-detector-reload-test.cpp b/tests/dll-unload/network-detector-reload-test.cpp new file mode 100644 index 000000000..55e033e55 --- /dev/null +++ b/tests/dll-unload/network-detector-reload-test.cpp @@ -0,0 +1,108 @@ +// +// Copyright (c) Microsoft Corporation. All rights reserved. +// SPDX-License-Identifier: Apache-2.0 +// +#include +#include +#include +#include +#include + +int main(int argc, char** argv) +{ + if (argc != 3 || (std::strcmp(argv[2], "native") != 0 && + std::strcmp(argv[2], "disabled") != 0 && + std::strcmp(argv[2], "legacy") != 0 && + std::strcmp(argv[2], "legacy-no-cost") != 0 && + std::strcmp(argv[2], "legacy-failures") != 0)) + { + std::fprintf(stderr, "Usage: network-detector-reload-test " + "\n"); + return 1; + } + APTTYPE apartment; + APTTYPEQUALIFIER qualifier; + if (CoGetApartmentType(&apartment, &qualifier) != CO_E_NOTINITIALIZED) + { + std::fprintf(stderr, "The host must not initialize a COM apartment.\n"); + return 1; + } + for (unsigned iteration = 1; iteration <= 5; ++iteration) + { +#ifdef _DEBUG + _CrtMemState before {}; + _CrtMemCheckpoint(&before); +#endif + std::printf("Loading SDK DLL, iteration %u\n", iteration); + std::fflush(stdout); + HMODULE module = LoadLibraryA(argv[1]); + if (module == nullptr) + { + std::fprintf(stderr, "LoadLibrary failed: %lu\n", GetLastError()); + return 1; + } + auto hasDetector = reinterpret_cast(GetProcAddress(module, "HasNetworkDetector")); + if (hasDetector != nullptr && !hasDetector()) + { + std::printf("Network detection is disabled in this SDK SKU.\n"); + return FreeLibrary(module) ? 77 : 1; + } + const char* entry = "ExerciseNetworkDetector"; + if (std::strcmp(argv[2], "disabled") == 0) + { + entry = "ExerciseDisabledNetworkDetection"; + } + else if (std::strcmp(argv[2], "legacy") == 0) + { + entry = "ExerciseLegacyNetworkDetector"; + } + else if (std::strcmp(argv[2], "legacy-no-cost") == 0) + { + entry = "ExerciseLegacyNetworkDetectorWithoutCost"; + } + else if (std::strcmp(argv[2], "legacy-failures") == 0) + { + entry = "ExerciseLegacyNetworkDetectorFailures"; + } + auto exercise = reinterpret_cast(GetProcAddress(module, entry)); + if (exercise == nullptr || !exercise()) + { + std::fprintf(stderr, "Network detector exercise failed on iteration %u.\n", iteration); + FreeLibrary(module); + return 1; + } + if (!FreeLibrary(module)) + { + std::fprintf(stderr, "FreeLibrary failed: %lu\n", GetLastError()); + return 1; + } + if (GetModuleHandleA(argv[1]) != nullptr) + { + std::fprintf(stderr, "The SDK DLL is still loaded.\n"); + return 1; + } + if (CoGetApartmentType(&apartment, &qualifier) != CO_E_NOTINITIALIZED) + { + std::fprintf(stderr, "The SDK changed the host COM apartment.\n"); + return 1; + } +#ifdef _DEBUG + _CrtMemState after {}; + _CrtMemState difference {}; + _CrtMemCheckpoint(&after); + _CrtMemDifference(&difference, &before, &after); + const size_t blocks = difference.lCounts[_NORMAL_BLOCK] + difference.lCounts[_CLIENT_BLOCK]; + const size_t bytes = difference.lSizes[_NORMAL_BLOCK] + difference.lSizes[_CLIENT_BLOCK]; + std::printf("After unload: %zu blocks, %zu bytes\n", blocks, bytes); + if (blocks != 0 || bytes != 0) + { + std::fprintf(stderr, "Reload iteration %u leaked %zu blocks / %zu bytes.\n", iteration, blocks, bytes); + _CrtMemDumpAllObjectsSince(&before); + return 1; + } +#endif + std::printf("SDK DLL unloaded, iteration %u\n", iteration); + std::fflush(stdout); + } + return 0; +} diff --git a/tests/unittests/DebugEventSourceTests.cpp b/tests/unittests/DebugEventSourceTests.cpp index cf725d6db..fbf4bf605 100644 --- a/tests/unittests/DebugEventSourceTests.cpp +++ b/tests/unittests/DebugEventSourceTests.cpp @@ -4,8 +4,11 @@ // #include "common/Common.hpp" +#include "callbacks/DebugSourceInternal.hpp" #include #include +#include +#include using namespace testing; using namespace MAT; @@ -238,3 +241,192 @@ TEST(DebugEventSourceTests, DispatchEvent_OneEventToCascadedAndToSource_Listener ASSERT_EQ(sequenceNumberToCountMap[1], uint64_t { 2 }); } +TEST(DebugEventSourceTests, PendingListeners_NoDispatch_ReturnsFalse) +{ + TestDebugEventListener listener; + EXPECT_FALSE(IsDebugEventListenerPending(&listener)); + EXPECT_FALSE(IsDebugEventListenerPending(nullptr)); +} + +TEST(DebugEventSourceTests, PendingListeners_NestedDispatchRemoval_RestoresOuterSnapshot) +{ + TestDebugEventSource source; + TestDebugEventListener first; + TestDebugEventListener second; + bool nested = false; + unsigned secondCalls = 0; + first.OnDebugEventOverride = [&](DebugEvent&) { + EXPECT_TRUE(IsDebugEventListenerPending(&second)); + EXPECT_FALSE(IsDebugEventListenerPending(&first)); + if (!nested) + { + nested = true; + source.RemoveEventListener(EVT_LOG_EVENT, second); + source.DispatchEvent(DebugEvent { EVT_LOG_EVENT }); + EXPECT_TRUE(IsDebugEventListenerPending(&second)); + } + }; + second.OnDebugEventOverride = [&](DebugEvent&) { + EXPECT_FALSE(IsDebugEventListenerPending(&second)); + ++secondCalls; + }; + source.AddEventListener(EVT_LOG_EVENT, first); + source.AddEventListener(EVT_LOG_EVENT, second); + + source.DispatchEvent(DebugEvent { EVT_LOG_EVENT }); + + EXPECT_EQ(secondCalls, 1u); + EXPECT_FALSE(IsDebugEventListenerPending(&second)); +} + +TEST(DebugEventSourceTests, PendingListeners_NestedDuplicateListeners_PreservesOuterOccurrences) +{ + TestDebugEventSource outer; + TestDebugEventSource inner; + TestDebugEventListener first; + TestDebugEventListener shared; + unsigned sharedCalls = 0; + first.OnDebugEventOverride = [&](DebugEvent&) { + inner.DispatchEvent(DebugEvent { EVT_LOG_EVENT }); + EXPECT_TRUE(IsDebugEventListenerPending(&shared)); + }; + shared.OnDebugEventOverride = [&](DebugEvent&) { + ++sharedCalls; + EXPECT_EQ(IsDebugEventListenerPending(&shared), sharedCalls < 4); + }; + outer.AddEventListener(EVT_LOG_EVENT, first); + outer.AddEventListener(EVT_LOG_EVENT, shared); + outer.AddEventListener(EVT_LOG_EVENT, shared); + inner.AddEventListener(EVT_LOG_EVENT, shared); + inner.AddEventListener(EVT_LOG_EVENT, shared); + + outer.DispatchEvent(DebugEvent { EVT_LOG_EVENT }); + + EXPECT_EQ(sharedCalls, 4u); + EXPECT_FALSE(IsDebugEventListenerPending(&shared)); +} + +TEST(DebugEventSourceTests, PendingListeners_AnotherThread_DoesNotSeeActiveScope) +{ + TestDebugEventSource source; + TestDebugEventListener first; + TestDebugEventListener second; + first.OnDebugEventOverride = [&](DebugEvent&) { + EXPECT_TRUE(IsDebugEventListenerPending(&second)); + std::thread other([&] { + EXPECT_FALSE(IsDebugEventListenerPending(&second)); + }); + other.join(); + EXPECT_TRUE(IsDebugEventListenerPending(&second)); + }; + source.AddEventListener(EVT_LOG_EVENT, first); + source.AddEventListener(EVT_LOG_EVENT, second); + + source.DispatchEvent(DebugEvent { EVT_LOG_EVENT }); + + EXPECT_FALSE(IsDebugEventListenerPending(&second)); +} + +namespace +{ + std::function pendingReleaseOverride; + + class PendingReleaseScope + { + public: + explicit PendingReleaseScope(std::function callback) + { + pendingReleaseOverride = std::move(callback); + SetDebugEventListenerPendingReleaseCallback([](DebugEventListener* listener) { + pendingReleaseOverride(listener); + }); + } + + ~PendingReleaseScope() + { + SetDebugEventListenerPendingReleaseCallback(nullptr); + pendingReleaseOverride = nullptr; + } + }; +} + +TEST(DebugEventSourceTests, PendingListeners_ExceptionRelease_ReentrantDispatchRestoresPendingState) +{ + TestDebugEventSource outer; + TestDebugEventSource inner; + TestDebugEventSource reentrant; + TestDebugEventListener first; + TestDebugEventListener throwing; + TestDebugEventListener shared; + TestDebugEventListener last; + unsigned releases = 0; + bool reentered = false; + PendingReleaseScope release([&](DebugEventListener* listener) { + ++releases; + EXPECT_EQ(listener, &shared); + EXPECT_TRUE(IsDebugEventListenerPending(&shared)); + EXPECT_TRUE(IsDebugEventListenerPending(&last)); + reentrant.DispatchEvent(DebugEvent { EVT_LOG_EVENT }); + EXPECT_TRUE(IsDebugEventListenerPending(&shared)); + EXPECT_TRUE(IsDebugEventListenerPending(&last)); + }); + first.OnDebugEventOverride = [&](DebugEvent&) { + EXPECT_THROW(inner.DispatchEvent(DebugEvent { EVT_LOG_EVENT }), std::runtime_error); + EXPECT_TRUE(IsDebugEventListenerPending(&shared)); + EXPECT_FALSE(IsDebugEventListenerPending(&throwing)); + }; + throwing.OnDebugEventOverride = [](DebugEvent&) { + throw std::runtime_error("listener failure"); + }; + shared.OnDebugEventOverride = [&](DebugEvent&) { + EXPECT_FALSE(IsDebugEventListenerPending(&shared)); + }; + last.OnDebugEventOverride = [&](DebugEvent&) { + EXPECT_FALSE(IsDebugEventListenerPending(&last)); + }; + TestDebugEventListener reentrantListener; + reentrantListener.OnDebugEventOverride = [&](DebugEvent&) { + reentered = true; + EXPECT_TRUE(IsDebugEventListenerPending(&shared)); + EXPECT_TRUE(IsDebugEventListenerPending(&last)); + }; + outer.AddEventListener(EVT_LOG_EVENT, first); + outer.AddEventListener(EVT_LOG_EVENT, shared); + outer.AddEventListener(EVT_LOG_EVENT, last); + inner.AddEventListener(EVT_LOG_EVENT, throwing); + inner.AddEventListener(EVT_LOG_EVENT, shared); + inner.AddEventListener(EVT_LOG_EVENT, shared); + reentrant.AddEventListener(EVT_LOG_EVENT, reentrantListener); + + outer.DispatchEvent(DebugEvent { EVT_LOG_EVENT }); + + EXPECT_EQ(releases, 2u); + EXPECT_TRUE(reentered); + EXPECT_FALSE(IsDebugEventListenerPending(&shared)); + EXPECT_FALSE(IsDebugEventListenerPending(&last)); +} + +TEST(DebugEventSourceTests, PendingListeners_ExceptionRelease_RemovesDuplicatesBeforeRelease) +{ + TestDebugEventSource source; + TestDebugEventListener throwing; + TestDebugEventListener pending; + unsigned releases = 0; + PendingReleaseScope release([&](DebugEventListener* listener) { + EXPECT_EQ(listener, &pending); + ++releases; + EXPECT_EQ(IsDebugEventListenerPending(&pending), releases == 1); + }); + throwing.OnDebugEventOverride = [](DebugEvent&) { + throw std::runtime_error("listener failure"); + }; + source.AddEventListener(EVT_LOG_EVENT, throwing); + source.AddEventListener(EVT_LOG_EVENT, pending); + source.AddEventListener(EVT_LOG_EVENT, pending); + + EXPECT_THROW(source.DispatchEvent(DebugEvent { EVT_LOG_EVENT }), std::runtime_error); + + EXPECT_EQ(releases, 2u); + EXPECT_FALSE(IsDebugEventListenerPending(&throwing)); + EXPECT_FALSE(IsDebugEventListenerPending(&pending)); +} diff --git a/tests/unittests/NetworkDetectorTests.cpp b/tests/unittests/NetworkDetectorTests.cpp index 0103a4c11..71ebb18ad 100644 --- a/tests/unittests/NetworkDetectorTests.cpp +++ b/tests/unittests/NetworkDetectorTests.cpp @@ -5,6 +5,9 @@ #if defined(_WIN32) && defined(HAVE_MAT_NETDETECT) #include "api/LogManagerFactory.hpp" #include "pal/desktop/NetworkDetector.hpp" +#include "common/network-detector-test-access.hpp" +#include "config/RuntimeConfig_Default.hpp" +#include "pal/NetworkInformationImpl.hpp" #include @@ -39,30 +42,47 @@ class StopDetectorOnNetworkChange : public DebugEventListener std::promise stopped; }; -TEST(NetworkDetectorTests, MapsWinRTNetworkCosts) +TEST(NetworkDetectorTests, MapsConnectivityHintNetworkCosts) { - EXPECT_EQ(MATW::MapNetworkCost(NetworkCostType_Unrestricted, false, false, false, false), NetworkCost_Unmetered); - EXPECT_EQ(MATW::MapNetworkCost(NetworkCostType_Fixed, false, false, false, false), NetworkCost_Metered); - EXPECT_EQ(MATW::MapNetworkCost(NetworkCostType_Variable, false, false, false, false), NetworkCost_Metered); - EXPECT_EQ(MATW::MapNetworkCost(NetworkCostType_Unknown, false, false, false, false), NetworkCost_Unknown); + EXPECT_EQ(MATW::MapNetworkCost(NetworkConnectivityCostHintUnrestricted, false, false, false), NetworkCost_Unmetered); + EXPECT_EQ(MATW::MapNetworkCost(NetworkConnectivityCostHintFixed, false, false, false), NetworkCost_Metered); + EXPECT_EQ(MATW::MapNetworkCost(NetworkConnectivityCostHintVariable, false, false, false), NetworkCost_Metered); + EXPECT_EQ(MATW::MapNetworkCost(NetworkConnectivityCostHintUnknown, false, false, false), NetworkCost_Unknown); } -TEST(NetworkDetectorTests, MapsRestrictiveWinRTNetworkStates) +TEST(NetworkDetectorTests, MapsRestrictiveConnectivityHints) { - EXPECT_EQ(MATW::MapNetworkCost(NetworkCostType_Unrestricted, true, false, false, false), NetworkCost_Roaming); - EXPECT_EQ(MATW::MapNetworkCost(NetworkCostType_Unrestricted, false, true, false, false), NetworkCost_Roaming); - EXPECT_EQ(MATW::MapNetworkCost(NetworkCostType_Unrestricted, false, false, true, false), NetworkCost_Roaming); - EXPECT_EQ(MATW::MapNetworkCost(NetworkCostType_Unrestricted, false, false, false, true), NetworkCost_Roaming); + EXPECT_EQ(MATW::MapNetworkCost(NetworkConnectivityCostHintUnrestricted, true, false, false), NetworkCost_Roaming); + EXPECT_EQ(MATW::MapNetworkCost(NetworkConnectivityCostHintUnrestricted, false, true, false), NetworkCost_Roaming); + EXPECT_EQ(MATW::MapNetworkCost(NetworkConnectivityCostHintUnrestricted, false, false, true), NetworkCost_Roaming); +} + +TEST(NetworkDetectorTests, MapsLegacyCostFlagsIncludingCombinedRestrictions) +{ + EXPECT_EQ(MATW::MapLegacyNetworkCost(NLM_CONNECTION_COST_UNKNOWN), NetworkCost_Unknown); + EXPECT_EQ(MATW::MapLegacyNetworkCost(NLM_CONNECTION_COST_UNRESTRICTED), NetworkCost_Unmetered); + EXPECT_EQ(MATW::MapLegacyNetworkCost(NLM_CONNECTION_COST_FIXED), NetworkCost_Metered); + EXPECT_EQ(MATW::MapLegacyNetworkCost(NLM_CONNECTION_COST_VARIABLE), NetworkCost_Metered); + for (const DWORD flag : { NLM_CONNECTION_COST_ROAMING, NLM_CONNECTION_COST_OVERDATALIMIT, + NLM_CONNECTION_COST_APPROACHINGDATALIMIT, NLM_CONNECTION_COST_CONGESTED }) + { + EXPECT_EQ(MATW::MapLegacyNetworkCost(NLM_CONNECTION_COST_UNRESTRICTED | flag), NetworkCost_Roaming); + EXPECT_EQ(MATW::MapLegacyNetworkCost(NLM_CONNECTION_COST_FIXED | flag), NetworkCost_Roaming); + } } TEST(NetworkDetectorTests, StartsReadsCostAndStopsWithoutNetworkListManager) { - ASSERT_EQ(GetModuleHandleW(L"netprofm.dll"), nullptr); + const auto moduleBefore = GetModuleHandleW(L"netprofm.dll"); MATW::NetworkDetector detector; + if (!MATW::NetworkDetectorTestAccess::HasNativeBackend(detector)) + { + GTEST_SKIP() << "IP Helper connectivity hints are unavailable on this Windows release."; + } ASSERT_TRUE(detector.Start()); EXPECT_TRUE(detector.isUp()); - EXPECT_EQ(GetModuleHandleW(L"netprofm.dll"), nullptr); + EXPECT_EQ(GetModuleHandleW(L"netprofm.dll"), moduleBefore); const auto cost = detector.GetCurrentNetworkCost(); EXPECT_THAT(cost, AnyOf( @@ -75,7 +95,7 @@ TEST(NetworkDetectorTests, StartsReadsCostAndStopsWithoutNetworkListManager) detector.Stop(); EXPECT_FALSE(detector.isUp()); EXPECT_FALSE(detector.QueueNetworkCostRefresh()); - EXPECT_EQ(GetModuleHandleW(L"netprofm.dll"), nullptr); + EXPECT_EQ(GetModuleHandleW(L"netprofm.dll"), moduleBefore); } TEST(NetworkDetectorTests, QueuedNetworkCallbackRaceDoesNotOutliveStop) @@ -96,6 +116,38 @@ TEST(NetworkDetectorTests, QueuedNetworkCallbackRaceDoesNotOutliveStop) callbackThread.join(); } +TEST(NetworkDetectorTests, ConfigurationDisablesDetectionWithoutLoadingNetworkListManager) +{ + const auto moduleBefore = GetModuleHandleW(L"netprofm.dll"); + ILogConfiguration configuration; + configuration[CFG_BOOL_ENABLE_NET_DETECT] = false; + RuntimeConfig_Default runtimeConfig(configuration); + for (unsigned iteration = 0; iteration < 10; ++iteration) + { + auto network = PAL::NetworkInformationImpl::Create(runtimeConfig); + ASSERT_NE(network, nullptr); + EXPECT_EQ(network->GetNetworkCost(), NetworkCost_Unmetered); + EXPECT_EQ(GetModuleHandleW(L"netprofm.dll"), moduleBefore); + } +} + +TEST(NetworkDetectorTests, FailedSubscriptionCleansUpAndCanRetry) +{ + MATW::NetworkDetector detector; + MATW::NetworkDetectorTestAccess::FailNativeSubscription(detector); + for (unsigned iteration = 0; iteration < 3; ++iteration) + { + EXPECT_FALSE(detector.Start()); + EXPECT_FALSE(detector.isUp()); + EXPECT_FALSE(detector.QueueNetworkCostRefresh()); + detector.Stop(); + } + MATW::NetworkDetectorTestAccess::UseLegacyBackend(detector); + EXPECT_TRUE(detector.Start()); + detector.Stop(); + EXPECT_FALSE(detector.isUp()); +} + TEST(NetworkDetectorTests, ConcurrentStopWaitsForStartupPublication) { for (int iteration = 0; iteration < 20; ++iteration) @@ -132,8 +184,295 @@ TEST(NetworkDetectorTests, NetworkChangeListenerCanStopDetector) ASSERT_TRUE(detector.Start()); ASSERT_EQ(stopped.wait_for(std::chrono::seconds(5)), std::future_status::ready); EXPECT_FALSE(detector.isUp()); + detector.Stop(); logManager->RemoveEventListener(EVT_NET_CHANGED, listener); EXPECT_EQ(LogManagerFactory::Destroy(logManager), STATUS_SUCCESS); } + +enum class NetworkDetectorBackend +{ + Native, + Legacy, + LegacyWithoutCost +}; + +class NetworkDetectorBackendTests : public TestWithParam +{ +protected: + void SelectBackend(MATW::NetworkDetector& detector) + { + if (GetParam() == NetworkDetectorBackend::Legacy) + { + MATW::NetworkDetectorTestAccess::UseLegacyBackend(detector); + } + else if (GetParam() == NetworkDetectorBackend::LegacyWithoutCost) + { + MATW::NetworkDetectorTestAccess::UseLegacyBackendWithoutCost(detector); + } + } +}; + +class BlockingStopDetectorOnNetworkChange : public DebugEventListener +{ +public: + BlockingStopDetectorOnNetworkChange(MATW::NetworkDetector& detector, bool stopBeforeRelease) : + detector(detector), stopBeforeRelease(stopBeforeRelease), released(release.get_future()) + { + } + + void OnDebugEvent(DebugEvent& event) override + { + if (event.type != EVT_NET_CHANGED || handled.exchange(true)) + { + return; + } + if (stopBeforeRelease) + { + detector.Stop(); + } + entered.set_value(); + released.wait(); + if (!stopBeforeRelease) + { + detector.Stop(); + } + } + + std::future GetEnteredFuture() { return entered.get_future(); } + void Release() { release.set_value(); } + +private: + MATW::NetworkDetector& detector; + bool stopBeforeRelease; + std::atomic handled{false}; + std::promise entered; + std::promise release; + std::future released; +}; + +TEST_P(NetworkDetectorBackendTests, ConcurrentExternalAndReentrantStopsDrainCallback) +{ + ILogConfiguration configuration; + configuration[CFG_BOOL_ENABLE_NET_DETECT] = false; + ILogManager* manager = LogManagerFactory::Create(configuration); + ASSERT_NE(manager, nullptr); + MATW::NetworkDetector detector; + SelectBackend(detector); + BlockingStopDetectorOnNetworkChange listener(detector, false); + auto entered = listener.GetEnteredFuture(); + manager->AddEventListener(EVT_NET_CHANGED, listener); + EXPECT_TRUE(detector.Start()); + EXPECT_EQ(entered.wait_for(std::chrono::seconds(5)), std::future_status::ready); + auto firstStop = std::async(std::launch::async, [&] { detector.Stop(); }); + const auto deadline = std::chrono::steady_clock::now() + std::chrono::seconds(5); + while (detector.isUp() && std::chrono::steady_clock::now() < deadline) + { + std::this_thread::yield(); + } + EXPECT_FALSE(detector.isUp()); + auto secondStop = std::async(std::launch::async, [&] { detector.Stop(); }); + EXPECT_EQ(firstStop.wait_for(std::chrono::milliseconds(50)), std::future_status::timeout); + EXPECT_EQ(secondStop.wait_for(std::chrono::milliseconds(50)), std::future_status::timeout); + listener.Release(); + EXPECT_EQ(firstStop.wait_for(std::chrono::seconds(5)), std::future_status::ready); + EXPECT_EQ(secondStop.wait_for(std::chrono::seconds(5)), std::future_status::ready); + firstStop.get(); + secondStop.get(); + manager->RemoveEventListener(EVT_NET_CHANGED, listener); + EXPECT_EQ(LogManagerFactory::Destroy(manager), STATUS_SUCCESS); +} + +TEST_P(NetworkDetectorBackendTests, ExternalStopDrainsCallbackAfterReentrantStop) +{ + ILogConfiguration configuration; + configuration[CFG_BOOL_ENABLE_NET_DETECT] = false; + ILogManager* manager = LogManagerFactory::Create(configuration); + ASSERT_NE(manager, nullptr); + MATW::NetworkDetector detector; + SelectBackend(detector); + BlockingStopDetectorOnNetworkChange listener(detector, true); + auto entered = listener.GetEnteredFuture(); + manager->AddEventListener(EVT_NET_CHANGED, listener); + EXPECT_TRUE(detector.Start()); + EXPECT_EQ(entered.wait_for(std::chrono::seconds(5)), std::future_status::ready); + EXPECT_FALSE(detector.isUp()); + auto stopped = std::async(std::launch::async, [&] { detector.Stop(); }); + EXPECT_EQ(stopped.wait_for(std::chrono::milliseconds(50)), std::future_status::timeout); + listener.Release(); + EXPECT_EQ(stopped.wait_for(std::chrono::seconds(5)), std::future_status::ready); + stopped.get(); + manager->RemoveEventListener(EVT_NET_CHANGED, listener); + EXPECT_EQ(LogManagerFactory::Destroy(manager), STATUS_SUCCESS); +} + +TEST_P(NetworkDetectorBackendTests, RepeatedStartReadAndStop) +{ + MATW::NetworkDetector detector; + SelectBackend(detector); + for (unsigned iteration = 0; iteration < 10; ++iteration) + { + ASSERT_TRUE(detector.Start()); + if (GetParam() != NetworkDetectorBackend::Native) + { + EXPECT_EQ(MATW::NetworkDetectorTestAccess::LegacySubscriptionCount(detector), 3u); + } + if (GetParam() == NetworkDetectorBackend::LegacyWithoutCost) + { + EXPECT_FALSE(MATW::NetworkDetectorTestAccess::HasLegacyCost(detector)); + EXPECT_EQ(detector.GetNetworkCost(), NetworkCost_Unknown); + EXPECT_EQ(MATW::NetworkDetectorTestAccess::LegacySubscriptionCount(detector), 3u); + } + EXPECT_EQ(detector.GetCurrentNetworkCost(), detector.GetNetworkCost()); + detector.Stop(); + EXPECT_EQ(MATW::NetworkDetectorTestAccess::LegacySubscriptionCount(detector), 0u); + EXPECT_FALSE(detector.isUp()); + EXPECT_FALSE(detector.QueueNetworkCostRefresh()); + } +} + +TEST_P(NetworkDetectorBackendTests, RestartDrainsPreviousCallback) +{ + ILogConfiguration configuration; + configuration[CFG_BOOL_ENABLE_NET_DETECT] = false; + ILogManager* manager = LogManagerFactory::Create(configuration); + ASSERT_NE(manager, nullptr); + MATW::NetworkDetector detector; + SelectBackend(detector); + BlockingStopDetectorOnNetworkChange listener(detector, true); + auto entered = listener.GetEnteredFuture(); + manager->AddEventListener(EVT_NET_CHANGED, listener); + EXPECT_TRUE(detector.Start()); + EXPECT_EQ(entered.wait_for(std::chrono::seconds(5)), std::future_status::ready); + auto restarted = std::async(std::launch::async, [&] { return detector.Start(); }); + EXPECT_EQ(restarted.wait_for(std::chrono::milliseconds(50)), std::future_status::timeout); + listener.Release(); + EXPECT_EQ(restarted.wait_for(std::chrono::seconds(5)), std::future_status::ready); + EXPECT_TRUE(restarted.get()); + EXPECT_TRUE(detector.isUp()); + detector.Stop(); + manager->RemoveEventListener(EVT_NET_CHANGED, listener); + EXPECT_EQ(LogManagerFactory::Destroy(manager), STATUS_SUCCESS); +} + +TEST_P(NetworkDetectorBackendTests, QueuedRefreshRaceDoesNotOutliveStop) +{ + MATW::NetworkDetector detector; + SelectBackend(detector); + ASSERT_TRUE(detector.Start()); + std::atomic keepQueuing{true}; + std::thread callbacks([&] { + while (keepQueuing.load(std::memory_order_acquire)) + { + detector.QueueNetworkCostRefresh(); + } + }); + detector.Stop(); + keepQueuing.store(false, std::memory_order_release); + callbacks.join(); + EXPECT_FALSE(detector.QueueNetworkCostRefresh()); +} + +TEST_P(NetworkDetectorBackendTests, CostReadsDoNotWaitOnAnEarlierListenerAfterRestart) +{ + MATW::NetworkDetector detector; + SelectBackend(detector); + ASSERT_TRUE(detector.Start()); + std::atomic keepReading{true}; + std::promise firstRead; + auto readCompleted = firstRead.get_future(); + std::thread reader([&] { + detector.GetCurrentNetworkCost(); + firstRead.set_value(); + while (keepReading.load(std::memory_order_acquire)) + { + detector.GetCurrentNetworkCost(); + } + }); + EXPECT_EQ(readCompleted.wait_for(std::chrono::seconds(5)), std::future_status::ready); + for (unsigned iteration = 0; iteration < 20; ++iteration) + { + detector.Stop(); + EXPECT_TRUE(detector.Start()); + } + keepReading.store(false, std::memory_order_release); + detector.Stop(); + reader.join(); + EXPECT_FALSE(detector.isUp()); +} + +TEST_P(NetworkDetectorBackendTests, ListenerCanStopDetector) +{ + ILogConfiguration configuration; + configuration[CFG_BOOL_ENABLE_NET_DETECT] = false; + ILogManager* manager = LogManagerFactory::Create(configuration); + ASSERT_NE(manager, nullptr); + MATW::NetworkDetector detector; + SelectBackend(detector); + StopDetectorOnNetworkChange listener(detector); + auto stopped = listener.GetStoppedFuture(); + manager->AddEventListener(EVT_NET_CHANGED, listener); + EXPECT_TRUE(detector.Start()); + EXPECT_EQ(stopped.wait_for(std::chrono::seconds(5)), std::future_status::ready); + detector.Stop(); + EXPECT_FALSE(detector.isUp()); + manager->RemoveEventListener(EVT_NET_CHANGED, listener); + EXPECT_EQ(LogManagerFactory::Destroy(manager), STATUS_SUCCESS); +} + +TEST(NetworkDetectorTests, LegacyCostQueryFailurePreservesConnectivity) +{ + MATW::NetworkDetector detector; + MATW::NetworkDetectorTestAccess::FailLegacyCostQuery(detector); + for (unsigned iteration = 0; iteration < 3; ++iteration) + { + ASSERT_TRUE(detector.Start()); + EXPECT_TRUE(detector.isUp()); + EXPECT_FALSE(MATW::NetworkDetectorTestAccess::HasLegacyCost(detector)); + EXPECT_EQ(detector.GetCurrentNetworkCost(), NetworkCost_Unknown); + EXPECT_EQ(MATW::NetworkDetectorTestAccess::LegacySubscriptionCount(detector), 3u); + detector.Stop(); + EXPECT_FALSE(MATW::NetworkDetectorTestAccess::HasLegacyManager(detector)); + } + MATW::NetworkDetectorTestAccess::UseLegacyBackendWithoutCost(detector); + EXPECT_TRUE(detector.Start()); + EXPECT_EQ(detector.GetCurrentNetworkCost(), NetworkCost_Unknown); + detector.Stop(); +} + +TEST(NetworkDetectorTests, PartialLegacySubscriptionFailureCleansUpAndCanRetry) +{ + MATW::NetworkDetector detector; + MATW::NetworkDetectorTestAccess::FailLegacySubscription(detector); + for (unsigned iteration = 0; iteration < 3; ++iteration) + { + EXPECT_FALSE(detector.Start()); + EXPECT_FALSE(detector.isUp()); + EXPECT_FALSE(detector.QueueNetworkCostRefresh()); + EXPECT_EQ(MATW::NetworkDetectorTestAccess::LegacySubscriptionCount(detector), 0u); + detector.Stop(); + } + MATW::NetworkDetectorTestAccess::UseLegacyBackend(detector); + EXPECT_TRUE(detector.Start()); + detector.Stop(); +} + +TEST(NetworkDetectorTests, DisconnectFailureStillCompletesApartmentShutdown) +{ + MATW::NetworkDetector detector; + MATW::NetworkDetectorTestAccess::FailLegacyDisconnect(detector); + for (unsigned iteration = 0; iteration < 3; ++iteration) + { + ASSERT_TRUE(detector.Start()); + detector.Stop(); + EXPECT_FALSE(detector.isUp()); + EXPECT_FALSE(detector.QueueNetworkCostRefresh()); + EXPECT_FALSE(MATW::NetworkDetectorTestAccess::HasLegacyCost(detector)); + EXPECT_EQ(MATW::NetworkDetectorTestAccess::LegacySubscriptionCount(detector), 0u); + } +} + +INSTANTIATE_TEST_SUITE_P(NativeAndLegacy, NetworkDetectorBackendTests, + Values(NetworkDetectorBackend::Native, NetworkDetectorBackend::Legacy, + NetworkDetectorBackend::LegacyWithoutCost)); #endif