Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -59,6 +59,11 @@ public IntPtr GetOrCreateComInterfaceForObject(object instance, CreateComInterfa
throw new PlatformNotSupportedException();
}

public IntPtr GetOrCreateComInterfaceForObject(object instance, CreateComInterfaceFlags flags, in Guid interfaceId)
{
throw new PlatformNotSupportedException();
}

public object GetOrCreateObjectForComInstance(IntPtr externalComObject, CreateObjectFlags flags)
{
throw new PlatformNotSupportedException();
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -835,6 +835,30 @@ public unsafe IntPtr GetOrCreateComInterfaceForObject(object instance, CreateCom
return managedObjectWrapper.ComIp;
}

/// <summary>
/// Gets or creates a COM representation of the supplied object and queries it for the specified interface.
/// </summary>
/// <param name="instance">The managed object to expose outside the .NET runtime.</param>
/// <param name="flags">A bitwise combination of the enumeration values that specifies how to create the COM representation.</param>
/// <param name="interfaceId">The identifier of the requested COM interface.</param>
/// <returns>A pointer to the requested COM interface. The caller is responsible for releasing the returned reference.</returns>
/// <exception cref="ArgumentNullException"><paramref name="instance"/> is <see langword="null"/>.</exception>
/// <exception cref="InvalidCastException">The COM representation does not support <paramref name="interfaceId"/>.</exception>
/// <remarks>
/// The COM representation is cached before querying for the requested interface, even if the query fails.
/// The requested interface is queried on every call.
/// </remarks>
public IntPtr GetOrCreateComInterfaceForObject(object instance, CreateComInterfaceFlags flags, in Guid interfaceId)
{
IntPtr unknown = GetOrCreateComInterfaceForObject(instance, flags);
int hr = Marshal.QueryInterface(unknown, in interfaceId, out IntPtr result);
Marshal.Release(unknown);

Marshal.ThrowExceptionForHR(hr);

return result;
}

private readonly struct CreateManagedObjectWrapperState(ComWrappers comWrappers, CreateComInterfaceFlags flags)
{
public readonly ComWrappers ComWrappers = comWrappers;
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -494,6 +494,7 @@ public StrategyBasedComWrappers() { }
protected static System.Runtime.InteropServices.Marshalling.IIUnknownCacheStrategy CreateDefaultCacheStrategy() { throw null; }
protected sealed override object CreateObject(nint externalComObject, System.Runtime.InteropServices.CreateObjectFlags flags) { throw null; }
unsafe protected sealed override object? CreateObject(nint externalComObject, System.Runtime.InteropServices.CreateObjectFlags flags, object? userState, out System.Runtime.InteropServices.CreatedWrapperFlags wrapperFlags) { throw null; }
public System.IntPtr GetOrCreateComInterfaceForObject<TInterface>(object instance, System.Runtime.InteropServices.CreateComInterfaceFlags flags) { throw null; }
protected virtual System.Runtime.InteropServices.Marshalling.IIUnknownInterfaceDetailsStrategy GetOrCreateInterfaceDetailsStrategy() { throw null; }
protected virtual System.Runtime.InteropServices.Marshalling.IIUnknownStrategy GetOrCreateIUnknownStrategy() { throw null; }
protected sealed override void ReleaseObjects(System.Collections.IEnumerable objects) { }
Expand Down Expand Up @@ -770,6 +771,7 @@ public struct ComInterfaceDispatch
public static T GetInstance<T>(ComInterfaceDispatch* dispatchPtr) where T : class { throw null; }
}
public System.IntPtr GetOrCreateComInterfaceForObject(object instance, System.Runtime.InteropServices.CreateComInterfaceFlags flags) { throw null; }
public System.IntPtr GetOrCreateComInterfaceForObject(object instance, System.Runtime.InteropServices.CreateComInterfaceFlags flags, in System.Guid interfaceId) { throw null; }
protected unsafe abstract ComInterfaceEntry* ComputeVtables(object obj, System.Runtime.InteropServices.CreateComInterfaceFlags flags, out int count);
public unsafe object GetOrCreateObjectForComInstance(System.IntPtr externalComObject, System.Runtime.InteropServices.CreateObjectFlags flags) { throw null; }
public unsafe object GetOrCreateObjectForComInstance(System.IntPtr externalComObject, System.Runtime.InteropServices.CreateObjectFlags flags, object? userState) { throw null; }
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -117,6 +117,9 @@
<resheader name="writer">
<value>System.Resources.ResXResourceWriter, System.Windows.Forms, Version=4.0.0.0, Culture=neutral, PublicKeyToken=b77a5c561934e089</value>
</resheader>
<data name="Argument_UnknownComInterface" xml:space="preserve">
<value>The interface details strategy does not provide COM interface details for type '{0}'.</value>
</data>
<data name="InvalidOperation_HCCountOverflow" xml:space="preserve">
<value>Handle collector count overflows or underflows.</value>
</data>
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -79,6 +79,33 @@ static IIUnknownInterfaceDetailsStrategy GetInteropStrategy()
/// <returns>The caching strategy to use for the new COM object.</returns>
protected virtual IIUnknownCacheStrategy CreateCacheStrategy() => CreateDefaultCacheStrategy();

/// <summary>
/// Gets or creates a COM representation of the supplied object and queries it for the interface represented by <typeparamref name="TInterface"/>.
/// </summary>
/// <typeparam name="TInterface">The managed type that represents the requested COM interface.</typeparam>
/// <param name="instance">The managed object to expose outside the .NET runtime.</param>
/// <param name="flags">A bitwise combination of the enumeration values that specifies how to create the COM representation.</param>
/// <returns>A pointer to the requested COM interface. The caller is responsible for releasing the returned reference.</returns>
/// <exception cref="ArgumentNullException"><paramref name="instance"/> is <see langword="null"/>.</exception>
/// <exception cref="ArgumentException">The interface details strategy does not provide details for <typeparamref name="TInterface"/>.</exception>
/// <exception cref="InvalidCastException">The COM representation does not support the interface represented by <typeparamref name="TInterface"/>.</exception>
/// <remarks>
/// The COM representation is cached before querying for the requested interface, even if the query fails.
/// The requested interface is queried on every call.
/// </remarks>
public IntPtr GetOrCreateComInterfaceForObject<TInterface>(object instance, CreateComInterfaceFlags flags)
{
ArgumentNullException.ThrowIfNull(instance);

IIUnknownDerivedDetails? details = GetOrCreateInterfaceDetailsStrategy().GetIUnknownDerivedDetails(typeof(TInterface).TypeHandle);
if (details is null)
{
throw new ArgumentException(SR.Format(SR.Argument_UnknownComInterface, typeof(TInterface)), nameof(TInterface));
}

return GetOrCreateComInterfaceForObject(instance, flags, details.Iid);
}

/// <inheritdoc cref="ComWrappers.ComputeVtables" />
protected sealed override unsafe ComInterfaceEntry* ComputeVtables(object obj, CreateComInterfaceFlags flags, out int count)
{
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -37,6 +37,20 @@ partial class DerivedComObject : ManagedObjectExposedToCom
{
}

[GeneratedComClass]
partial class QueryCountingComObject : ManagedObjectExposedToCom, ICustomQueryInterface
{
public int QueryCount { get; private set; }
public bool FailQuery { get; set; }

public CustomQueryInterfaceResult GetInterface(ref Guid iid, out nint ppv)
{
QueryCount++;
ppv = 0;
return FailQuery ? CustomQueryInterfaceResult.Failed : CustomQueryInterfaceResult.NotHandled;
}
}

[GeneratedComInterface]
[Guid("781E56C2-A530-4A8F-90FE-01244426E0CC")]
partial interface IActivationFactory
Expand All @@ -54,6 +68,107 @@ public unsafe class GeneratedComClassTests
{
private const int E_NOINTERFACE = unchecked((int)0x80004002);

[Theory]
[InlineData(false)]
[InlineData(true)]
public void GetOrCreateComInterfaceForObject_NullInstance(bool useGenericOverload)
{
StrategyBasedComWrappers wrappers = new();
Assert.Throws<ArgumentNullException>("instance", () =>
useGenericOverload
? wrappers.GetOrCreateComInterfaceForObject<IGetAndSetInt>(null, CreateComInterfaceFlags.None)
: wrappers.GetOrCreateComInterfaceForObject(null, CreateComInterfaceFlags.None, Guid.Empty));
}

[Theory]
[InlineData(false)]
[InlineData(true)]
public void GetOrCreateComInterfaceForObject_UnsupportedInterfacePreservesCache(bool useGenericOverload)
{
ManagedObjectExposedToCom obj = new();
InterfaceDetailsComWrappers wrappers = new();
Guid iid = StrategyBasedComWrappers.DefaultIUnknownInterfaceDetailsStrategy.GetIUnknownDerivedDetails(typeof(IActivationFactory).TypeHandle).Iid;
for (int i = 0; i < 2; i++)
{
Assert.Throws<InvalidCastException>(() =>
useGenericOverload
? wrappers.GetOrCreateComInterfaceForObject<IActivationFactory>(obj, CreateComInterfaceFlags.None)
: wrappers.GetOrCreateComInterfaceForObject(obj, CreateComInterfaceFlags.None, in iid));
}

Assert.Equal(1, wrappers.Details.ExposedTypeLookups);
nint ptr = wrappers.GetOrCreateComInterfaceForObject<IGetAndSetInt>(obj, CreateComInterfaceFlags.None);
Assert.Equal(0, Marshal.Release(ptr));
Assert.Equal(1, wrappers.Details.ExposedTypeLookups);
GC.KeepAlive(obj);
GC.KeepAlive(wrappers);
}

[Fact]
public void GetOrCreateComInterfaceForObject_UnknownInterfaceType()
{
StrategyBasedComWrappers wrappers = new();
ManagedObjectExposedToCom obj = new();
Assert.Throws<ArgumentException>("TInterface", () => wrappers.GetOrCreateComInterfaceForObject<IDisposable>(obj, CreateComInterfaceFlags.None));
Assert.Throws<ArgumentException>("TInterface", () => wrappers.GetOrCreateComInterfaceForObject<object>(obj, CreateComInterfaceFlags.None));
Assert.Throws<ArgumentException>("TInterface", () => wrappers.GetOrCreateComInterfaceForObject<int>(obj, CreateComInterfaceFlags.None));
}

[Theory]
[InlineData(false, false)]
[InlineData(true, false)]
[InlineData(false, true)]
[InlineData(true, true)]
public void GetOrCreateComInterfaceForObject_QueriesOnEveryCall(bool useGenericOverload, bool failQuery)
{
QueryCountingComObject obj = new() { FailQuery = failQuery };
StrategyBasedComWrappers wrappers = new();
Guid iid = StrategyBasedComWrappers.DefaultIUnknownInterfaceDetailsStrategy.GetIUnknownDerivedDetails(typeof(IGetAndSetInt).TypeHandle).Iid;
for (int i = 0; i < 2; i++)
{
if (failQuery)
{
Assert.Throws<InvalidCastException>(() => GetInterface());
}
else
{
Assert.Equal(0, Marshal.Release(GetInterface()));
}
}
Assert.Equal(2, obj.QueryCount);
nint unknown = wrappers.GetOrCreateComInterfaceForObject(obj, CreateComInterfaceFlags.None);
Assert.Equal(0, Marshal.Release(unknown));
GC.KeepAlive(obj);
GC.KeepAlive(wrappers);

nint GetInterface() => useGenericOverload
? wrappers.GetOrCreateComInterfaceForObject<IGetAndSetInt>(obj, CreateComInterfaceFlags.None)
: wrappers.GetOrCreateComInterfaceForObject(obj, CreateComInterfaceFlags.None, in iid);
}

private sealed class InterfaceDetailsComWrappers : StrategyBasedComWrappers
{
public CountingInterfaceDetailsStrategy Details { get; } = new();

protected override IIUnknownInterfaceDetailsStrategy GetOrCreateInterfaceDetailsStrategy() => Details;
}

private sealed class CountingInterfaceDetailsStrategy : IIUnknownInterfaceDetailsStrategy
{
public int ExposedTypeLookups { get; private set; }

public IComExposedDetails GetComExposedTypeDetails(RuntimeTypeHandle type)
{
ExposedTypeLookups++;
return StrategyBasedComWrappers.DefaultIUnknownInterfaceDetailsStrategy.GetComExposedTypeDetails(type);
}

public IIUnknownDerivedDetails GetIUnknownDerivedDetails(RuntimeTypeHandle type)
{
return StrategyBasedComWrappers.DefaultIUnknownInterfaceDetailsStrategy.GetIUnknownDerivedDetails(type);
}
}

[Fact]
public void ComInstanceProvidesInterfaceForDirectlyImplementedComInterface()
{
Expand Down
Loading