From 891826c5ab3b6275bcabff329734c7699d2cfbc8 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Marc-Andr=C3=A9=20Moreau?= Date: Thu, 24 Sep 2026 06:47:48 -0400 Subject: [PATCH] Release WTS plugin reference on registration failure Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> --- dotnet/Devolutions.MsRdpEx/RdpInstance.cs | 17 +++- dotnet/MsRdpEx_Test/RdpInstancePluginTests.cs | 84 +++++++++++++++++++ 2 files changed, 100 insertions(+), 1 deletion(-) create mode 100644 dotnet/MsRdpEx_Test/RdpInstancePluginTests.cs diff --git a/dotnet/Devolutions.MsRdpEx/RdpInstance.cs b/dotnet/Devolutions.MsRdpEx/RdpInstance.cs index 06a43be..e3b00c3 100644 --- a/dotnet/Devolutions.MsRdpEx/RdpInstance.cs +++ b/dotnet/Devolutions.MsRdpEx/RdpInstance.cs @@ -53,7 +53,22 @@ public bool GetShadowBitmap(ref IntPtr phDC, ref IntPtr phBitmap, ref IntPtr pBi public object WTSPlugin { - set { iface.SetWTSPluginObject(Marshal.GetIUnknownForObject(value)); } + set + { + IntPtr plugin = Marshal.GetIUnknownForObject(value); + // The native setter takes ownership only when the call succeeds. + bool transferred = false; + try + { + iface.SetWTSPluginObject(plugin); + transferred = true; + } + finally + { + if (!transferred) + Marshal.Release(plugin); + } + } } } } diff --git a/dotnet/MsRdpEx_Test/RdpInstancePluginTests.cs b/dotnet/MsRdpEx_Test/RdpInstancePluginTests.cs new file mode 100644 index 0000000..17af8c6 --- /dev/null +++ b/dotnet/MsRdpEx_Test/RdpInstancePluginTests.cs @@ -0,0 +1,84 @@ +using System.Reflection; +using System.Runtime.InteropServices; + +namespace MsRdpEx.Tests +{ + public class RdpInstancePluginTests + { + [Fact] + public void WTSPlugin_WhenRegistrationFails_ReleasesOwnedReference() + { + var plugin = new object(); + IntPtr reference = Marshal.GetIUnknownForObject(plugin); + int initialCount = GetReferenceCount(reference); + try + { + var failure = new COMException("Plugin registration failed."); + int calls = 0; + var instance = DispatchProxy.Create(); + ((PluginInstanceProxy)instance).Register = pointer => + { + Assert.Equal(reference, pointer); + calls++; + throw failure; + }; + + Assert.Same(failure, Assert.Throws(() => new RdpInstance(instance).WTSPlugin = plugin)); + Assert.Equal(1, calls); + Assert.Equal(initialCount, GetReferenceCount(reference)); + } + finally + { + if (GetReferenceCount(reference) > initialCount) + Marshal.Release(reference); + Marshal.Release(reference); + } + } + + [Fact] + public void WTSPlugin_WhenRegistrationSucceeds_TransfersOwnedReference() + { + var plugin = new object(); + IntPtr reference = Marshal.GetIUnknownForObject(plugin); + int initialCount = GetReferenceCount(reference); + IntPtr transferred = IntPtr.Zero; + try + { + var instance = DispatchProxy.Create(); + ((PluginInstanceProxy)instance).Register = pointer => transferred = pointer; + + new RdpInstance(instance).WTSPlugin = plugin; + + Assert.Equal(reference, transferred); + Assert.Equal(initialCount + 1, GetReferenceCount(reference)); + } + finally + { + if (transferred != IntPtr.Zero && GetReferenceCount(reference) > initialCount) + Marshal.Release(transferred); + Marshal.Release(reference); + } + } + + private static int GetReferenceCount(IntPtr pointer) + { + int count = Marshal.AddRef(pointer); + Marshal.Release(pointer); + return count; + } + + public class PluginInstanceProxy : DispatchProxy + { + public Action? Register { get; set; } + + protected override object? Invoke(MethodInfo? targetMethod, object?[]? args) + { + if (targetMethod?.Name != nameof(IMsRdpExInstance.SetWTSPluginObject)) + throw new NotSupportedException(targetMethod?.Name); + + Register!((IntPtr)args![0]!); + return null; + } + } + } +}