diff --git a/src/Windows/Avalonia.Win32.Automation/Avalonia.Win32.Automation.csproj b/src/Windows/Avalonia.Win32.Automation/Avalonia.Win32.Automation.csproj index c56131c495..0d30f8e74b 100644 --- a/src/Windows/Avalonia.Win32.Automation/Avalonia.Win32.Automation.csproj +++ b/src/Windows/Avalonia.Win32.Automation/Avalonia.Win32.Automation.csproj @@ -15,5 +15,6 @@ + diff --git a/src/Windows/Avalonia.Win32.Automation/Marshalling/SafeArrayMarshaller.cs b/src/Windows/Avalonia.Win32.Automation/Marshalling/SafeArrayMarshaller.cs index ef3af4a3b2..a82ad213c4 100644 --- a/src/Windows/Avalonia.Win32.Automation/Marshalling/SafeArrayMarshaller.cs +++ b/src/Windows/Avalonia.Win32.Automation/Marshalling/SafeArrayMarshaller.cs @@ -10,7 +10,7 @@ internal static class SafeArrayMarshaller where T : notnull { public static SafeArrayRef ConvertToUnmanaged(T[]? managed) => managed is null ? new SafeArrayRef() - : SafeArrayRef.TryCreate(managed, out var result, out _) ? result.Value + : SafeArrayRef.TryCreate(managed, out var result, out _) ? result.Value : throw new NotImplementedException($"SafeArray marshalling for '{managed?.GetType().Name}' is not implemented."); public static T[]? ConvertToManaged(SafeArrayRef unmanaged) => SafeArrayRef.ToArray(unmanaged); diff --git a/src/Windows/Avalonia.Win32.Automation/Marshalling/SafeArrayRef.cs b/src/Windows/Avalonia.Win32.Automation/Marshalling/SafeArrayRef.cs index f27c67e36e..b4c2e53710 100644 --- a/src/Windows/Avalonia.Win32.Automation/Marshalling/SafeArrayRef.cs +++ b/src/Windows/Avalonia.Win32.Automation/Marshalling/SafeArrayRef.cs @@ -8,6 +8,7 @@ using System.Diagnostics.CodeAnalysis; using System.Linq; using System.Runtime.CompilerServices; using System.Runtime.InteropServices; +using System.Runtime.InteropServices.Marshalling; // ReSharper disable InconsistentNaming namespace Avalonia.Win32.Automation.Marshalling; @@ -160,6 +161,25 @@ internal unsafe partial struct SafeArrayRef }); } + /// + /// Creates a SAFEARRAY from a typed array. + /// + /// + /// COM interfaces have to come through here rather than through the non-generic overload: the + /// element type is what tells us to wrap each item, and it isn't recoverable from the values. + /// + public static bool TryCreate(T[]? managed, [NotNullWhen(true)] out SafeArrayRef? safearray, out VarEnum varEnum) + { + if (managed is not null && typeof(T).IsInterface) + { + safearray = CreateFromComObjects(managed); + varEnum = VarEnum.VT_UNKNOWN; + return true; + } + + return TryCreate((IEnumerable?)managed, out safearray, out varEnum); + } + public static bool TryCreate(IEnumerable? managed, [NotNullWhen(true)] out SafeArrayRef? safearray, out VarEnum varEnum) { safearray = default; @@ -181,38 +201,6 @@ internal unsafe partial struct SafeArrayRef return CreateFromSpan(collectionSpan, varEnum); } - static SafeArrayRef CreateFromSpan(ReadOnlySpan span, VarEnum varEnum) - { - var bound = new SAFEARRAYBOUND { cElements = (uint)span.Length, lLbound = 0 }; - var safearray = SafeArrayCreate(varEnum, 1, bound); - if (span.Length == 0) - { - return new SafeArrayRef - { - _ptr = safearray - }; - } - - var lockResult = SafeArrayLock(safearray); - if (lockResult != 0) throw new Win32Exception(lockResult); - - try - { - // We assume it has the same length. - var output = new Span(safearray->pvData, (int)safearray->rgsabound[0].cElements); - span.CopyTo(output); - } - finally - { - SafeArrayUnlock(safearray); - } - - return new SafeArrayRef - { - _ptr = safearray - }; - } - static SafeArrayRef CreateFromStrings(IReadOnlyList strings, VarEnum varEnum) { Debug.Assert(varEnum == VarEnum.VT_BSTR); // other types not supported yet @@ -251,28 +239,6 @@ internal unsafe partial struct SafeArrayRef } } - static SafeArrayRef CreateFromObjects(IReadOnlyList objects, VarEnum varEnum) - { - Debug.Assert(varEnum == VarEnum.VT_UNKNOWN); // other types not supported yet - var pointers = ArrayPool.Shared.Rent(objects.Count); - try - { - for (int i = 0; i < objects.Count; i++) - { - if (ComWrappers.TryGetComInstance(objects[i], out var pointer)) - { - pointers[i] = pointer; - } - } - - return CreateFromSpan(pointers, varEnum); - } - finally - { - ArrayPool.Shared.Return(pointers); - } - } - safearray = managed switch { IReadOnlyCollection ints => CreateFromCollection(ints, varEnum = VarEnum.VT_I1), @@ -295,14 +261,63 @@ internal unsafe partial struct SafeArrayRef IReadOnlyList strings => CreateFromStrings(strings, varEnum = VarEnum.VT_BSTR), - IReadOnlyList objects => CreateFromObjects(objects, varEnum = VarEnum.VT_UNKNOWN), - _ => null }; return safearray is not null; } + private static SafeArrayRef CreateFromSpan(ReadOnlySpan span, VarEnum varEnum) + { + var bound = new SAFEARRAYBOUND { cElements = (uint)span.Length, lLbound = 0 }; + var safearray = SafeArrayCreate(varEnum, 1, bound); + if (span.Length == 0) + { + return new SafeArrayRef + { + _ptr = safearray + }; + } + + var lockResult = SafeArrayLock(safearray); + if (lockResult != 0) throw new Win32Exception(lockResult); + + try + { + // We assume it has the same length. + var output = new Span(safearray->pvData, (int)safearray->rgsabound[0].cElements); + span.CopyTo(output); + } + finally + { + SafeArrayUnlock(safearray); + } + + return new SafeArrayRef + { + _ptr = safearray + }; + } + + private static SafeArrayRef CreateFromComObjects(T[] objects) + { + var pointers = ArrayPool.Shared.Rent(objects.Length); + + try + { + for (var i = 0; i < objects.Length; i++) + { + pointers[i] = (IntPtr)ComInterfaceMarshaller.ConvertToUnmanaged(objects[i]); + } + + return CreateFromSpan(pointers.AsSpan(0, objects.Length), VarEnum.VT_UNKNOWN); + } + finally + { + ArrayPool.Shared.Return(pointers); + } + } + [LibraryImport("oleaut32.dll")] private static unsafe partial SAFEARRAY* SafeArrayCreate(VarEnum vt, uint cDims, in SAFEARRAYBOUND rgsabound); diff --git a/tests/Avalonia.IntegrationTests.Win32/Automation/SafeArrayRefTests.cs b/tests/Avalonia.IntegrationTests.Win32/Automation/SafeArrayRefTests.cs new file mode 100644 index 0000000000..8794a03b76 --- /dev/null +++ b/tests/Avalonia.IntegrationTests.Win32/Automation/SafeArrayRefTests.cs @@ -0,0 +1,83 @@ +using System; +using System.Runtime.CompilerServices; +using System.Runtime.InteropServices; +using System.Runtime.InteropServices.Marshalling; +using Avalonia.Win32.Automation.Marshalling; +using Xunit; + +namespace Avalonia.IntegrationTests.Win32.Automation; + +public unsafe class SafeArrayRefTests +{ + [Fact] + public void ConvertToUnmanaged_Sizes_String_Array_To_Input() + { + var safeArray = SafeArrayMarshaller.ConvertToUnmanaged(["foo", "bar", "baz"]); + + Assert.Equal(3, GetLength(safeArray)); + + SafeArrayMarshaller.Free(safeArray); + } + + [Fact] + public void ConvertToUnmanaged_Sizes_Provider_Array_To_Input() + { + ITestProvider[] providers = [new TestProvider(1), new TestProvider(2)]; + + var safeArray = SafeArrayMarshaller.ConvertToUnmanaged(providers); + + Assert.Equal(providers.Length, GetLength(safeArray)); + + SafeArrayMarshaller.Free(safeArray); + } + + [Fact] + public void ConvertToUnmanaged_Fills_Provider_Array_With_Com_Wrappers() + { + ITestProvider[] providers = [new TestProvider(1), new TestProvider(2)]; + + var safeArray = SafeArrayMarshaller.ConvertToUnmanaged(providers); + + // Check the length before reading the entries. A wrongly sized array holds slots that were + // never written, and dereferencing those below would take down the test host. + Assert.Equal(providers.Length, GetLength(safeArray)); + Assert.All(GetEntries(safeArray), x => Assert.NotEqual(IntPtr.Zero, x)); + + var roundTripped = SafeArrayMarshaller.ConvertToManaged(safeArray); + Assert.NotNull(roundTripped); + Assert.Equal([1, 2], Array.ConvertAll(roundTripped, x => x.GetValue())); + + SafeArrayMarshaller.Free(safeArray); + } + + private static int GetLength(SafeArrayRef safeArray) + { + var ptr = (SafeArrayRef.SAFEARRAY*)Unsafe.As(ref safeArray); + return (int)ptr->rgsabound[0].cElements; + } + + private static IntPtr[] GetEntries(SafeArrayRef safeArray) + { + var ptr = (SafeArrayRef.SAFEARRAY*)Unsafe.As(ref safeArray); + var result = new IntPtr[(int)ptr->rgsabound[0].cElements]; + new Span(ptr->pvData, result.Length).CopyTo(result); + return result; + } +} + +[GeneratedComInterface] +[Guid("6ADEBBF3-6C63-4C1B-9C4F-9EE7C3B8A6D1")] +internal partial interface ITestProvider +{ + int GetValue(); +} + +[GeneratedComClass] +internal partial class TestProvider : ITestProvider +{ + private readonly int _value; + + public TestProvider(int value) => _value = value; + + public int GetValue() => _value; +}