#include "ScriptHookPlugin.h" #include "Logging.h" #include "RegistryConfig.h" #include "../Shared/TriggerProtocol.h" #include #include #include namespace { std::wstring GetLocalClientHostname() { // PhysicalDnsHostname liefert den Hostteil ohne DNS-Domaene und ist nicht von // Variablen in der Remotesitzung abhaengig. Auf normalen Workstations entspricht // das dem lokalen Rechnernamen. GetComputerNameW dient als Rueckfall. wchar_t hostname[256]{}; DWORD hostnameChars = static_cast(sizeof(hostname) / sizeof(hostname[0])); if (GetComputerNameExW(ComputerNamePhysicalDnsHostname, hostname, &hostnameChars)) { return std::wstring(hostname, hostnameChars); } wchar_t fallback[MAX_COMPUTERNAME_LENGTH + 1]{}; DWORD fallbackChars = static_cast(sizeof(fallback) / sizeof(fallback[0])); if (GetComputerNameW(fallback, &fallbackChars)) { return std::wstring(fallback, fallbackChars); } return {}; } std::string WideToUtf8(const std::wstring& value) { if (value.empty()) { return {}; } const int required = WideCharToMultiByte( CP_UTF8, WC_ERR_INVALID_CHARS, value.data(), static_cast(value.size()), nullptr, 0, nullptr, nullptr); if (required <= 0) { return {}; } std::string result(static_cast(required), '\0'); const int written = WideCharToMultiByte( CP_UTF8, WC_ERR_INVALID_CHARS, value.data(), static_cast(value.size()), result.data(), required, nullptr, nullptr); if (written != required) { return {}; } return result; } class TriggerChannelCallback final : public IWTSVirtualChannelCallback { public: explicit TriggerChannelCallback(IWTSVirtualChannel* channel) : channel_(channel) { if (channel_ != nullptr) { channel_->AddRef(); } } 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(this); AddRef(); return S_OK; } return E_NOINTERFACE; } ULONG STDMETHODCALLTYPE AddRef() override { return static_cast(InterlockedIncrement(&refCount_)); } ULONG STDMETHODCALLTYPE Release() override { ULONG remaining = static_cast(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::kGetClientHostnameMessageLength && std::memcmp(buffer, hook::protocol::kGetClientHostnameMessage, hook::protocol::kGetClientHostnameMessageLength) == 0) { const std::wstring hostname = GetLocalClientHostname(); const std::string hostnameUtf8 = WideToUtf8(hostname); if (channel_ == nullptr || hostnameUtf8.empty()) { hook::Log(config.enableLogging, L"DVC-Hostname: Lokaler Client-Hostname konnte nicht ermittelt oder beantwortet werden."); return S_OK; } std::string response; response.reserve(hook::protocol::kClientHostnameResponsePrefixLength + hostnameUtf8.size()); response.append(hook::protocol::kClientHostnameResponsePrefix, hook::protocol::kClientHostnameResponsePrefixLength); response.append(hostnameUtf8); const HRESULT hr = channel_->Write( static_cast(response.size()), reinterpret_cast(response.data()), nullptr); if (FAILED(hr)) { hook::Log(config.enableLogging, L"DVC-Hostname: Antwort konnte nicht geschrieben werden. HRESULT=" + std::to_wstring(static_cast(hr)) + L"."); } else { hook::Log(config.enableLogging, L"DVC-Hostname: Client-Hostname an Terminalserver beantwortet: " + hostname); } 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 Sicherheitsgruenden ignoriert."); return S_OK; } hook::Log(config.enableLogging, L"DVC-Trigger: START vom Terminalserver empfangen."); bool startSucceeded = false; if (!config.enabled) { hook::Log(config.enableLogging, L"DVC-Trigger: Plugin-Funktion ist deaktiviert; Start ignoriert."); } else if (!config.triggerEnabled) { hook::Log(config.enableLogging, L"DVC-Trigger: Server-Trigger ist in der Client-Konfiguration deaktiviert."); } else { // Der Terminalserver liefert absichtlich weder Pfad noch Argumente. // Beides wird ausschliesslich aus HKCU auf dem lokalen Client gelesen. hook::ProcessLauncher launcher; startSucceeded = launcher.Start( config.triggerProgram, config.enableLogging, L"Server-Trigger-Programm"); } // Der Server wartet auf diese Quittung, bevor er den kurzlebigen DVC // schliesst. Das verhindert die beobachtete Race Condition, bei der // WTSVirtualChannelWrite erfolgreich war, der Kanal aber geschlossen // wurde, bevor OnDataReceived die START-Nachricht erhalten hatte. if (channel_ != nullptr) { const char* response = startSucceeded ? hook::protocol::kStartOkResponse : hook::protocol::kStartFailedResponse; const std::size_t responseLength = startSucceeded ? hook::protocol::kStartOkResponseLength : hook::protocol::kStartFailedResponseLength; const HRESULT hr = channel_->Write( static_cast(responseLength), reinterpret_cast(const_cast(response)), nullptr); if (FAILED(hr)) { hook::Log(config.enableLogging, L"DVC-Trigger: START-Quittung konnte nicht geschrieben werden. HRESULT=" + std::to_wstring(static_cast(hr)) + L"."); } else { hook::Log(config.enableLogging, startSucceeded ? L"DVC-Trigger: START_OK an Terminalserver gesendet." : L"DVC-Trigger: START_FAILED an Terminalserver gesendet."); } } 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() { if (channel_ != nullptr) { channel_->Release(); channel_ = nullptr; } } LONG refCount_ = 1; IWTSVirtualChannel* channel_ = nullptr; }; } 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(this); } else if (riid == __uuidof(IWTSListenerCallback)) { *ppvObject = static_cast(this); } else { return E_NOINTERFACE; } AddRef(); return S_OK; } ULONG STDMETHODCALLTYPE ScriptHookPlugin::AddRef() { return static_cast(InterlockedIncrement(&refCount_)); } ULONG STDMETHODCALLTYPE ScriptHookPlugin::Release() { ULONG remaining = static_cast(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(hook::protocol::kChannelName), 0, static_cast(this), &listener_); if (FAILED(hr)) { hook::Log(config.enableLogging, L"DVC-Listener konnte nicht erstellt werden. HRESULT=" + std::to_wstring(static_cast(hr)) + L". Der automatische Start bei RDP-Verbindung bleibt trotzdem aktiv."); ReleaseDvcObjects(); return S_OK; } hook::Log(config.enableLogging, L"DVC-Listener fuer '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; } 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(channel); if (channelCallback == nullptr) { hook::Log(config.enableLogging, L"DVC-Trigger: Callback konnte nicht angelegt werden."); return E_OUTOFMEMORY; } *callback = static_cast(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; } }