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; + } + } + } +}