Files
Mstsc-Script-Hook/Plugin/ScriptHookPlugin.cpp
T

286 lines
8.1 KiB
C++

#include "ScriptHookPlugin.h"
#include "Logging.h"
#include "RegistryConfig.h"
#include "../Shared/TriggerProtocol.h"
#include <cstring>
#include <new>
#include <string>
namespace
{
class TriggerChannelCallback final : public IWTSVirtualChannelCallback
{
public:
TriggerChannelCallback() = default;
HRESULT STDMETHODCALLTYPE QueryInterface(REFIID riid, void** ppvObject) override
{
if (ppvObject == nullptr)
{
return E_POINTER;
}
*ppvObject = nullptr;
if (riid == __uuidof(IUnknown) || riid == __uuidof(IWTSVirtualChannelCallback))
{
*ppvObject = static_cast<IWTSVirtualChannelCallback*>(this);
AddRef();
return S_OK;
}
return E_NOINTERFACE;
}
ULONG STDMETHODCALLTYPE AddRef() override
{
return static_cast<ULONG>(InterlockedIncrement(&refCount_));
}
ULONG STDMETHODCALLTYPE Release() override
{
ULONG remaining = static_cast<ULONG>(InterlockedDecrement(&refCount_));
if (remaining == 0)
{
delete this;
}
return remaining;
}
HRESULT STDMETHODCALLTYPE OnDataReceived(ULONG cbSize, BYTE* buffer) override
{
const hook::HookConfig config = hook::LoadConfig();
if (buffer == nullptr || cbSize == 0)
{
hook::Log(config.enableLogging, L"DVC-Trigger: Leere Nachricht empfangen; ignoriert.");
return S_OK;
}
if (cbSize != hook::protocol::kStartMessageLength ||
std::memcmp(buffer, hook::protocol::kStartMessage, hook::protocol::kStartMessageLength) != 0)
{
hook::Log(config.enableLogging,
L"DVC-Trigger: Unbekannte Nachricht empfangen; aus Sicherheitsgründen ignoriert.");
return S_OK;
}
hook::Log(config.enableLogging, L"DVC-Trigger: START vom Terminalserver empfangen.");
if (!config.enabled)
{
hook::Log(config.enableLogging, L"DVC-Trigger: Plugin-Funktion ist deaktiviert; Start ignoriert.");
return S_OK;
}
if (!config.triggerEnabled)
{
hook::Log(config.enableLogging, L"DVC-Trigger: Server-Trigger ist in der Client-Konfiguration deaktiviert.");
return S_OK;
}
// Der Terminalserver liefert absichtlich weder Pfad noch Argumente.
// Beides wird ausschließlich aus HKCU auf dem lokalen Client gelesen.
hook::ProcessLauncher launcher;
launcher.Start(config.triggerProgram, config.enableLogging, L"Server-Trigger-Programm");
return S_OK;
}
HRESULT STDMETHODCALLTYPE OnClose() override
{
const hook::HookConfig config = hook::LoadConfig();
hook::Log(config.enableLogging, L"DVC-Trigger: Kanal geschlossen.");
return S_OK;
}
private:
~TriggerChannelCallback() = default;
LONG refCount_ = 1;
};
}
ScriptHookPlugin::ScriptHookPlugin() = default;
ScriptHookPlugin::~ScriptHookPlugin()
{
ReleaseDvcObjects();
}
HRESULT STDMETHODCALLTYPE ScriptHookPlugin::QueryInterface(REFIID riid, void** ppvObject)
{
if (ppvObject == nullptr)
{
return E_POINTER;
}
*ppvObject = nullptr;
if (riid == __uuidof(IUnknown) || riid == __uuidof(IWTSPlugin))
{
*ppvObject = static_cast<IWTSPlugin*>(this);
}
else if (riid == __uuidof(IWTSListenerCallback))
{
*ppvObject = static_cast<IWTSListenerCallback*>(this);
}
else
{
return E_NOINTERFACE;
}
AddRef();
return S_OK;
}
ULONG STDMETHODCALLTYPE ScriptHookPlugin::AddRef()
{
return static_cast<ULONG>(InterlockedIncrement(&refCount_));
}
ULONG STDMETHODCALLTYPE ScriptHookPlugin::Release()
{
ULONG remaining = static_cast<ULONG>(InterlockedDecrement(&refCount_));
if (remaining == 0)
{
delete this;
}
return remaining;
}
HRESULT STDMETHODCALLTYPE ScriptHookPlugin::Initialize(IWTSVirtualChannelManager* channelManager)
{
const hook::HookConfig config = hook::LoadConfig();
hook::Log(config.enableLogging, L"IWTSPlugin::Initialize aufgerufen.");
if (channelManager == nullptr)
{
hook::Log(config.enableLogging, L"DVC-Initialisierung fehlgeschlagen: channelManager ist NULL.");
return E_POINTER;
}
ReleaseDvcObjects();
channelManager->AddRef();
channelManager_ = channelManager;
HRESULT hr = channelManager_->CreateListener(
const_cast<LPSTR>(hook::protocol::kChannelName),
0,
static_cast<IWTSListenerCallback*>(this),
&listener_);
if (FAILED(hr))
{
hook::Log(config.enableLogging,
L"DVC-Listener konnte nicht erstellt werden. HRESULT=" + std::to_wstring(static_cast<unsigned long>(hr)) +
L". Der automatische Start bei RDP-Verbindung bleibt trotzdem aktiv.");
ReleaseDvcObjects();
// Der DVC-Trigger ist eine optionale Zusatzfunktion. Ein nicht verfügbarer
// DVC darf den bestehenden Connected()/Disconnected()-Hook nicht deaktivieren.
return S_OK;
}
hook::Log(config.enableLogging,
L"DVC-Listener für 'plandent::mstsc-script-hook' wurde registriert.");
return S_OK;
}
HRESULT STDMETHODCALLTYPE ScriptHookPlugin::Connected()
{
const hook::HookConfig config = hook::LoadConfig();
hook::Log(config.enableLogging, L"RDP-Verbindung hergestellt (IWTSPlugin::Connected).");
if (!config.enabled || !config.startOnConnect)
{
hook::Log(config.enableLogging, L"Automatischer Start bei RDP-Verbindung ist deaktiviert.");
return S_OK;
}
// Nur Registry lesen + CreateProcess aufrufen. Es wird niemals auf das Programm gewartet.
connectLauncher_.Start(config.connectProgram, config.enableLogging, L"RDP-Verbindungsprogramm");
return S_OK;
}
HRESULT STDMETHODCALLTYPE ScriptHookPlugin::Disconnected(DWORD disconnectCode)
{
const hook::HookConfig config = hook::LoadConfig();
hook::Log(config.enableLogging,
L"RDP-Verbindung getrennt (IWTSPlugin::Disconnected), Code=" + std::to_wstring(disconnectCode) + L".");
if (config.stopOnDisconnect)
{
connectLauncher_.Stop(config.enableLogging, L"RDP-Verbindungsprogramm");
}
return S_OK;
}
HRESULT STDMETHODCALLTYPE ScriptHookPlugin::Terminated()
{
const hook::HookConfig config = hook::LoadConfig();
hook::Log(config.enableLogging, L"IWTSPlugin::Terminated aufgerufen.");
if (config.stopOnDisconnect)
{
connectLauncher_.Stop(config.enableLogging, L"RDP-Verbindungsprogramm");
}
else
{
connectLauncher_.Detach(config.enableLogging, L"RDP-Verbindungsprogramm");
}
ReleaseDvcObjects();
return S_OK;
}
HRESULT STDMETHODCALLTYPE ScriptHookPlugin::OnNewChannelConnection(
IWTSVirtualChannel* channel,
BSTR data,
BOOL* accept,
IWTSVirtualChannelCallback** callback)
{
UNREFERENCED_PARAMETER(data);
const hook::HookConfig config = hook::LoadConfig();
if (accept == nullptr || callback == nullptr)
{
return E_POINTER;
}
*accept = FALSE;
*callback = nullptr;
if (channel == nullptr)
{
return E_POINTER;
}
auto* channelCallback = new (std::nothrow) TriggerChannelCallback();
if (channelCallback == nullptr)
{
hook::Log(config.enableLogging, L"DVC-Trigger: Callback konnte nicht angelegt werden.");
return E_OUTOFMEMORY;
}
*callback = static_cast<IWTSVirtualChannelCallback*>(channelCallback);
*accept = TRUE;
hook::Log(config.enableLogging, L"DVC-Trigger: Verbindung vom Terminalserver akzeptiert.");
return S_OK;
}
void ScriptHookPlugin::ReleaseDvcObjects()
{
if (listener_ != nullptr)
{
listener_->Release();
listener_ = nullptr;
}
if (channelManager_ != nullptr)
{
channelManager_->Release();
channelManager_ = nullptr;
}
}