diff --git a/dll/MsRdpEx.cpp b/dll/MsRdpEx.cpp index ac668b8..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, (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); + instance->Release(); return hr; } diff --git a/dll/RdpDvcClient.cpp b/dll/RdpDvcClient.cpp index be17ed2..83ae637 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,24 +354,35 @@ 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]; MsRdpEx_GuidBinToStr((GUID*)&riid, iid, 0); if (riid == IID_IWTSPlugin) { - IUnknown* wtsPlugin = NULL; - IMsRdpExInstance* rdpInstance = (IMsRdpExInstance*)m_instance; - rdpInstance->GetWTSPluginObject((void**)&wtsPlugin); + MsRdpEx_WTSPluginReference* wtsPlugin = NULL; + 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); + hr = wtsPlugin->Get()->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 1414a4b..e1c4aa6 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: @@ -32,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()); @@ -50,9 +78,8 @@ class CMsRdpExInstance : public IMsRdpExInstance m_pMsRdpExtendedSettings->Release(); } - if (m_WTSPlugin) { - m_WTSPlugin->Release(); - } + if (plugin) + plugin->Release(); } // IUnknown interface @@ -535,20 +562,45 @@ 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) { - IUnknown* previousPlugin = m_WTSPlugin; - m_WTSPlugin = (IUnknown*)pvObject; - if (previousPlugin) { - previousPlugin->Release(); + MsRdpEx_WTSPluginReference* replacement = NULL; + if (pvObject) { + replacement = new (std::nothrow) MsRdpEx_WTSPluginReference((IUnknown*)pvObject); + if (!replacement) + return E_OUTOFMEMORY; + } + + 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); + if (previous) + previous->Release(); return S_OK; } + MsRdpEx_WTSPluginReference* AcquireWTSPluginObject() + { + AcquireSRWLockShared(&m_WTSPluginLock); + MsRdpEx_WTSPluginReference* plugin = m_WTSPlugin; + if (plugin) + plugin->AddRef(); + ReleaseSRWLockShared(&m_WTSPluginLock); + return plugin; + } + public: GUID m_sessionId; ULONG m_refCount; @@ -562,7 +614,9 @@ 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; + bool m_WTSPluginClosing = false; LONG m_GdiReconnectPending = 0; LONG m_GdiReconnectAttempts = 0; LONG m_HardwareCaptureFrameReceived = 0; @@ -772,6 +826,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) { @@ -1109,6 +1164,68 @@ CMsRdpExInstance* MsRdpEx_InstanceManager_FindBySessionId(GUID* sessionId) return found ? obj : NULL; } +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; + *plugin = NULL; + if (!sessionId) + return E_INVALIDARG; + + AcquireSRWLockShared(&g_InstanceManagerLock); + MsRdpEx_InstanceManager* ctx = g_InstanceManager; + if (!ctx) { + ReleaseSRWLockShared(&g_InstanceManagerLock); + 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); + ReleaseSRWLockShared(&g_InstanceManagerLock); + return found ? S_OK : REGDB_E_CLASSNOTREG; +} + CMsRdpExtendedSettings* MsRdpEx_FindExtendedSettingsBySessionId(GUID* sessionId) { CMsRdpExInstance* instance = NULL; @@ -1194,24 +1311,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 9bb7a9c..c2889e1 100644 --- a/dll/RdpInstanceInternal.h +++ b/dll/RdpInstanceInternal.h @@ -3,8 +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); +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/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 9811ef0..cd1631d 100644 --- a/tests/logging/CMakeLists.txt +++ b/tests/logging/CMakeLists.txt @@ -70,6 +70,19 @@ foreach(scenario replacement same-pointer clearing destruction) 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 addref-reentrant release-reentrant destructor-reentrant manager-shutdown allocation-failure query-failure) + add_test(NAME logging.dvc-factory.${scenario} + COMMAND MsRdpEx_DvcFactoryLifetimeTest + "$" ${scenario}) + set_tests_properties(logging.dvc-factory.${scenario} PROPERTIES TIMEOUT 45) +endforeach() + add_executable(MsRdpEx_DetachedInstanceTest DetachedInstanceTest.cpp) target_include_directories(MsRdpEx_DetachedInstanceTest PRIVATE "${PROJECT_SOURCE_DIR}/dll") target_compile_features(MsRdpEx_DetachedInstanceTest PRIVATE cxx_std_17) diff --git a/tests/logging/DvcFactoryLifetimeTest.cpp b/tests/logging/DvcFactoryLifetimeTest.cpp new file mode 100644 index 0000000..ece19d1 --- /dev/null +++ b/tests/logging/DvcFactoryLifetimeTest.cpp @@ -0,0 +1,298 @@ +#include + +#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 (queryEntered) { + 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; + refsAfterReplacement = refs; + } + *object = static_cast(this); + AddRef(); + return S_OK; + } + + ULONG STDMETHODCALLTYPE AddRef() override + { + if (clearOnAddRef) { + clearOnAddRef = false; + instance->SetWTSPluginObject(NULL); + } + return ++refs; + } + ULONG STDMETHODCALLTYPE Release() override + { + if (clearOnRelease) { + clearOnRelease = false; + releaseSetterHr = instance->SetWTSPluginObject(replacement); + if (replacement) + releaseClearHr = instance->SetWTSPluginObject(NULL); + releaseGetterHr = instance->GetWTSPluginObject(&borrowedAfterRelease); + } + 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; + 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; +}; + +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")); + using FailAllocation = void(*)(); + auto failAllocation = reinterpret_cast( + GetProcAddress(module, "FailNextPluginHolderAllocation")); + Check(create && registerInstance && unregisterInstance && createFactory && failAllocation, + "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 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(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") { + 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 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"); + } + + 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 ffb18be..f79636e 100644 --- a/tests/logging/GatewayShutdownFixture.def +++ b/tests/logging/GatewayShutdownFixture.def @@ -1,6 +1,10 @@ EXPORTS PrepareGatewayShutdown CreatePluginReferenceInstance + RegisterPluginReferenceInstance + UnregisterPluginReferenceInstance + CreatePluginReferenceFactory + FailNextPluginHolderAllocation TryRegisterDetachedInstance TryRemoveDetachedInstance RegisterDetachedOutputWindowClass diff --git a/tests/logging/PluginReferenceFixture.cpp b/tests/logging/PluginReferenceFixture.cpp index db7f3c8..cf756d1 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; +} + ATOM WINAPI Hook_RegisterClassExW(WNDCLASSEXW* wndClass); extern "C" IMsRdpExInstance* CreatePluginReferenceInstance() @@ -7,6 +25,28 @@ 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 MsRdpEx_DllGetClassObject(sessionId, IID_IClassFactory, (void**)factory); +} + extern "C" bool TryRegisterDetachedInstance(IMsRdpExInstance* instance) { return MsRdpEx_InstanceManager_Add((CMsRdpExInstance*)instance); diff --git a/tests/logging/README.md b/tests/logging/README.md index 1ff14b2..ba84f66 100644 --- a/tests/logging/README.md +++ b/tests/logging/README.md @@ -83,6 +83,18 @@ releases it when replaced (including by another reference to the same object), cleared, or destroyed. The plugin getter returns a borrowed pointer. The fixture export is not included in the production DLL. +## DVC class factory lifetime + +`logging.dvc-factory.*` obtains a class factory through the production +`DllGetClassObject` route without connecting RDP. The cases cover use after +session removal, plugin replacement and clearing, failed plugin queries, +reentrant plugin `AddRef`/`Release`, and manager shutdown during an in-flight +query. A test-only allocation failure checks that the native setter leaves the +caller's reference untouched on failure, including same-pointer replacement. +During instance destruction, reentrant getters observe an empty slot and +setters reject new registrations without consuming their input reference. +The fixture-only exports and allocator override are not included in the DLL. + ## Detached output-window lifetime `logging.detached-instance.lifetime` registers an output window class through