From fdee02f29dc850050b18b485a3976b6b14c3e2c5 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Marc-Andr=C3=A9=20Moreau?= Date: Thu, 24 Sep 2026 06:55:12 -0400 Subject: [PATCH 1/4] Fix DVC plugin class factory lifetime Capture the session ID rather than borrowing an instance pointer, and acquire the registered plugin while synchronized with removal and replacement. Exercise stale factories and replacement during plugin QueryInterface without opening an RDP session. Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> --- dll/MsRdpEx.cpp | 2 +- dll/RdpDvcClient.cpp | 80 ++++++------ dll/RdpDvcClient.h | 2 +- dll/RdpInstance.cpp | 43 +++++++ dll/RdpInstanceInternal.h | 2 + tests/logging/CMakeLists.txt | 13 ++ tests/logging/DvcFactoryLifetimeTest.cpp | 152 +++++++++++++++++++++++ tests/logging/GatewayShutdownFixture.def | 3 + tests/logging/PluginReferenceFixture.cpp | 23 ++++ 9 files changed, 282 insertions(+), 38 deletions(-) create mode 100644 tests/logging/DvcFactoryLifetimeTest.cpp diff --git a/dll/MsRdpEx.cpp b/dll/MsRdpEx.cpp index ac668b8..ba74dd6 100644 --- a/dll/MsRdpEx.cpp +++ b/dll/MsRdpEx.cpp @@ -42,7 +42,7 @@ HRESULT STDAPICALLTYPE DllGetClassObject(REFCLSID rclsid, REFIID riid, LPVOID* p CMsRdpExInstance* instance = MsRdpEx_InstanceManager_FindBySessionId((GUID*) pclsid); if (instance) { - hr = DllGetClassObject_DvcPlugin(rclsid, riid, ppv, (void*) instance); + hr = DllGetClassObject_DvcPlugin(rclsid, riid, ppv); MsRdpEx_LogPrint(DEBUG, "DllGetClassObject_DvcPlugin(%s, %s) with instance %p, hr = 0x%08X", clsid, iid, hr, instance); return hr; } diff --git a/dll/RdpDvcClient.cpp b/dll/RdpDvcClient.cpp index be17ed2..e50130e 100644 --- a/dll/RdpDvcClient.cpp +++ b/dll/RdpDvcClient.cpp @@ -7,6 +7,10 @@ #include +#include + +#include "RdpInstanceInternal.h" + // // CRdpDvcClient class // @@ -298,24 +302,18 @@ CRdpDvcPlugin::~CRdpDvcPlugin() // CDvcPluginClassFactory class -class CDvcPluginClassFactory : IClassFactory +class CDvcPluginClassFactory : public IClassFactory { public: - CDvcPluginClassFactory(CMsRdpExInstance* instance) - { - m_instance = instance; - } - - ~CDvcPluginClassFactory() - { - - } + explicit CDvcPluginClassFactory(REFCLSID sessionId) : m_sessionId(sessionId) {} // IUnknown interface public: HRESULT STDMETHODCALLTYPE QueryInterface(REFIID riid, LPVOID* ppvObject) { - HRESULT hr = E_NOINTERFACE; + if (!ppvObject) + return E_POINTER; + *ppvObject = NULL; char iid[MSRDPEX_GUID_STRING_SIZE]; MsRdpEx_GuidBinToStr((GUID*)&riid, iid, 0); @@ -323,17 +321,16 @@ class CDvcPluginClassFactory : IClassFactory MsRdpEx_LogPrint(DEBUG, "CDvcPluginClassFactory::QueryInterface(%s)", iid); if (riid == IID_IUnknown) { - *ppvObject = (LPVOID)((IUnknown*)this); - InterlockedIncrement(&m_refCount); - return S_OK; + *ppvObject = static_cast(this); } - if (riid == IID_IClassFactory) { - *ppvObject = (LPVOID)((IClassFactory*)this); - InterlockedIncrement(&m_refCount); - return S_OK; + else if (riid == IID_IClassFactory) { + *ppvObject = static_cast(this); } - return hr; + if (!*ppvObject) + return E_NOINTERFACE; + AddRef(); + return S_OK; } ULONG STDMETHODCALLTYPE AddRef() @@ -357,6 +354,12 @@ class CDvcPluginClassFactory : IClassFactory public: HRESULT STDMETHODCALLTYPE CreateInstance(IUnknown* pUnkOuter, REFIID riid, LPVOID* ppvObject) { + if (!ppvObject) + return E_POINTER; + *ppvObject = NULL; + if (pUnkOuter) + return CLASS_E_NOAGGREGATION; + HRESULT hr = E_NOINTERFACE; char iid[MSRDPEX_GUID_STRING_SIZE]; @@ -364,17 +367,22 @@ class CDvcPluginClassFactory : IClassFactory if (riid == IID_IWTSPlugin) { IUnknown* wtsPlugin = NULL; - IMsRdpExInstance* rdpInstance = (IMsRdpExInstance*)m_instance; - rdpInstance->GetWTSPluginObject((void**)&wtsPlugin); + hr = MsRdpEx_InstanceManager_AcquireWTSPluginBySessionId(&m_sessionId, &wtsPlugin); - if (wtsPlugin) { + if (SUCCEEDED(hr) && wtsPlugin) { MsRdpEx_LogPrint(DEBUG, "CDvcPluginClassFactory using registered WTSPlugin"); hr = wtsPlugin->QueryInterface(riid, ppvObject); + wtsPlugin->Release(); } - else { + else if (SUCCEEDED(hr)) { MsRdpEx_LogPrint(DEBUG, "CDvcPluginClassFactory using built-in WTSPlugin"); - CRdpDvcPlugin* dvcPlugin = new CRdpDvcPlugin(); - hr = dvcPlugin->QueryInterface(riid, ppvObject); + CRdpDvcPlugin* dvcPlugin = new (std::nothrow) CRdpDvcPlugin(); + if (!dvcPlugin) + hr = E_OUTOFMEMORY; + else { + hr = dvcPlugin->QueryInterface(riid, ppvObject); + dvcPlugin->Release(); + } } } @@ -390,20 +398,20 @@ class CDvcPluginClassFactory : IClassFactory } private: - ULONG m_refCount = 0; - CMsRdpExInstance* m_instance = NULL; + ULONG m_refCount = 1; + GUID m_sessionId; }; -HRESULT STDAPICALLTYPE DllGetClassObject_DvcPlugin(REFCLSID rclsid, REFIID riid, LPVOID* ppv, void* instance) +HRESULT STDAPICALLTYPE DllGetClassObject_DvcPlugin(REFCLSID rclsid, REFIID riid, LPVOID* ppv) { - HRESULT hr = E_NOINTERFACE; - - if (riid == (REFIID) IID_IClassFactory) - { - CDvcPluginClassFactory* classFactory = new CDvcPluginClassFactory((CMsRdpExInstance*) instance); - *ppv = (LPVOID) classFactory; - hr = S_OK; - } + if (!ppv) + return E_POINTER; + *ppv = NULL; + CDvcPluginClassFactory* classFactory = new (std::nothrow) CDvcPluginClassFactory(rclsid); + if (!classFactory) + return E_OUTOFMEMORY; + HRESULT hr = classFactory->QueryInterface(riid, ppv); + classFactory->Release(); return hr; } diff --git a/dll/RdpDvcClient.h b/dll/RdpDvcClient.h index 1ae977f..63c5b15 100644 --- a/dll/RdpDvcClient.h +++ b/dll/RdpDvcClient.h @@ -75,6 +75,6 @@ class CRdpDvcPlugin : IWTSVirtualChannel* m_pChannel = NULL; }; -HRESULT STDAPICALLTYPE DllGetClassObject_DvcPlugin(REFCLSID rclsid, REFIID riid, LPVOID* ppv, void* instance); +HRESULT STDAPICALLTYPE DllGetClassObject_DvcPlugin(REFCLSID rclsid, REFIID riid, LPVOID* ppv); #endif /* MSRDPEX_DVC_CLIENT_H */ diff --git a/dll/RdpInstance.cpp b/dll/RdpInstance.cpp index 7f79db3..5aacb61 100644 --- a/dll/RdpInstance.cpp +++ b/dll/RdpInstance.cpp @@ -541,14 +541,26 @@ class CMsRdpExInstance : public IMsRdpExInstance HRESULT STDMETHODCALLTYPE SetWTSPluginObject(LPVOID pvObject) { + AcquireSRWLockExclusive(&m_WTSPluginLock); IUnknown* previousPlugin = m_WTSPlugin; m_WTSPlugin = (IUnknown*)pvObject; + ReleaseSRWLockExclusive(&m_WTSPluginLock); if (previousPlugin) { previousPlugin->Release(); } return S_OK; } + IUnknown* AcquireWTSPluginObject() + { + AcquireSRWLockShared(&m_WTSPluginLock); + IUnknown* plugin = m_WTSPlugin; + if (plugin) + plugin->AddRef(); + ReleaseSRWLockShared(&m_WTSPluginLock); + return plugin; + } + public: GUID m_sessionId; ULONG m_refCount; @@ -563,6 +575,7 @@ class CMsRdpExInstance : public IMsRdpExInstance int32_t m_LastMousePosX = 0; int32_t m_LastMousePosY = 0; IUnknown* m_WTSPlugin = NULL; + SRWLOCK m_WTSPluginLock = SRWLOCK_INIT; LONG m_GdiReconnectPending = 0; LONG m_GdiReconnectAttempts = 0; LONG m_HardwareCaptureFrameReceived = 0; @@ -1105,6 +1118,36 @@ CMsRdpExInstance* MsRdpEx_InstanceManager_FindBySessionId(GUID* sessionId) return found ? obj : NULL; } +HRESULT MsRdpEx_InstanceManager_AcquireWTSPluginBySessionId(const GUID* sessionId, IUnknown** plugin) +{ + if (!plugin) + return E_POINTER; + *plugin = NULL; + if (!sessionId) + return E_INVALIDARG; + + MsRdpEx_InstanceManager* ctx = g_InstanceManager; + if (!ctx) + return REGDB_E_CLASSNOTREG; + + bool found = false; + MsRdpEx_ArrayListIt* it = MsRdpEx_ArrayList_It(ctx->instances, MSRDPEX_ITERATOR_FLAG_EXCLUSIVE); + + while (!MsRdpEx_ArrayListIt_Done(it)) + { + CMsRdpExInstance* instance = (CMsRdpExInstance*)MsRdpEx_ArrayListIt_Next(it); + if (MsRdpEx_GuidIsEqual(&instance->m_sessionId, sessionId)) + { + *plugin = instance->AcquireWTSPluginObject(); + found = true; + break; + } + } + + MsRdpEx_ArrayListIt_Finish(it); + return found ? S_OK : REGDB_E_CLASSNOTREG; +} + CMsRdpExtendedSettings* MsRdpEx_FindExtendedSettingsBySessionId(GUID* sessionId) { CMsRdpExInstance* instance = NULL; diff --git a/dll/RdpInstanceInternal.h b/dll/RdpInstanceInternal.h index 9bb7a9c..0fc5fac 100644 --- a/dll/RdpInstanceInternal.h +++ b/dll/RdpInstanceInternal.h @@ -5,6 +5,8 @@ IMsRdpExInstance* MsRdpEx_InstanceManager_AcquireByOutputPresenterHwnd(HWND hWnd); IMsRdpExInstance* MsRdpEx_InstanceManager_AcquireByInputCaptureHwnd(HWND hWnd); +// A non-null plugin returned on S_OK owns one reference for the caller to release. +HRESULT MsRdpEx_InstanceManager_AcquireWTSPluginBySessionId(const GUID* sessionId, IUnknown** plugin); UINT MsRdpEx_Instance_GetGdiReconnectMessage(); UINT_PTR MsRdpEx_Instance_GetHardwareCaptureWatchdogTimerId(); bool MsRdpEx_Instance_RequestGdiReconnect(IMsRdpExInstance* instance); diff --git a/tests/logging/CMakeLists.txt b/tests/logging/CMakeLists.txt index d63ee93..dc320c2 100644 --- a/tests/logging/CMakeLists.txt +++ b/tests/logging/CMakeLists.txt @@ -69,3 +69,16 @@ foreach(scenario replacement same-pointer clearing destruction) "$" ${scenario}) set_tests_properties(logging.plugin-reference.${scenario} PROPERTIES TIMEOUT 45) endforeach() + +add_executable(MsRdpEx_DvcFactoryLifetimeTest DvcFactoryLifetimeTest.cpp) +target_include_directories(MsRdpEx_DvcFactoryLifetimeTest PRIVATE "${PROJECT_SOURCE_DIR}/dll") +target_compile_features(MsRdpEx_DvcFactoryLifetimeTest PRIVATE cxx_std_17) +target_compile_options(MsRdpEx_DvcFactoryLifetimeTest PRIVATE /W4 /EHsc) +target_link_libraries(MsRdpEx_DvcFactoryLifetimeTest PRIVATE uuid.lib) +add_dependencies(MsRdpEx_DvcFactoryLifetimeTest MsRdpEx_GatewayShutdownTestDll) +foreach(scenario removed replacement) + add_test(NAME logging.dvc-factory.${scenario} + COMMAND MsRdpEx_DvcFactoryLifetimeTest + "$" ${scenario}) + set_tests_properties(logging.dvc-factory.${scenario} PROPERTIES TIMEOUT 45) +endforeach() diff --git a/tests/logging/DvcFactoryLifetimeTest.cpp b/tests/logging/DvcFactoryLifetimeTest.cpp new file mode 100644 index 0000000..9b0c9a7 --- /dev/null +++ b/tests/logging/DvcFactoryLifetimeTest.cpp @@ -0,0 +1,152 @@ +#include + +#include +#include +#include +#include + +static void Check(bool condition, const char* message) +{ + if (!condition) throw std::runtime_error(message); +} + +class CountedPlugin : public IWTSPlugin +{ +public: + HRESULT STDMETHODCALLTYPE QueryInterface(REFIID iid, void** object) override + { + if (!object) return E_POINTER; + *object = NULL; + if (iid != IID_IUnknown && iid != IID_IWTSPlugin) return E_NOINTERFACE; + if (replaceOnQuery) { + replaceOnQuery = false; + if (FAILED(instance->SetWTSPluginObject(replacement))) return E_FAIL; + refsAfterReplacement = refs; + } + *object = static_cast(this); + AddRef(); + return S_OK; + } + + ULONG STDMETHODCALLTYPE AddRef() override { return ++refs; } + ULONG STDMETHODCALLTYPE Release() override { return --refs; } + HRESULT STDMETHODCALLTYPE Initialize(IWTSVirtualChannelManager*) override { return S_OK; } + HRESULT STDMETHODCALLTYPE Connected() override { return S_OK; } + HRESULT STDMETHODCALLTYPE Disconnected(DWORD) override { return S_OK; } + HRESULT STDMETHODCALLTYPE Terminated() override { return S_OK; } + + ULONG refs = 1; + IMsRdpExInstance* instance = NULL; + IWTSPlugin* replacement = NULL; + bool replaceOnQuery = false; + ULONG refsAfterReplacement = 0; +}; + +int wmain(int argc, wchar_t** argv) +{ + try { + Check(argc == 3, "Expected test DLL path and scenario"); + Check(SetEnvironmentVariableW(L"MSRDPEX_HOOK_ENABLED", L"0") != 0, + "Could not disable ActiveX hooks"); + Check(SetEnvironmentVariableW(L"MSRDPEX_LOG_ENABLED", L"0") != 0, + "Could not disable logging"); + HMODULE module = LoadLibraryW(argv[1]); + Check(module != NULL, "Could not load test DLL"); + + using Create = IMsRdpExInstance*(*)(); + using Register = bool(*)(IMsRdpExInstance*); + using Factory = HRESULT(*)(REFCLSID, IClassFactory**); + auto create = reinterpret_cast( + GetProcAddress(module, "CreatePluginReferenceInstance")); + auto registerInstance = reinterpret_cast( + GetProcAddress(module, "RegisterPluginReferenceInstance")); + auto unregisterInstance = reinterpret_cast( + GetProcAddress(module, "UnregisterPluginReferenceInstance")); + auto createFactory = reinterpret_cast( + GetProcAddress(module, "CreatePluginReferenceFactory")); + Check(create && registerInstance && unregisterInstance && createFactory, + "DVC factory fixture exports missing"); + + IMsRdpExInstance* instance = create(); + Check(instance != NULL, "Could not create plugin instance"); + GUID sessionId = {}; + Check(SUCCEEDED(instance->GetSessionId(&sessionId)), "Could not get session ID"); + Check(registerInstance(instance), "Could not register plugin instance"); + + IClassFactory* factory = NULL; + Check(SUCCEEDED(createFactory(sessionId, &factory)) && factory, + "Could not create DVC plugin class factory"); + IUnknown* identity = NULL; + Check(SUCCEEDED(factory->QueryInterface(IID_IUnknown, (void**)&identity)) && identity, + "Class factory does not support IUnknown"); + Check(identity->Release() == 1, "Class factory initial reference is unbalanced"); + void* unsupported = factory; + Check(factory->QueryInterface(IID_IWTSPlugin, &unsupported) == E_NOINTERFACE && !unsupported, + "Unsupported factory interface did not clear the output"); + Check(factory->QueryInterface(IID_IUnknown, NULL) == E_POINTER, + "Null factory QueryInterface output was accepted"); + Check(factory->CreateInstance(NULL, IID_IWTSPlugin, NULL) == E_POINTER, + "Null CreateInstance output was accepted"); + IWTSPlugin* invalid = reinterpret_cast(factory); + Check(factory->CreateInstance(identity, IID_IWTSPlugin, (void**)&invalid) + == CLASS_E_NOAGGREGATION && !invalid, + "Factory accepted aggregation"); + + CountedPlugin first; + first.AddRef(); + Check(SUCCEEDED(instance->SetWTSPluginObject(&first)), "Initial plugin setter failed"); + IWTSPlugin* returned = NULL; + Check(SUCCEEDED(factory->CreateInstance(NULL, IID_IWTSPlugin, (void**)&returned)) + && returned == static_cast(&first), + "Factory did not return the registered plugin"); + Check(returned->Release() == 2, "Factory did not balance its temporary plugin reference"); + + const std::wstring scenario = argv[2]; + if (scenario == L"removed") { + Check(unregisterInstance(instance), "Could not remove plugin instance"); + Check(instance->Release() == 0, "Instance remained alive after removal"); + Check(first.refs == 1, "Instance did not release the plugin"); + returned = reinterpret_cast(factory); + Check(factory->CreateInstance(NULL, IID_IWTSPlugin, (void**)&returned) + == REGDB_E_CLASSNOTREG && !returned, + "Factory used a removed instance or fell back to a built-in plugin"); + Check(first.Release() == 0, "Test plugin reference not balanced"); + } else if (scenario == L"replacement") { + CountedPlugin second; + second.AddRef(); + first.instance = instance; + first.replacement = &second; + first.replaceOnQuery = true; + Check(SUCCEEDED(factory->CreateInstance(NULL, IID_IWTSPlugin, (void**)&returned)) + && returned == static_cast(&first), + "Factory failed when the plugin replaced itself during QueryInterface"); + Check(first.refsAfterReplacement == 2, + "Factory did not retain the old plugin across replacement"); + Check(returned->Release() == 1, "Factory leaked its temporary old plugin reference"); + Check(SUCCEEDED(factory->CreateInstance(NULL, IID_IWTSPlugin, (void**)&returned)) + && returned == static_cast(&second), + "Existing factory did not return the replacement plugin"); + Check(returned->Release() == 2, "Factory did not balance the replacement reference"); + Check(SUCCEEDED(instance->SetWTSPluginObject(NULL)), "Plugin clearing failed"); + Check(second.refs == 1, "Clearing did not release the replacement plugin"); + Check(SUCCEEDED(factory->CreateInstance(NULL, IID_IWTSPlugin, (void**)&returned)) + && returned && returned != static_cast(&second), + "Live session without a plugin did not use the built-in plugin"); + returned->Release(); + Check(unregisterInstance(instance), "Could not remove plugin instance"); + Check(instance->Release() == 0, "Instance remained alive after removal"); + Check(first.Release() == 0 && second.Release() == 0, + "Test plugin references not balanced"); + } else { + throw std::runtime_error("Unknown DVC factory scenario"); + } + + Check(factory->Release() == 0, "Class factory retained an extra reference"); + Check(FreeLibrary(module) != FALSE, "Could not unload test DLL"); + std::wcout << L"PASS DVC factory lifetime: " << scenario << L'\n'; + return 0; + } catch (const std::exception& error) { + std::cerr << error.what() << '\n'; + return 1; + } +} diff --git a/tests/logging/GatewayShutdownFixture.def b/tests/logging/GatewayShutdownFixture.def index 2735ae0..3284527 100644 --- a/tests/logging/GatewayShutdownFixture.def +++ b/tests/logging/GatewayShutdownFixture.def @@ -1,3 +1,6 @@ EXPORTS PrepareGatewayShutdown CreatePluginReferenceInstance + RegisterPluginReferenceInstance + UnregisterPluginReferenceInstance + CreatePluginReferenceFactory diff --git a/tests/logging/PluginReferenceFixture.cpp b/tests/logging/PluginReferenceFixture.cpp index 819dd03..4457ddb 100644 --- a/tests/logging/PluginReferenceFixture.cpp +++ b/tests/logging/PluginReferenceFixture.cpp @@ -1,6 +1,29 @@ #include "../../dll/RdpInstance.cpp" +#include "../../dll/RdpDvcClient.h" extern "C" IMsRdpExInstance* CreatePluginReferenceInstance() { return CMsRdpExInstance_New(NULL); } + +extern "C" bool RegisterPluginReferenceInstance(IMsRdpExInstance* instance) +{ + if (!MsRdpEx_InstanceManager_Get()) + return false; + if (MsRdpEx_InstanceManager_Add((CMsRdpExInstance*)instance)) + return true; + MsRdpEx_InstanceManager_Release(); + return false; +} + +extern "C" bool UnregisterPluginReferenceInstance(IMsRdpExInstance* instance) +{ + bool removed = MsRdpEx_InstanceManager_Remove((CMsRdpExInstance*)instance); + MsRdpEx_InstanceManager_Release(); + return removed; +} + +extern "C" HRESULT CreatePluginReferenceFactory(REFCLSID sessionId, IClassFactory** factory) +{ + return DllGetClassObject_DvcPlugin(sessionId, IID_IClassFactory, (void**)factory); +} From 003cafbd5280c8fd3a269bd0eca3f644fab1d335 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Marc-Andr=C3=A9=20Moreau?= Date: Thu, 24 Sep 2026 08:29:06 -0400 Subject: [PATCH 2/4] Prevent reentrant DVC plugin deadlock and manager teardown race Pin plugins with an internal nothrow reference holder, so COM AddRef and Release run outside locks. Coordinate session lookups with manager teardown, exercise the production factory route, and cover reentrant callbacks plus in-flight shutdown. Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> --- dll/MsRdpEx.cpp | 4 +- dll/RdpDvcClient.cpp | 4 +- dll/RdpInstance.cpp | 100 +++++++++++++++++++---- dll/RdpInstanceInternal.h | 19 ++++- tests/logging/CMakeLists.txt | 2 +- tests/logging/DvcFactoryLifetimeTest.cpp | 84 ++++++++++++++++++- tests/logging/PluginReferenceFixture.cpp | 3 +- 7 files changed, 191 insertions(+), 25 deletions(-) diff --git a/dll/MsRdpEx.cpp b/dll/MsRdpEx.cpp index ba74dd6..2c4b374 100644 --- a/dll/MsRdpEx.cpp +++ b/dll/MsRdpEx.cpp @@ -9,6 +9,7 @@ #include #include "RdpDvcClient.h" +#include "RdpInstanceInternal.h" #include #include @@ -39,11 +40,12 @@ HRESULT STDAPICALLTYPE DllGetClassObject(REFCLSID rclsid, REFIID riid, LPVOID* p MsRdpEx_GuidBinToStr(pclsid, clsid, 0); MsRdpEx_GuidBinToStr(piid, iid, 0); - CMsRdpExInstance* instance = MsRdpEx_InstanceManager_FindBySessionId((GUID*) pclsid); + IMsRdpExInstance* instance = MsRdpEx_InstanceManager_AcquireBySessionId(pclsid); if (instance) { hr = DllGetClassObject_DvcPlugin(rclsid, riid, ppv); MsRdpEx_LogPrint(DEBUG, "DllGetClassObject_DvcPlugin(%s, %s) with instance %p, hr = 0x%08X", clsid, iid, hr, instance); + instance->Release(); return hr; } diff --git a/dll/RdpDvcClient.cpp b/dll/RdpDvcClient.cpp index e50130e..83ae637 100644 --- a/dll/RdpDvcClient.cpp +++ b/dll/RdpDvcClient.cpp @@ -366,12 +366,12 @@ class CDvcPluginClassFactory : public IClassFactory MsRdpEx_GuidBinToStr((GUID*)&riid, iid, 0); if (riid == IID_IWTSPlugin) { - IUnknown* wtsPlugin = NULL; + MsRdpEx_WTSPluginReference* wtsPlugin = NULL; hr = MsRdpEx_InstanceManager_AcquireWTSPluginBySessionId(&m_sessionId, &wtsPlugin); if (SUCCEEDED(hr) && wtsPlugin) { MsRdpEx_LogPrint(DEBUG, "CDvcPluginClassFactory using registered WTSPlugin"); - hr = wtsPlugin->QueryInterface(riid, ppvObject); + hr = wtsPlugin->Get()->QueryInterface(riid, ppvObject); wtsPlugin->Release(); } else if (SUCCEEDED(hr)) { diff --git a/dll/RdpInstance.cpp b/dll/RdpInstance.cpp index 5aacb61..d3fc332 100644 --- a/dll/RdpInstance.cpp +++ b/dll/RdpInstance.cpp @@ -15,6 +15,27 @@ extern "C" const GUID IID_IMsRdpExInstance; +MsRdpEx_WTSPluginReference::MsRdpEx_WTSPluginReference(IUnknown* plugin) : m_plugin(plugin) {} + +IUnknown* MsRdpEx_WTSPluginReference::Get() const +{ + return m_plugin; +} + +void MsRdpEx_WTSPluginReference::AddRef() +{ + InterlockedIncrement(&m_refCount); +} + +void MsRdpEx_WTSPluginReference::Release() +{ + if (InterlockedDecrement(&m_refCount) == 0) { + IUnknown* plugin = m_plugin; + delete this; + plugin->Release(); + } +} + class CMsRdpExInstance : public IMsRdpExInstance { public: @@ -50,9 +71,8 @@ class CMsRdpExInstance : public IMsRdpExInstance m_pMsRdpExtendedSettings->Release(); } - if (m_WTSPlugin) { + if (m_WTSPlugin) m_WTSPlugin->Release(); - } } // IUnknown interface @@ -535,26 +555,36 @@ class CMsRdpExInstance : public IMsRdpExInstance public: HRESULT STDMETHODCALLTYPE GetWTSPluginObject(LPVOID* ppvObject) { - *ppvObject = m_WTSPlugin; + AcquireSRWLockShared(&m_WTSPluginLock); + *ppvObject = m_WTSPlugin ? m_WTSPlugin->Get() : NULL; + ReleaseSRWLockShared(&m_WTSPluginLock); return S_OK; } HRESULT STDMETHODCALLTYPE SetWTSPluginObject(LPVOID pvObject) { + MsRdpEx_WTSPluginReference* replacement = NULL; + if (pvObject) { + replacement = new (std::nothrow) MsRdpEx_WTSPluginReference((IUnknown*)pvObject); + if (!replacement) { + ((IUnknown*)pvObject)->Release(); + return E_OUTOFMEMORY; + } + } + AcquireSRWLockExclusive(&m_WTSPluginLock); - IUnknown* previousPlugin = m_WTSPlugin; - m_WTSPlugin = (IUnknown*)pvObject; + MsRdpEx_WTSPluginReference* previous = m_WTSPlugin; + m_WTSPlugin = replacement; ReleaseSRWLockExclusive(&m_WTSPluginLock); - if (previousPlugin) { - previousPlugin->Release(); - } + if (previous) + previous->Release(); return S_OK; } - IUnknown* AcquireWTSPluginObject() + MsRdpEx_WTSPluginReference* AcquireWTSPluginObject() { AcquireSRWLockShared(&m_WTSPluginLock); - IUnknown* plugin = m_WTSPlugin; + MsRdpEx_WTSPluginReference* plugin = m_WTSPlugin; if (plugin) plugin->AddRef(); ReleaseSRWLockShared(&m_WTSPluginLock); @@ -574,7 +604,7 @@ class CMsRdpExInstance : public IMsRdpExInstance CMsRdpExtendedSettings* m_pMsRdpExtendedSettings = NULL; int32_t m_LastMousePosX = 0; int32_t m_LastMousePosY = 0; - IUnknown* m_WTSPlugin = NULL; + MsRdpEx_WTSPluginReference* m_WTSPlugin = NULL; SRWLOCK m_WTSPluginLock = SRWLOCK_INIT; LONG m_GdiReconnectPending = 0; LONG m_GdiReconnectAttempts = 0; @@ -785,6 +815,7 @@ void MsRdpEx_InstanceManager_Free(MsRdpEx_InstanceManager* ctx); static int g_RefCount = 0; static MsRdpEx_InstanceManager* g_InstanceManager = NULL; +static SRWLOCK g_InstanceManagerLock = SRWLOCK_INIT; bool MsRdpEx_InstanceManager_Add(CMsRdpExInstance* instance) { @@ -1118,7 +1149,35 @@ CMsRdpExInstance* MsRdpEx_InstanceManager_FindBySessionId(GUID* sessionId) return found ? obj : NULL; } -HRESULT MsRdpEx_InstanceManager_AcquireWTSPluginBySessionId(const GUID* sessionId, IUnknown** plugin) +IMsRdpExInstance* MsRdpEx_InstanceManager_AcquireBySessionId(const GUID* sessionId) +{ + AcquireSRWLockShared(&g_InstanceManagerLock); + MsRdpEx_InstanceManager* ctx = g_InstanceManager; + if (!ctx) { + ReleaseSRWLockShared(&g_InstanceManagerLock); + return NULL; + } + + IMsRdpExInstance* instance = NULL; + MsRdpEx_ArrayListIt* it = MsRdpEx_ArrayList_It(ctx->instances, MSRDPEX_ITERATOR_FLAG_EXCLUSIVE); + while (!MsRdpEx_ArrayListIt_Done(it)) + { + CMsRdpExInstance* candidate = (CMsRdpExInstance*)MsRdpEx_ArrayListIt_Next(it); + if (MsRdpEx_GuidIsEqual(&candidate->m_sessionId, sessionId)) + { + candidate->AddRef(); + instance = candidate; + break; + } + } + + MsRdpEx_ArrayListIt_Finish(it); + ReleaseSRWLockShared(&g_InstanceManagerLock); + return instance; +} + +HRESULT MsRdpEx_InstanceManager_AcquireWTSPluginBySessionId( + const GUID* sessionId, MsRdpEx_WTSPluginReference** plugin) { if (!plugin) return E_POINTER; @@ -1126,9 +1185,12 @@ HRESULT MsRdpEx_InstanceManager_AcquireWTSPluginBySessionId(const GUID* sessionI if (!sessionId) return E_INVALIDARG; + AcquireSRWLockShared(&g_InstanceManagerLock); MsRdpEx_InstanceManager* ctx = g_InstanceManager; - if (!ctx) + if (!ctx) { + ReleaseSRWLockShared(&g_InstanceManagerLock); return REGDB_E_CLASSNOTREG; + } bool found = false; MsRdpEx_ArrayListIt* it = MsRdpEx_ArrayList_It(ctx->instances, MSRDPEX_ITERATOR_FLAG_EXCLUSIVE); @@ -1145,6 +1207,7 @@ HRESULT MsRdpEx_InstanceManager_AcquireWTSPluginBySessionId(const GUID* sessionI } MsRdpEx_ArrayListIt_Finish(it); + ReleaseSRWLockShared(&g_InstanceManagerLock); return found ? S_OK : REGDB_E_CLASSNOTREG; } @@ -1233,24 +1296,31 @@ void MsRdpEx_InstanceManager_Free(MsRdpEx_InstanceManager* ctx) MsRdpEx_InstanceManager* MsRdpEx_InstanceManager_Get() { + AcquireSRWLockExclusive(&g_InstanceManagerLock); if (!g_InstanceManager) g_InstanceManager = MsRdpEx_InstanceManager_New(); g_RefCount++; - return g_InstanceManager; + MsRdpEx_InstanceManager* ctx = g_InstanceManager; + ReleaseSRWLockExclusive(&g_InstanceManagerLock); + return ctx; } void MsRdpEx_InstanceManager_Release() { + AcquireSRWLockExclusive(&g_InstanceManagerLock); g_RefCount--; if (g_RefCount < 0) g_RefCount = 0; + MsRdpEx_InstanceManager* ctx = NULL; if (g_InstanceManager && (g_RefCount < 1)) { - MsRdpEx_InstanceManager_Free(g_InstanceManager); + ctx = g_InstanceManager; g_InstanceManager = NULL; } + ReleaseSRWLockExclusive(&g_InstanceManagerLock); + MsRdpEx_InstanceManager_Free(ctx); } diff --git a/dll/RdpInstanceInternal.h b/dll/RdpInstanceInternal.h index 0fc5fac..c2889e1 100644 --- a/dll/RdpInstanceInternal.h +++ b/dll/RdpInstanceInternal.h @@ -3,10 +3,25 @@ #include +class MsRdpEx_WTSPluginReference +{ +public: + explicit MsRdpEx_WTSPluginReference(IUnknown* plugin); + IUnknown* Get() const; + void AddRef(); + void Release(); + +private: + volatile LONG m_refCount = 1; + IUnknown* m_plugin; +}; + IMsRdpExInstance* MsRdpEx_InstanceManager_AcquireByOutputPresenterHwnd(HWND hWnd); IMsRdpExInstance* MsRdpEx_InstanceManager_AcquireByInputCaptureHwnd(HWND hWnd); -// A non-null plugin returned on S_OK owns one reference for the caller to release. -HRESULT MsRdpEx_InstanceManager_AcquireWTSPluginBySessionId(const GUID* sessionId, IUnknown** plugin); +IMsRdpExInstance* MsRdpEx_InstanceManager_AcquireBySessionId(const GUID* sessionId); +// The returned holder owns one internal reference; call Release when done. +HRESULT MsRdpEx_InstanceManager_AcquireWTSPluginBySessionId( + const GUID* sessionId, MsRdpEx_WTSPluginReference** plugin); UINT MsRdpEx_Instance_GetGdiReconnectMessage(); UINT_PTR MsRdpEx_Instance_GetHardwareCaptureWatchdogTimerId(); bool MsRdpEx_Instance_RequestGdiReconnect(IMsRdpExInstance* instance); diff --git a/tests/logging/CMakeLists.txt b/tests/logging/CMakeLists.txt index dc320c2..ede3ac1 100644 --- a/tests/logging/CMakeLists.txt +++ b/tests/logging/CMakeLists.txt @@ -76,7 +76,7 @@ target_compile_features(MsRdpEx_DvcFactoryLifetimeTest PRIVATE cxx_std_17) target_compile_options(MsRdpEx_DvcFactoryLifetimeTest PRIVATE /W4 /EHsc) target_link_libraries(MsRdpEx_DvcFactoryLifetimeTest PRIVATE uuid.lib) add_dependencies(MsRdpEx_DvcFactoryLifetimeTest MsRdpEx_GatewayShutdownTestDll) -foreach(scenario removed replacement) +foreach(scenario removed replacement addref-reentrant release-reentrant manager-shutdown) add_test(NAME logging.dvc-factory.${scenario} COMMAND MsRdpEx_DvcFactoryLifetimeTest "$" ${scenario}) diff --git a/tests/logging/DvcFactoryLifetimeTest.cpp b/tests/logging/DvcFactoryLifetimeTest.cpp index 9b0c9a7..b367738 100644 --- a/tests/logging/DvcFactoryLifetimeTest.cpp +++ b/tests/logging/DvcFactoryLifetimeTest.cpp @@ -3,6 +3,7 @@ #include #include #include +#include #include static void Check(bool condition, const char* message) @@ -18,6 +19,10 @@ class CountedPlugin : public IWTSPlugin if (!object) return E_POINTER; *object = NULL; if (iid != IID_IUnknown && iid != IID_IWTSPlugin) return E_NOINTERFACE; + if (queryEntered) { + SetEvent(queryEntered); + if (WaitForSingleObject(queryResume, 5000) != WAIT_OBJECT_0) return E_FAIL; + } if (replaceOnQuery) { replaceOnQuery = false; if (FAILED(instance->SetWTSPluginObject(replacement))) return E_FAIL; @@ -28,8 +33,22 @@ class CountedPlugin : public IWTSPlugin return S_OK; } - ULONG STDMETHODCALLTYPE AddRef() override { return ++refs; } - ULONG STDMETHODCALLTYPE Release() override { return --refs; } + ULONG STDMETHODCALLTYPE AddRef() override + { + if (clearOnAddRef) { + clearOnAddRef = false; + instance->SetWTSPluginObject(NULL); + } + return ++refs; + } + ULONG STDMETHODCALLTYPE Release() override + { + if (clearOnRelease) { + clearOnRelease = false; + instance->SetWTSPluginObject(NULL); + } + return --refs; + } HRESULT STDMETHODCALLTYPE Initialize(IWTSVirtualChannelManager*) override { return S_OK; } HRESULT STDMETHODCALLTYPE Connected() override { return S_OK; } HRESULT STDMETHODCALLTYPE Disconnected(DWORD) override { return S_OK; } @@ -39,7 +58,11 @@ class CountedPlugin : public IWTSPlugin IMsRdpExInstance* instance = NULL; IWTSPlugin* replacement = NULL; bool replaceOnQuery = false; + bool clearOnAddRef = false; + bool clearOnRelease = false; ULONG refsAfterReplacement = 0; + HANDLE queryEntered = NULL; + HANDLE queryResume = NULL; }; int wmain(int argc, wchar_t** argv) @@ -137,6 +160,63 @@ int wmain(int argc, wchar_t** argv) Check(instance->Release() == 0, "Instance remained alive after removal"); Check(first.Release() == 0 && second.Release() == 0, "Test plugin references not balanced"); + } else if (scenario == L"addref-reentrant") { + first.instance = instance; + first.clearOnAddRef = true; + Check(SUCCEEDED(factory->CreateInstance(NULL, IID_IWTSPlugin, (void**)&returned)) + && returned == static_cast(&first), + "Plugin AddRef could not reenter the setter"); + void* borrowed = &first; + Check(SUCCEEDED(instance->GetWTSPluginObject(&borrowed)) && !borrowed, + "Reentrant AddRef did not clear the registered plugin"); + Check(returned->Release() == 1, "Reentrant AddRef leaked the plugin reference"); + Check(unregisterInstance(instance), "Could not remove plugin instance"); + Check(instance->Release() == 0, "Instance remained alive after removal"); + Check(first.Release() == 0, "Test plugin reference not balanced"); + } else if (scenario == L"release-reentrant") { + CountedPlugin second; + second.AddRef(); + first.instance = instance; + first.clearOnRelease = true; + Check(SUCCEEDED(instance->SetWTSPluginObject(&second)), + "Plugin Release could not reenter the setter"); + void* borrowed = &first; + Check(SUCCEEDED(instance->GetWTSPluginObject(&borrowed)) && !borrowed, + "Reentrant Release did not clear the replacement plugin"); + Check(first.refs == 1 && second.refs == 1, + "Reentrant Release did not balance the plugin references"); + Check(unregisterInstance(instance), "Could not remove plugin instance"); + Check(instance->Release() == 0, "Instance remained alive after removal"); + Check(first.Release() == 0 && second.Release() == 0, + "Test plugin references not balanced"); + } else if (scenario == L"manager-shutdown") { + HANDLE entered = CreateEventW(NULL, TRUE, FALSE, NULL); + HANDLE resume = CreateEventW(NULL, TRUE, FALSE, NULL); + Check(entered && resume, "Could not create query synchronization events"); + first.queryEntered = entered; + first.queryResume = resume; + HRESULT workerHr = E_FAIL; + std::thread worker([&]() { + IWTSPlugin* plugin = NULL; + workerHr = factory->CreateInstance(NULL, IID_IWTSPlugin, (void**)&plugin); + if (plugin) plugin->Release(); + }); + DWORD enteredResult = WaitForSingleObject(entered, 5000); + bool removed = unregisterInstance(instance); + ULONG remaining = instance->Release(); + SetEvent(resume); + worker.join(); + Check(enteredResult == WAIT_OBJECT_0, "Factory query did not start"); + Check(removed && remaining == 0, "Could not tear down manager during factory query"); + Check(workerHr == S_OK && first.refs == 1, + "In-flight factory query lost its plugin during manager teardown"); + returned = reinterpret_cast(factory); + Check(factory->CreateInstance(NULL, IID_IWTSPlugin, (void**)&returned) + == REGDB_E_CLASSNOTREG && !returned, + "Factory used manager after shutdown"); + CloseHandle(entered); + CloseHandle(resume); + Check(first.Release() == 0, "Test plugin reference not balanced"); } else { throw std::runtime_error("Unknown DVC factory scenario"); } diff --git a/tests/logging/PluginReferenceFixture.cpp b/tests/logging/PluginReferenceFixture.cpp index 4457ddb..b5f599f 100644 --- a/tests/logging/PluginReferenceFixture.cpp +++ b/tests/logging/PluginReferenceFixture.cpp @@ -1,5 +1,4 @@ #include "../../dll/RdpInstance.cpp" -#include "../../dll/RdpDvcClient.h" extern "C" IMsRdpExInstance* CreatePluginReferenceInstance() { @@ -25,5 +24,5 @@ extern "C" bool UnregisterPluginReferenceInstance(IMsRdpExInstance* instance) extern "C" HRESULT CreatePluginReferenceFactory(REFCLSID sessionId, IClassFactory** factory) { - return DllGetClassObject_DvcPlugin(sessionId, IID_IClassFactory, (void**)factory); + return MsRdpEx_DllGetClassObject(sessionId, IID_IClassFactory, (void**)factory); } From ec384d194906a37944170becb4b7dd17044b0aa5 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Marc-Andr=C3=A9=20Moreau?= Date: Thu, 24 Sep 2026 08:34:26 -0400 Subject: [PATCH 3/4] Preserve caller ownership when plugin registration fails Do not release the transferred reference on holder allocation failure: managed callers release it on failure. Document success-only transfer and test forced OOM, same-pointer retry, and plugin QueryInterface failure. Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> --- dll/RdpInstance.cpp | 4 +-- include/MsRdpEx/RdpInstance.h | 3 +- tests/logging/CMakeLists.txt | 2 +- tests/logging/DvcFactoryLifetimeTest.cpp | 42 +++++++++++++++++++++++- tests/logging/GatewayShutdownFixture.def | 1 + tests/logging/PluginReferenceFixture.cpp | 18 ++++++++++ 6 files changed, 64 insertions(+), 6 deletions(-) diff --git a/dll/RdpInstance.cpp b/dll/RdpInstance.cpp index d3fc332..67e0d4c 100644 --- a/dll/RdpInstance.cpp +++ b/dll/RdpInstance.cpp @@ -566,10 +566,8 @@ class CMsRdpExInstance : public IMsRdpExInstance MsRdpEx_WTSPluginReference* replacement = NULL; if (pvObject) { replacement = new (std::nothrow) MsRdpEx_WTSPluginReference((IUnknown*)pvObject); - if (!replacement) { - ((IUnknown*)pvObject)->Release(); + if (!replacement) return E_OUTOFMEMORY; - } } AcquireSRWLockExclusive(&m_WTSPluginLock); diff --git a/include/MsRdpEx/RdpInstance.h b/include/MsRdpEx/RdpInstance.h index d512dff..8c713f4 100644 --- a/include/MsRdpEx/RdpInstance.h +++ b/include/MsRdpEx/RdpInstance.h @@ -37,7 +37,8 @@ struct __declspec(novtable) virtual void __stdcall UnlockShadowBitmap() = 0; virtual void __stdcall GetLastMousePosition(int32_t* posX, int32_t* posY) = 0; virtual void __stdcall SetLastMousePosition(int32_t posX, int32_t posY) = 0; - // Get returns a borrowed pointer; Set consumes one owned reference, even for the same pointer. + // Get returns a borrowed pointer; Set consumes one owned reference on success, + // even for the same pointer. On failure, the caller retains its reference. virtual HRESULT __stdcall GetWTSPluginObject(LPVOID* ppvObject) = 0; virtual HRESULT __stdcall SetWTSPluginObject(LPVOID pvObject) = 0; virtual void __stdcall SetCursor(HCURSOR cursor) = 0; diff --git a/tests/logging/CMakeLists.txt b/tests/logging/CMakeLists.txt index ede3ac1..aba772e 100644 --- a/tests/logging/CMakeLists.txt +++ b/tests/logging/CMakeLists.txt @@ -76,7 +76,7 @@ target_compile_features(MsRdpEx_DvcFactoryLifetimeTest PRIVATE cxx_std_17) target_compile_options(MsRdpEx_DvcFactoryLifetimeTest PRIVATE /W4 /EHsc) target_link_libraries(MsRdpEx_DvcFactoryLifetimeTest PRIVATE uuid.lib) add_dependencies(MsRdpEx_DvcFactoryLifetimeTest MsRdpEx_GatewayShutdownTestDll) -foreach(scenario removed replacement addref-reentrant release-reentrant manager-shutdown) +foreach(scenario removed replacement addref-reentrant release-reentrant manager-shutdown allocation-failure query-failure) add_test(NAME logging.dvc-factory.${scenario} COMMAND MsRdpEx_DvcFactoryLifetimeTest "$" ${scenario}) diff --git a/tests/logging/DvcFactoryLifetimeTest.cpp b/tests/logging/DvcFactoryLifetimeTest.cpp index b367738..1c8cc86 100644 --- a/tests/logging/DvcFactoryLifetimeTest.cpp +++ b/tests/logging/DvcFactoryLifetimeTest.cpp @@ -23,6 +23,7 @@ class CountedPlugin : public IWTSPlugin SetEvent(queryEntered); if (WaitForSingleObject(queryResume, 5000) != WAIT_OBJECT_0) return E_FAIL; } + if (failQuery) return E_NOINTERFACE; if (replaceOnQuery) { replaceOnQuery = false; if (FAILED(instance->SetWTSPluginObject(replacement))) return E_FAIL; @@ -60,6 +61,7 @@ class CountedPlugin : public IWTSPlugin bool replaceOnQuery = false; bool clearOnAddRef = false; bool clearOnRelease = false; + bool failQuery = false; ULONG refsAfterReplacement = 0; HANDLE queryEntered = NULL; HANDLE queryResume = NULL; @@ -87,7 +89,10 @@ int wmain(int argc, wchar_t** argv) GetProcAddress(module, "UnregisterPluginReferenceInstance")); auto createFactory = reinterpret_cast( GetProcAddress(module, "CreatePluginReferenceFactory")); - Check(create && registerInstance && unregisterInstance && createFactory, + using FailAllocation = void(*)(); + auto failAllocation = reinterpret_cast( + GetProcAddress(module, "FailNextPluginHolderAllocation")); + Check(create && registerInstance && unregisterInstance && createFactory && failAllocation, "DVC factory fixture exports missing"); IMsRdpExInstance* instance = create(); @@ -217,6 +222,41 @@ int wmain(int argc, wchar_t** argv) CloseHandle(entered); CloseHandle(resume); Check(first.Release() == 0, "Test plugin reference not balanced"); + } else if (scenario == L"allocation-failure") { + CountedPlugin second; + second.AddRef(); + failAllocation(); + Check(instance->SetWTSPluginObject(&second) == E_OUTOFMEMORY, + "Plugin holder allocation failure was not returned"); + void* borrowed = NULL; + Check(SUCCEEDED(instance->GetWTSPluginObject(&borrowed)) && borrowed == &first, + "Failed setter replaced the registered plugin"); + Check(first.refs == 2 && second.refs == 2, + "Failed setter consumed or leaked a caller-owned reference"); + Check(second.Release() == 1, "Caller could not release the failed setter's reference"); + first.AddRef(); + failAllocation(); + Check(instance->SetWTSPluginObject(&first) == E_OUTOFMEMORY && first.refs == 3, + "Failed same-pointer setter consumed the caller's reference"); + Check(first.Release() == 2, "Caller could not release failed same-pointer reference"); + Check(SUCCEEDED(factory->CreateInstance(NULL, IID_IWTSPlugin, (void**)&returned)) + && returned == static_cast(&first), + "Failed same-pointer setter invalidated the registered plugin"); + Check(returned->Release() == 2, "Factory leaked a plugin reference after setter failure"); + Check(unregisterInstance(instance), "Could not remove plugin instance"); + Check(instance->Release() == 0, "Instance remained alive after removal"); + Check(first.Release() == 0 && second.Release() == 0, + "Test plugin references not balanced"); + } else if (scenario == L"query-failure") { + first.failQuery = true; + returned = reinterpret_cast(factory); + Check(factory->CreateInstance(NULL, IID_IWTSPlugin, (void**)&returned) + == E_NOINTERFACE && !returned, + "Failed plugin QueryInterface did not clear output and propagate the error"); + Check(first.refs == 2, "Failed plugin QueryInterface leaked its holder reference"); + Check(unregisterInstance(instance), "Could not remove plugin instance"); + Check(instance->Release() == 0, "Instance remained alive after removal"); + Check(first.Release() == 0, "Test plugin reference not balanced"); } else { throw std::runtime_error("Unknown DVC factory scenario"); } diff --git a/tests/logging/GatewayShutdownFixture.def b/tests/logging/GatewayShutdownFixture.def index 3284527..ae75de4 100644 --- a/tests/logging/GatewayShutdownFixture.def +++ b/tests/logging/GatewayShutdownFixture.def @@ -4,3 +4,4 @@ RegisterPluginReferenceInstance UnregisterPluginReferenceInstance CreatePluginReferenceFactory + FailNextPluginHolderAllocation diff --git a/tests/logging/PluginReferenceFixture.cpp b/tests/logging/PluginReferenceFixture.cpp index b5f599f..b5a28e5 100644 --- a/tests/logging/PluginReferenceFixture.cpp +++ b/tests/logging/PluginReferenceFixture.cpp @@ -1,5 +1,23 @@ #include "../../dll/RdpInstance.cpp" +#include + +static thread_local bool g_FailNextPluginHolderAllocation = false; + +void* __cdecl operator new(size_t size, const std::nothrow_t&) noexcept +{ + if (g_FailNextPluginHolderAllocation && size == sizeof(MsRdpEx_WTSPluginReference)) { + g_FailNextPluginHolderAllocation = false; + return nullptr; + } + return std::malloc(size ? size : 1); +} + +extern "C" void FailNextPluginHolderAllocation() +{ + g_FailNextPluginHolderAllocation = true; +} + extern "C" IMsRdpExInstance* CreatePluginReferenceInstance() { return CMsRdpExInstance_New(NULL); From 42124aefbb18be9970662145389402fead07cf70 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Marc-Andr=C3=A9=20Moreau?= Date: Thu, 24 Sep 2026 09:07:08 -0400 Subject: [PATCH 4/4] Detach WTS plugin before instance destructor callbacks Close the plugin slot before releasing the final holder so reentrant COM callbacks cannot access a freed holder or register a replacement on an instance being destroyed. Cover reentrant getters and setters during teardown. Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> --- dll/RdpInstance.cpp | 17 ++++++++++++-- tests/logging/CMakeLists.txt | 2 +- tests/logging/DvcFactoryLifetimeTest.cpp | 28 +++++++++++++++++++++++- 3 files changed, 43 insertions(+), 4 deletions(-) diff --git a/dll/RdpInstance.cpp b/dll/RdpInstance.cpp index 67e0d4c..7816fb2 100644 --- a/dll/RdpInstance.cpp +++ b/dll/RdpInstance.cpp @@ -53,6 +53,13 @@ class CMsRdpExInstance : public IMsRdpExInstance ~CMsRdpExInstance() { + // Plugin Release may call back into the getter or setter. + AcquireSRWLockExclusive(&m_WTSPluginLock); + m_WTSPluginClosing = true; + MsRdpEx_WTSPluginReference* plugin = m_WTSPlugin; + m_WTSPlugin = NULL; + ReleaseSRWLockExclusive(&m_WTSPluginLock); + if (m_hOutputPresenterWnd) KillTimer(m_hOutputPresenterWnd, MsRdpEx_Instance_GetHardwareCaptureWatchdogTimerId()); @@ -71,8 +78,8 @@ class CMsRdpExInstance : public IMsRdpExInstance m_pMsRdpExtendedSettings->Release(); } - if (m_WTSPlugin) - m_WTSPlugin->Release(); + if (plugin) + plugin->Release(); } // IUnknown interface @@ -571,6 +578,11 @@ class CMsRdpExInstance : public IMsRdpExInstance } AcquireSRWLockExclusive(&m_WTSPluginLock); + if (m_WTSPluginClosing) { + ReleaseSRWLockExclusive(&m_WTSPluginLock); + delete replacement; // Failure leaves the incoming COM reference with the caller. + return E_UNEXPECTED; + } MsRdpEx_WTSPluginReference* previous = m_WTSPlugin; m_WTSPlugin = replacement; ReleaseSRWLockExclusive(&m_WTSPluginLock); @@ -604,6 +616,7 @@ class CMsRdpExInstance : public IMsRdpExInstance int32_t m_LastMousePosY = 0; MsRdpEx_WTSPluginReference* m_WTSPlugin = NULL; SRWLOCK m_WTSPluginLock = SRWLOCK_INIT; + bool m_WTSPluginClosing = false; LONG m_GdiReconnectPending = 0; LONG m_GdiReconnectAttempts = 0; LONG m_HardwareCaptureFrameReceived = 0; diff --git a/tests/logging/CMakeLists.txt b/tests/logging/CMakeLists.txt index aba772e..8240e05 100644 --- a/tests/logging/CMakeLists.txt +++ b/tests/logging/CMakeLists.txt @@ -76,7 +76,7 @@ target_compile_features(MsRdpEx_DvcFactoryLifetimeTest PRIVATE cxx_std_17) target_compile_options(MsRdpEx_DvcFactoryLifetimeTest PRIVATE /W4 /EHsc) target_link_libraries(MsRdpEx_DvcFactoryLifetimeTest PRIVATE uuid.lib) add_dependencies(MsRdpEx_DvcFactoryLifetimeTest MsRdpEx_GatewayShutdownTestDll) -foreach(scenario removed replacement addref-reentrant release-reentrant manager-shutdown allocation-failure query-failure) +foreach(scenario removed replacement addref-reentrant release-reentrant destructor-reentrant manager-shutdown allocation-failure query-failure) add_test(NAME logging.dvc-factory.${scenario} COMMAND MsRdpEx_DvcFactoryLifetimeTest "$" ${scenario}) diff --git a/tests/logging/DvcFactoryLifetimeTest.cpp b/tests/logging/DvcFactoryLifetimeTest.cpp index 1c8cc86..ece19d1 100644 --- a/tests/logging/DvcFactoryLifetimeTest.cpp +++ b/tests/logging/DvcFactoryLifetimeTest.cpp @@ -46,7 +46,10 @@ class CountedPlugin : public IWTSPlugin { if (clearOnRelease) { clearOnRelease = false; - instance->SetWTSPluginObject(NULL); + releaseSetterHr = instance->SetWTSPluginObject(replacement); + if (replacement) + releaseClearHr = instance->SetWTSPluginObject(NULL); + releaseGetterHr = instance->GetWTSPluginObject(&borrowedAfterRelease); } return --refs; } @@ -62,6 +65,10 @@ class CountedPlugin : public IWTSPlugin bool clearOnAddRef = false; bool clearOnRelease = false; bool failQuery = false; + HRESULT releaseSetterHr = E_FAIL; + HRESULT releaseClearHr = E_FAIL; + HRESULT releaseGetterHr = E_FAIL; + void* borrowedAfterRelease = reinterpret_cast(1); ULONG refsAfterReplacement = 0; HANDLE queryEntered = NULL; HANDLE queryResume = NULL; @@ -190,8 +197,27 @@ int wmain(int argc, wchar_t** argv) "Reentrant Release did not clear the replacement plugin"); Check(first.refs == 1 && second.refs == 1, "Reentrant Release did not balance the plugin references"); + Check(first.releaseSetterHr == S_OK && first.releaseGetterHr == S_OK + && !first.borrowedAfterRelease, + "Replacement callback could not access the cleared plugin slot"); + Check(unregisterInstance(instance), "Could not remove plugin instance"); + Check(instance->Release() == 0, "Instance remained alive after removal"); + Check(first.Release() == 0 && second.Release() == 0, + "Test plugin references not balanced"); + } else if (scenario == L"destructor-reentrant") { + CountedPlugin second; + second.AddRef(); + first.instance = instance; + first.replacement = &second; + first.clearOnRelease = true; Check(unregisterInstance(instance), "Could not remove plugin instance"); Check(instance->Release() == 0, "Instance remained alive after removal"); + Check(first.releaseSetterHr == E_UNEXPECTED && first.releaseClearHr == E_UNEXPECTED + && first.releaseGetterHr == S_OK && !first.borrowedAfterRelease, + "Destructor callback accessed or replaced the detached plugin slot"); + Check(first.refs == 1 && second.refs == 2, + "Destructor callback consumed or leaked a plugin reference"); + Check(second.Release() == 1, "Caller could not release rejected replacement"); Check(first.Release() == 0 && second.Release() == 0, "Test plugin references not balanced"); } else if (scenario == L"manager-shutdown") {