Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
6 changes: 4 additions & 2 deletions dll/MsRdpEx.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -9,6 +9,7 @@
#include <MsRdpEx/Detours.h>

#include "RdpDvcClient.h"
#include "RdpInstanceInternal.h"

#include <stdarg.h>
#include <comutil.h>
Expand Down Expand Up @@ -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;
}

Expand Down
84 changes: 46 additions & 38 deletions dll/RdpDvcClient.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,10 @@

#include <MsRdpEx/RdpInstance.h>

#include <new>

#include "RdpInstanceInternal.h"

//
// CRdpDvcClient class
//
Expand Down Expand Up @@ -298,42 +302,35 @@ 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);

MsRdpEx_LogPrint(DEBUG, "CDvcPluginClassFactory::QueryInterface(%s)", iid);

if (riid == IID_IUnknown) {
*ppvObject = (LPVOID)((IUnknown*)this);
InterlockedIncrement(&m_refCount);
return S_OK;
*ppvObject = static_cast<IUnknown*>(this);
}
if (riid == IID_IClassFactory) {
*ppvObject = (LPVOID)((IClassFactory*)this);
InterlockedIncrement(&m_refCount);
return S_OK;
else if (riid == IID_IClassFactory) {
*ppvObject = static_cast<IClassFactory*>(this);
}

return hr;
if (!*ppvObject)
return E_NOINTERFACE;
AddRef();
return S_OK;
}

ULONG STDMETHODCALLTYPE AddRef()
Expand All @@ -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();
}
}
}

Expand All @@ -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;
}
2 changes: 1 addition & 1 deletion dll/RdpDvcClient.h
Original file line number Diff line number Diff line change
Expand Up @@ -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 */
146 changes: 135 additions & 11 deletions dll/RdpInstance.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand All @@ -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());

Expand All @@ -50,9 +78,8 @@ class CMsRdpExInstance : public IMsRdpExInstance
m_pMsRdpExtendedSettings->Release();
}

if (m_WTSPlugin) {
m_WTSPlugin->Release();
}
if (plugin)
plugin->Release();
}

// IUnknown interface
Expand Down Expand Up @@ -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;
Expand All @@ -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;
Expand Down Expand Up @@ -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)
{
Expand Down Expand Up @@ -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;
Expand Down Expand Up @@ -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);
}
Loading