File: DownlevelLibraryImportGenerator.cs
Web Access
Project: DownlevelLibraryImportGenerator.csproj (Microsoft.Interop.LibraryImportGenerator.Downlevel)
// 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}";
    }
}