// Licensed to the .NET Foundation under one or more agreements.
// The .NET Foundation licenses this file to you under the MIT license.
using System;
using System.Collections.Generic;
using System.Collections.Immutable;
using System.Diagnostics;
using System.Globalization;
using System.Runtime.InteropServices;
using System.Threading;
using Microsoft.CodeAnalysis;
using Microsoft.CodeAnalysis.CSharp;
using Microsoft.CodeAnalysis.CSharp.Syntax;
using SourceGenerators;
[assembly: System.Resources.NeutralResourcesLanguage("en-US")]
namespace Microsoft.Interop
{
[Generator]
public sealed class DownlevelLibraryImportGenerator : IIncrementalGenerator
{
internal sealed record IncrementalStubGenerationContext(
SignatureContext SignatureContext,
ContainingSyntaxContext ContainingSyntaxContext,
DeclarationHeader StubMethodSyntaxTemplate,
string MethodName,
MethodSignatureDiagnosticLocations DiagnosticLocation,
SequenceEqualImmutableArray<string> ForwardedAttributes,
LibraryImportData LibraryImportData,
EnvironmentFlags EnvironmentFlags);
public static class StepNames
{
public const string CalculateStubInformation = nameof(CalculateStubInformation);
public const string GenerateSingleStub = nameof(GenerateSingleStub);
}
// Internal definitions of LibraryImportAttribute and StringMarshalling that are injected
// into user projects targeting downlevel frameworks that don't have these types in the BCL.
// These definitions only expose the properties and enum values supported downlevel.
private const string GeneratedInteropTypes = """
// <auto-generated/>
#nullable enable
namespace System.Runtime.InteropServices
{
[global::System.AttributeUsage(global::System.AttributeTargets.Method, AllowMultiple = false, Inherited = false)]
[global::Microsoft.CodeAnalysis.Embedded]
internal sealed partial class LibraryImportAttribute : global::System.Attribute
{
public LibraryImportAttribute(string libraryName) { }
public string? EntryPoint { get; set; }
public StringMarshalling StringMarshalling { get; set; }
public bool SetLastError { get; set; }
}
[global::Microsoft.CodeAnalysis.Embedded]
internal enum StringMarshalling
{
Utf16 = 2,
}
}
""";
public void Initialize(IncrementalGeneratorInitializationContext context)
{
// Inject internal definitions of LibraryImportAttribute and StringMarshalling into the compilation
// so that users targeting downlevel frameworks can apply [LibraryImport] to their methods.
context.RegisterPostInitializationOutput(static ctx =>
{
ctx.AddEmbeddedAttributeDefinition();
ctx.AddSource("LibraryImportInteropTypes.g.cs", GeneratedInteropTypes);
});
// Collect all methods adorned with LibraryImportAttribute and filter out invalid ones
// (diagnostics for invalid methods are reported by the analyzer)
var methodsToGenerate = context.SyntaxProvider
.ForAttributeWithMetadataName(
TypeNames.LibraryImportAttribute,
static (node, ct) => node is MethodDeclarationSyntax,
static (context, ct) => context.TargetSymbol is IMethodSymbol methodSymbol
? new { Syntax = (MethodDeclarationSyntax)context.TargetNode, Symbol = methodSymbol }
: null)
.Where(
static modelData => modelData is not null
&& Analyzers.DownlevelLibraryImportDiagnosticsAnalyzer.GetDiagnosticIfInvalidMethodForGeneration(modelData.Syntax, modelData.Symbol) is null);
IncrementalValueProvider<StubEnvironment> stubEnvironment = context.CreateStubEnvironmentProvider();
IncrementalValuesProvider<string> generateSingleStub = methodsToGenerate
.Combine(stubEnvironment)
.Select(static (data, ct) => new
{
data.Left.Syntax,
data.Left.Symbol,
Environment = data.Right,
})
.Select(
static (data, ct) => CalculateStubInformation(data.Syntax, data.Symbol, data.Environment, ct)
)
.WithTrackingName(StepNames.CalculateStubInformation)
.Select(
static (data, ct) => GenerateSource(data)
)
.WithComparer(StringComparer.Ordinal)
.WithTrackingName(StepNames.GenerateSingleStub);
context.RegisterConcatenatedOutputs(generateSingleStub, "LibraryImports.g.cs");
}
private static List<string> GenerateForwardedAttributes(AttributeData? defaultDllImportSearchPathsAttribute)
{
// Manually rehydrate the forwarded attributes with fully qualified types so we don't have to worry about any using directives.
List<string> attributes = [];
if (defaultDllImportSearchPathsAttribute is not null)
{
string searchPaths = ((int)defaultDllImportSearchPathsAttribute.ConstructorArguments[0].Value!).ToString(CultureInfo.InvariantCulture);
attributes.Add($"{TypeNames.GlobalAlias}{TypeNames.DefaultDllImportSearchPathsAttribute}(({TypeNames.GlobalAlias}{TypeNames.DllImportSearchPath}){searchPaths})");
}
return attributes;
}
private static string PrintGeneratedSource(
IncrementalStubGenerationContext stub,
ManagedToNativeStubGenerator stubGenerator)
{
var writer = new IndentedTextWriter();
foreach (string attribute in stub.SignatureContext.AdditionalAttributes)
{
writer.WriteLine($"[{attribute}]");
}
DeclarationHeader userDeclaredMethod = stub.StubMethodSyntaxTemplate;
writer.WriteLine($"{string.Join(" ", userDeclaredMethod.Modifiers)} {stub.SignatureContext.StubReturnType} {userDeclaredMethod.Identifier}({string.Join(", ", stub.SignatureContext.StubParameters)})");
// Create stub function. The generated body performs unmanaged operations (pointers, fixed,
// stackalloc, calling the extern local P/Invoke), so it is wrapped in an explicit unsafe block
// rather than relying on an unsafe modifier on the containing type.
using (writer.WriteBlock())
{
const string InnerPInvokeName = "__PInvoke";
writer.WriteLine("unsafe");
using (writer.WriteBlock())
{
stubGenerator.GenerateStubStatements(writer, InnerPInvokeName);
writer.WriteLine("// Local P/Invoke");
WriteTargetDllImport(writer, stubGenerator, stub, InnerPInvokeName);
}
}
return writer.ToString();
}
private static LibraryImportCompilationData? ProcessLibraryImportAttribute(AttributeData attrData)
{
// Found the LibraryImport, but it has an error so report the error.
// This is most likely an issue with targeting an incorrect TFM.
if (attrData.AttributeClass?.TypeKind is null or TypeKind.Error)
{
return null;
}
if (attrData.ConstructorArguments.Length == 0)
{
return null;
}
ImmutableDictionary<string, TypedConstant> namedArguments = ImmutableDictionary.CreateRange(attrData.NamedArguments);
string? entryPoint = null;
if (namedArguments.TryGetValue(nameof(LibraryImportCompilationData.EntryPoint), out TypedConstant entryPointValue))
{
if (entryPointValue.Value is not string)
{
return null;
}
entryPoint = (string)entryPointValue.Value!;
}
return new LibraryImportCompilationData(attrData.ConstructorArguments[0].Value!.ToString())
{
EntryPoint = entryPoint,
}.WithValuesFromNamedArguments(namedArguments);
}
private static IncrementalStubGenerationContext CalculateStubInformation(
MethodDeclarationSyntax originalSyntax,
IMethodSymbol symbol,
StubEnvironment environment,
CancellationToken ct)
{
ct.ThrowIfCancellationRequested();
INamedTypeSymbol? defaultDllImportSearchPathsAttrType = environment.DefaultDllImportSearchPathsAttrType;
// Get any attributes of interest on the method
AttributeData? generatedDllImportAttr = null;
AttributeData? defaultDllImportSearchPathsAttribute = null;
foreach (AttributeData attr in symbol.GetAttributes())
{
if (attr.AttributeClass is not null
&& attr.AttributeClass.ToDisplayString() == TypeNames.LibraryImportAttribute)
{
generatedDllImportAttr = attr;
}
else if (defaultDllImportSearchPathsAttrType is not null && SymbolEqualityComparer.Default.Equals(attr.AttributeClass, defaultDllImportSearchPathsAttrType))
{
defaultDllImportSearchPathsAttribute = attr;
}
}
Debug.Assert(generatedDllImportAttr is not null);
var locations = new MethodSignatureDiagnosticLocations(originalSyntax);
// Process the LibraryImport attribute
LibraryImportCompilationData libraryImportData =
ProcessLibraryImportAttribute(generatedDllImportAttr!) ??
new LibraryImportCompilationData("INVALID_CSHARP_SYNTAX");
// Create a diagnostics bag that discards all diagnostics.
// Diagnostics are now reported by the analyzer, not the generator.
var discardedDiagnostics = new GeneratorDiagnosticsBag(new DiagnosticDescriptorProvider(), locations, SR.ResourceManager, typeof(FxResources.Microsoft.Interop.LibraryImportGenerator.Downlevel.SR));
// Create the stub.
var signatureContext = SignatureContext.Create(
symbol,
DownlevelLibraryImportGeneratorHelpers.CreateMarshallingInfoParser(environment, discardedDiagnostics, symbol, libraryImportData),
environment,
new CodeEmitOptions(SkipInit: false),
typeof(DownlevelLibraryImportGenerator).Assembly);
ContainingSyntaxContext containingTypeContext = originalSyntax.GetContainingSyntaxContext();
DeclarationHeader methodSyntaxTemplate = ContainingTypeUtilities.GetDeclarationHeader(originalSyntax);
List<string> additionalAttributes = GenerateForwardedAttributes(defaultDllImportSearchPathsAttribute);
return new IncrementalStubGenerationContext(
signatureContext,
containingTypeContext,
methodSyntaxTemplate,
symbol.Name,
locations,
new SequenceEqualImmutableArray<string>(additionalAttributes.ToImmutableArray(), StringComparer.Ordinal),
LibraryImportData.From(libraryImportData),
environment.EnvironmentFlags
);
}
private static string GenerateSource(
IncrementalStubGenerationContext pinvokeStub)
{
// Note: Diagnostics are now reported by the analyzer, so we pass a discarding diagnostics bag
var discardedDiagnostics = new GeneratorDiagnosticsBag(new DiagnosticDescriptorProvider(), pinvokeStub.DiagnosticLocation, SR.ResourceManager, typeof(FxResources.Microsoft.Interop.LibraryImportGenerator.Downlevel.SR));
// Generate stub code
var stubGenerator = new ManagedToNativeStubGenerator(
pinvokeStub.SignatureContext.ElementTypeInformation,
pinvokeStub.LibraryImportData.SetLastError,
discardedDiagnostics,
DownlevelLibraryImportGeneratorHelpers.GeneratorResolver,
new CodeEmitOptions(SkipInit: false));
// Check if the generator should produce a forwarder stub - regular DllImport.
// This is done if the signature is blittable or if some parameters cannot be marshalled.
if (stubGenerator.NoMarshallingRequired
|| stubGenerator.HasForwardedTypes
|| pinvokeStub.LibraryImportData.SetLastError)
{
return PrintForwarderStub(pinvokeStub.StubMethodSyntaxTemplate, pinvokeStub);
}
return pinvokeStub.ContainingSyntaxContext.WrapMemberInContainingSyntax(PrintGeneratedSource(pinvokeStub, stubGenerator));
}
private static string PrintForwarderStub(DeclarationHeader userDeclaredMethod, IncrementalStubGenerationContext stub)
{
var writer = new IndentedTextWriter();
ImmutableArray<string> modifiers = CodeWriterHelpers.AddModifier(userDeclaredMethod.Modifiers, "extern");
writer.WriteLine($"[{CreateDllImportAttribute(stub.LibraryImportData, stub.MethodName, forwardSetLastError: true)}]");
writer.WriteLine($"{string.Join(" ", modifiers)} {stub.SignatureContext.StubReturnType} {userDeclaredMethod.Identifier}({string.Join(", ", stub.SignatureContext.StubParameters)});");
return stub.ContainingSyntaxContext.WrapMemberInContainingSyntax(writer.ToString());
}
private static void WriteTargetDllImport(
IndentedTextWriter writer,
ManagedToNativeStubGenerator stubGenerator,
IncrementalStubGenerationContext stub,
string stubTargetName)
{
GeneratedMethodSignature signature = stubGenerator.GenerateTargetMethodSignatureData();
writer.WriteLine($"[{CreateDllImportAttribute(stub.LibraryImportData, stub.MethodName, forwardSetLastError: false)}]");
if (!string.IsNullOrEmpty(signature.ReturnTypeAttributes))
{
writer.WriteLine($"[return: {signature.ReturnTypeAttributes}]");
}
if (!stub.ForwardedAttributes.Array.IsEmpty)
{
writer.WriteLine($"[{string.Join(", ", stub.ForwardedAttributes.Array)}]");
}
writer.WriteLine($"static extern unsafe {signature.ReturnType} {stubTargetName}{signature.ParameterList};");
}
private static string CreateDllImportAttribute(LibraryImportData target, string methodName, bool forwardSetLastError)
{
var arguments = new List<string>
{
CodeWriterHelpers.StringLiteral(target.ModuleName),
$"{nameof(DllImportAttribute.EntryPoint)} = {CodeWriterHelpers.StringLiteral(target.EntryPoint ?? methodName)}",
$"{nameof(DllImportAttribute.ExactSpelling)} = true"
};
// Forward the charset to either interop boundary so runtime-marshalled types use the requested encoding.
if (target.IsUserDefined.HasFlag(InteropAttributeMember.StringMarshalling)
&& target.StringMarshalling == StringMarshalling.Utf16)
{
arguments.Add($"{nameof(DllImportAttribute.CharSet)} = {CreateEnumExpression(CharSet.Unicode)}");
}
if (forwardSetLastError && target.IsUserDefined.HasFlag(InteropAttributeMember.SetLastError))
{
arguments.Add($"{nameof(DllImportAttribute.SetLastError)} = {(target.SetLastError ? "true" : "false")}");
}
return $"{TypeNames.GlobalAlias}{TypeNames.DllImportAttribute}({string.Join(", ", arguments)})";
}
private static string CreateEnumExpression<T>(T value) where T : Enum
=> $"{TypeNames.GlobalAlias}{typeof(T).FullName}.{value}";
}
}