From 6a502e4e7d94fd93711bdbf4553c53bf4712b7ef Mon Sep 17 00:00:00 2001 From: AriusII Date: Mon, 21 Sep 2026 19:18:47 +0200 Subject: [PATCH] Enforce recursive Client API type boundaries --- .../AnalyzerReleases.Unshipped.md | 2 +- .../CheatEngineLuaGenerator.cs | 337 +++++++++++++++++- .../CheatEngineLuaGeneratorTests.cs | 176 ++++++++- .../Infrastructure/GeneratorRun.cs | 15 +- .../CheatEngine.Client.Tests.csproj | 4 + .../PublicClientSignatureBoundaryTests.cs | 250 +++++++++++-- 6 files changed, 748 insertions(+), 36 deletions(-) diff --git a/source-generators/CheatEngine.Client.SourceGenerators.Lua/AnalyzerReleases.Unshipped.md b/source-generators/CheatEngine.Client.SourceGenerators.Lua/AnalyzerReleases.Unshipped.md index afc3277..c776cd5 100644 --- a/source-generators/CheatEngine.Client.SourceGenerators.Lua/AnalyzerReleases.Unshipped.md +++ b/source-generators/CheatEngine.Client.SourceGenerators.Lua/AnalyzerReleases.Unshipped.md @@ -13,4 +13,4 @@ CECLUA1103 | CheatEngine.Client.Lua | Error | Generated Lua operation signatures must be bounded. CECLUA1104 | CheatEngine.Client.Lua | Error | Non-scalar Lua operation results require a mapper. CECLUA1105 | CheatEngine.Client.Lua | Error | A Lua operation mapper must match the SDK result. - CECLUA1106 | CheatEngine.Client.Lua | Error | A Lua mapper must project a safe Client result. + CECLUA1106 | CheatEngine.Client.Lua | Error | A Lua mapper must project a safe, recursively closed Client result and source graph. diff --git a/source-generators/CheatEngine.Client.SourceGenerators.Lua/CheatEngineLuaGenerator.cs b/source-generators/CheatEngine.Client.SourceGenerators.Lua/CheatEngineLuaGenerator.cs index 3089817..876983b 100644 --- a/source-generators/CheatEngine.Client.SourceGenerators.Lua/CheatEngineLuaGenerator.cs +++ b/source-generators/CheatEngine.Client.SourceGenerators.Lua/CheatEngineLuaGenerator.cs @@ -26,6 +26,38 @@ public sealed class CheatEngineLuaGenerator : IIncrementalGenerator private const string LuaFunctionAttributeMetadataName = "CheatEngine.SDK.Annotations.Lua.LuaFunctionAttribute"; private const string LuaGlobalAttributeMetadataName = "CheatEngine.SDK.Annotations.Lua.LuaGlobalAttribute"; private const string LuaResultMapperMetadataName = "CheatEngine.Client.Lua.ILuaResultMapper"; + private const string LuaClassAttributeMetadataName = "CheatEngine.SDK.Annotations.Lua.LuaClassAttribute"; + private const string CheatEngineSdkAssemblyPrefix = "CheatEngine.SDK"; + private const string CheatEngineSdkObjectContractMetadataName = + "CheatEngine.SDK.Engine.Objects.ICEObject"; + private const int ClientBoundaryMaximumDepth = 32; + private const int ClientBoundaryMaximumNodes = 256; + + private static readonly HashSet ApprovedSdkClientResultTypes = new(StringComparer.Ordinal) + { + "CheatEngine.SDK.Engine.AddressList.MemoryRecordId", + "CheatEngine.SDK.Engine.Enums.FastScanMethod", + "CheatEngine.SDK.Engine.Enums.VariableType", + "CheatEngine.SDK.Engine.Inspection.AddressResolutionOptions", + "CheatEngine.SDK.Engine.Inspection.MemoryRegionInfo", + "CheatEngine.SDK.Engine.Inspection.ModuleInfo", + "CheatEngine.SDK.Engine.Inspection.ModuleName", + "CheatEngine.SDK.Engine.Inspection.ModuleSectionInfo", + "CheatEngine.SDK.Engine.Inspection.SymbolExpression", + "CheatEngine.SDK.Engine.Inspection.SymbolInfo", + "CheatEngine.SDK.Engine.Inspection.TargetProcessId", + "CheatEngine.SDK.Engine.Runtime.CheatEngineArchitecture", + "CheatEngine.SDK.Engine.Runtime.CheatEngineVersion", + "CheatEngine.SDK.Engine.Runtime.PointerSize", + "CheatEngine.SDK.Engine.Runtime.RuntimeCapabilityAvailability", + "CheatEngine.SDK.Engine.Runtime.RuntimeCapabilityId", + "CheatEngine.SDK.Engine.Runtime.TargetAbi", + "CheatEngine.SDK.Engine.Scanning.Aob.AobPattern", + "CheatEngine.SDK.Engine.Scanning.Aob.AobScanOptions", + "CheatEngine.SDK.Engine.Scanning.Values.FirstScanRequest", + "CheatEngine.SDK.Engine.Scanning.Values.NextScanRequest", + "CheatEngine.SDK.Engine.Values.Address" + }; /// public void Initialize(IncrementalGeneratorInitializationContext context) @@ -428,23 +460,296 @@ SpecialType.System_UInt64 or SpecialType.System_Single or SpecialType.System_Dou }; } - private static bool IsSafeMapperSource(ITypeSymbol type) + private static bool TryFindClientBoundaryViolation(ITypeSymbol type, ClientBoundaryRole role, + out string violation) + { + HashSet visited = new(SymbolEqualityComparer.Default); + int visitedCount = 0; + string rootPath = role == ClientBoundaryRole.MapperSource ? "mapper source" : "mapper result"; + return TryFindClientBoundaryViolation(type, role, rootPath, visited, ref visitedCount, 0, out violation); + } + + private static bool TryFindClientBoundaryViolation(ITypeSymbol type, ClientBoundaryRole role, string path, + HashSet visited, ref int visitedCount, int depth, out string violation) + { + if (depth > ClientBoundaryMaximumDepth || ++visitedCount > ClientBoundaryMaximumNodes) + { + violation = path + " exceeds the supported Client result type graph budget."; + return true; + } + + if (!visited.Add(type)) + { + violation = string.Empty; + return false; + } + + if (type.TypeKind == TypeKind.Error || type.TypeKind == TypeKind.Dynamic) + { + violation = path + " uses an unresolved or dynamic type."; + return true; + } + + if (type is IFunctionPointerTypeSymbol) + { + violation = path + " exposes a function pointer."; + return true; + } + + if (type is IPointerTypeSymbol) + { + violation = path + " exposes a pointer."; + return true; + } + + if (type.IsRefLikeType) + { + violation = path + " exposes a ref-like type."; + return true; + } + + if (type is IArrayTypeSymbol array) + { + return TryFindClientBoundaryViolation(array.ElementType, role, path + "[]", visited, ref visitedCount, + depth + 1, out violation); + } + + if (type is ITypeParameterSymbol typeParameter) + { + return TryFindConstraintViolation(typeParameter, role, path, visited, ref visitedCount, depth, out violation); + } + + if (type is not INamedTypeSymbol named) + { + violation = path + " uses an unsupported type shape '" + TypeName(type) + "'."; + return true; + } + + if (TryFindNamedTypeViolation(named, role, path, out violation)) + { + return true; + } + + foreach (ITypeParameterSymbol parameter in named.OriginalDefinition.TypeParameters + .OrderBy(static candidate => candidate.Ordinal)) + { + if (TryFindConstraintViolation(parameter, role, path + "." + parameter.Name, visited, ref visitedCount, + depth + 1, out violation)) + { + return true; + } + } + + if (named.IsTupleType) + { + foreach (IFieldSymbol element in named.TupleElements) + { + if (TryFindClientBoundaryViolation(element.Type, role, path + "." + element.Name, visited, + ref visitedCount, depth + 1, out violation)) + { + return true; + } + } + } + + foreach (ITypeSymbol argument in named.TypeArguments) + { + if (TryFindClientBoundaryViolation(argument, role, path + "<" + TypeName(argument) + ">", visited, + ref visitedCount, depth + 1, out violation)) + { + return true; + } + } + + if (IsFrameworkOrApprovedSdkValue(named) || + (role == ClientBoundaryRole.MapperSource && IsCheatEngineSdkType(named))) + { + violation = string.Empty; + return false; + } + + return TryFindUserDefinedDtoViolation(named, role, path, visited, ref visitedCount, depth, out violation); + } + + private static bool TryFindNamedTypeViolation(INamedTypeSymbol type, ClientBoundaryRole role, string path, + out string violation) + { + string metadataName = type.OriginalDefinition.ToDisplayString(); + if (metadataName is "CheatEngine.SDK.Lua.State.LuaState" or "CheatEngine.SDK.Lua.References.LuaRef" or + "CheatEngine.SDK.Engine.Objects.CEObject" or "CheatEngine.SDK.Engine.Objects.Owned" || + type.ContainingNamespace.ToDisplayString().Contains(".Interop", StringComparison.Ordinal) || + HasAttribute(type, LuaClassAttributeMetadataName) || ImplementsSdkObjectContract(type)) + { + violation = path + " exposes forbidden SDK lifetime or interop type '" + metadataName + "'."; + return true; + } + + if (type.TypeKind == TypeKind.Delegate) + { + violation = path + " exposes a delegate or callback."; + return true; + } + + if (role == ClientBoundaryRole.ClientResult && + type.ContainingAssembly?.Name.StartsWith(CheatEngineSdkAssemblyPrefix, StringComparison.Ordinal) == true && + !ApprovedSdkClientResultTypes.Contains(metadataName)) + { + violation = path + " exposes non-approved SDK type '" + metadataName + "'."; + return true; + } + + if (type.SpecialType == SpecialType.System_Object || metadataName == "System.Type") + { + violation = path + " exposes an unqualified object or runtime type."; + return true; + } + + violation = string.Empty; + return false; + } + + private static bool TryFindConstraintViolation(ITypeParameterSymbol typeParameter, ClientBoundaryRole role, + string path, HashSet visited, ref int visitedCount, int depth, out string violation) { - return !IsRawLuaOrOwnershipType(type) && type.TypeKind != TypeKind.Pointer && !type.IsRefLikeType; + foreach (ITypeSymbol constraint in typeParameter.ConstraintTypes.OrderBy(TypeName, StringComparer.Ordinal)) + { + if (TryFindClientBoundaryViolation(constraint, role, path + " constraint", visited, ref visitedCount, + depth + 1, out violation)) + { + return true; + } + } + + violation = string.Empty; + return false; } - private static bool IsSafeClientResult(ITypeSymbol type) + private static bool TryFindUserDefinedDtoViolation(INamedTypeSymbol type, ClientBoundaryRole role, string path, + HashSet visited, ref int visitedCount, int depth, out string violation) { - return IsSafeMapperSource(type) && !HasAttribute(type, "CheatEngine.SDK.Annotations.Lua.LuaClassAttribute"); + if (type.BaseType is { SpecialType: not SpecialType.System_Object } baseType && + TryFindClientBoundaryViolation(baseType, role, path + ".base", visited, ref visitedCount, depth + 1, + out violation)) + { + return true; + } + + foreach (INamedTypeSymbol implementedInterface in type.Interfaces.OrderBy(TypeName, StringComparer.Ordinal)) + { + if (TryFindClientBoundaryViolation(implementedInterface, role, path + ".interface", visited, + ref visitedCount, depth + 1, out violation)) + { + return true; + } + } + + foreach (ISymbol member in type.GetMembers().OrderBy(static candidate => candidate.MetadataName, + StringComparer.Ordinal)) + { + switch (member) + { + case IFieldSymbol { IsStatic: false } field: + if (TryFindClientBoundaryViolation(field.Type, role, path + "." + field.Name, visited, + ref visitedCount, depth + 1, out violation)) + { + return true; + } + + break; + case IPropertySymbol { IsStatic: false } property: + if (TryFindClientBoundaryViolation(property.Type, role, path + "." + property.Name, visited, + ref visitedCount, depth + 1, out violation) || + TryFindParameterViolation(property.Parameters, role, path + "." + property.Name, visited, + ref visitedCount, depth + 1, out violation)) + { + return true; + } + + break; + case IEventSymbol { IsStatic: false } @event: + if (TryFindClientBoundaryViolation(@event.Type, role, path + "." + @event.Name, visited, + ref visitedCount, depth + 1, out violation)) + { + return true; + } + + break; + case IMethodSymbol { IsStatic: false } method when !method.IsImplicitlyDeclared && + method.DeclaredAccessibility != Accessibility.Private: + violation = string.Empty; + if (method.ReturnsByRef || method.ReturnsByRefReadonly || + TryFindClientBoundaryViolation(method.ReturnType, role, path + "." + method.Name, visited, + ref visitedCount, depth + 1, out violation) || + TryFindParameterViolation(method.Parameters, role, path + "." + method.Name, visited, + ref visitedCount, depth + 1, out violation)) + { + violation = string.IsNullOrEmpty(violation) + ? path + "." + method.Name + " exposes a by-reference return." + : violation; + return true; + } + + break; + } + } + + violation = string.Empty; + return false; + } + + private static bool TryFindParameterViolation(ImmutableArray parameters, ClientBoundaryRole role, + string path, HashSet visited, ref int visitedCount, int depth, out string violation) + { + foreach (IParameterSymbol parameter in parameters.OrderBy(static candidate => candidate.Ordinal)) + { + if (parameter.RefKind != RefKind.None) + { + violation = path + " parameter '" + parameter.Name + "' is passed by reference."; + return true; + } + + if (TryFindClientBoundaryViolation(parameter.Type, role, path + " parameter '" + parameter.Name + "'", + visited, ref visitedCount, depth + 1, out violation)) + { + return true; + } + } + + violation = string.Empty; + return false; } - private static bool IsRawLuaOrOwnershipType(ITypeSymbol type) + private static bool IsFrameworkOrApprovedSdkValue(INamedTypeSymbol type) { string metadataName = type.OriginalDefinition.ToDisplayString(); - return string.Equals(metadataName, "CheatEngine.SDK.Lua.State.LuaState", StringComparison.Ordinal) || - string.Equals(metadataName, "CheatEngine.SDK.Lua.References.LuaRef", StringComparison.Ordinal) || - string.Equals(metadataName, "CheatEngine.SDK.Engine.Objects.CEObject", StringComparison.Ordinal) || - string.Equals(metadataName, "CheatEngine.SDK.Engine.Objects.Owned", StringComparison.Ordinal); + if (ApprovedSdkClientResultTypes.Contains(metadataName)) + { + return true; + } + + string? assemblyName = type.ContainingAssembly?.Name; + return assemblyName is not null && + !assemblyName.StartsWith(CheatEngineSdkAssemblyPrefix, StringComparison.Ordinal) && + (assemblyName.StartsWith("System", StringComparison.Ordinal) || + assemblyName.StartsWith("Microsoft", StringComparison.Ordinal)); + } + + private static bool IsCheatEngineSdkType(INamedTypeSymbol type) + { + return type.ContainingAssembly?.Name.StartsWith(CheatEngineSdkAssemblyPrefix, StringComparison.Ordinal) == true; + } + + private static bool ImplementsSdkObjectContract(INamedTypeSymbol type) + { + return type.AllInterfaces.Any(@interface => + string.Equals(@interface.OriginalDefinition.ToDisplayString(), CheatEngineSdkObjectContractMetadataName, + StringComparison.Ordinal)); + } + + private enum ClientBoundaryRole + { + MapperSource, + ClientResult } private static string TypeDeclaration(INamedTypeSymbol type, bool isStatic) @@ -793,9 +1098,17 @@ method.PartialImplementationPart is not null || method.ReturnsByRef || method.Re { return Invalid(OperationDiagnosticDescriptors.InvalidMapper, location, mapper.Name, method.Name); } - else if (!IsSafeMapperSource(sourceResult) || !IsSafeClientResult(result)) + else if (TryFindClientBoundaryViolation(sourceResult, ClientBoundaryRole.MapperSource, + out string sourceViolation)) + { + return Invalid(OperationDiagnosticDescriptors.UnsafeMappedType, location, method.Name, + sourceViolation); + } + else if (TryFindClientBoundaryViolation(result, ClientBoundaryRole.ClientResult, + out string resultViolation)) { - return Invalid(OperationDiagnosticDescriptors.UnsafeMappedType, location, method.Name); + return Invalid(OperationDiagnosticDescriptors.UnsafeMappedType, location, method.Name, + resultViolation); } string operationTypeName = method.Name + "LuaOperation"; @@ -918,7 +1231,7 @@ private static class OperationDiagnosticDescriptors public static readonly DiagnosticDescriptor UnsafeMappedType = new( "CECLUA1106", "Lua operation mapper must project a safe Client result", - "Lua operation '{0}' maps a raw Lua, ownership, pointer, ref-like, or borrowed Lua-class value across the Client boundary", + "Lua operation '{0}' maps an unsafe value across the Client boundary: {1}", "CheatEngine.Client.Lua", DiagnosticSeverity.Error, true); } diff --git a/tests/CheatEngine.Client.SourceGenerators.Lua.Tests/CheatEngineLuaGeneratorTests.cs b/tests/CheatEngine.Client.SourceGenerators.Lua.Tests/CheatEngineLuaGeneratorTests.cs index 3455677..1e97831 100644 --- a/tests/CheatEngine.Client.SourceGenerators.Lua.Tests/CheatEngineLuaGeneratorTests.cs +++ b/tests/CheatEngine.Client.SourceGenerators.Lua.Tests/CheatEngineLuaGeneratorTests.cs @@ -117,7 +117,7 @@ public void ModuleAdapterSnapshotUsesOneAdmittedOperationAndPreflightsEveryExpor Assert.Empty(run.Diagnostics); string generated = run.GeneratedText("PluginLuaModule.CheatEngineLuaModule.g.cs"); Assert.Equal( - """ + NormalizeLineEndings(""" // #nullable enable @@ -135,7 +135,7 @@ public PluginLuaModule() "plugin", global::System.Collections.Immutable.ImmutableArray.Create(new global::CheatEngine.Client.Lua.LuaExportDescriptor("status"), new global::CheatEngine.Client.Lua.LuaExportDescriptor("ping"))); - """ + "\n", + """ + "\n"), NormalizeLineEndings(generated[..generated.IndexOf("\t/// ", StringComparison.Ordinal)])); Assert.Equal(2, Count(generated, "LuaRuntime.AcquireOperation()")); Assert.Contains("EnsureExportsAreVacant(operation.State);", generated, StringComparison.Ordinal); @@ -364,6 +364,144 @@ internal static partial class Globals Assert.Empty(run.GeneratedSources); } + [Fact] + public void MapperRejectsForbiddenTypesNestedInArraysGenericsTuplesAndDtos() + { + string[] sources = + [ + CreateMappedOperationSource("global::CheatEngine.SDK.Lua.References.LuaRef[]"), + CreateMappedOperationSource( + "global::System.Collections.Immutable.ImmutableArray>"), + CreateMappedOperationSource("(int Code, global::CheatEngine.SDK.Engine.Objects.CEObject Handle)"), + CreateMappedOperationSource( + "global::CheatEngine.SDK.Engine.Objects.Owned"), + CreateMappedOperationSource("LeakingSnapshot", + "internal sealed class LeakingSnapshot { private global::CheatEngine.SDK.Lua.References.LuaRef _reference; }"), + CreateMappedOperationSource("UnsafeCallback", + "internal delegate void UnsafeCallback(global::CheatEngine.SDK.Lua.References.LuaRef reference);"), + CreateMappedOperationSource("ConstrainedSnapshot", + "internal interface IUnsafeConstraint { global::CheatEngine.SDK.Lua.References.LuaRef Reference { get; } }" + Environment.NewLine + + "internal sealed class ConstrainedSnapshot where T : IUnsafeConstraint { }") + ]; + + foreach (string source in sources) + { + GeneratorRun run = GeneratorRun.Execute(source); + Diagnostic diagnostic = Assert.Single(run.Diagnostics.Where(static candidate => candidate.Id == "CECLUA1106")); + Assert.Equal("CECLUA1106", diagnostic.Id); + Assert.Empty(run.GeneratedSources); + } + } + + [Fact] + public void RecursiveBoundaryDiagnosticsAreDeterministic() + { + string source = CreateMappedOperationSource("ConstrainedSnapshot", + "internal interface IUnsafeConstraint { global::CheatEngine.SDK.Lua.References.LuaRef Reference { get; } }" + + Environment.NewLine + + "internal sealed class ConstrainedSnapshot where T : IUnsafeConstraint { }"); + + Diagnostic first = Assert.Single(GeneratorRun.Execute(source).Diagnostics + .Where(static candidate => candidate.Id == "CECLUA1106")); + Diagnostic second = Assert.Single(GeneratorRun.Execute(source).Diagnostics + .Where(static candidate => candidate.Id == "CECLUA1106")); + + Assert.Equal(first.GetMessage(System.Globalization.CultureInfo.InvariantCulture), + second.GetMessage(System.Globalization.CultureInfo.InvariantCulture)); + Assert.Contains("forbidden SDK lifetime or interop type", + first.GetMessage(System.Globalization.CultureInfo.InvariantCulture), StringComparison.Ordinal); + } + + [Fact] + public void MapperRejectsForbiddenSourcesAndFunctionPointerResults() + { + GeneratorRun unsafeSource = GeneratorRun.Execute( + CreateMappedOperationSource( + "int", + string.Empty, + "global::CheatEngine.SDK.Lua.References.LuaRef[]")); + Assert.Contains(unsafeSource.Diagnostics, static candidate => candidate.Id == "CECLUA1106"); + Assert.Empty(unsafeSource.GeneratedSources); + + GeneratorRun functionPointer = GeneratorRun.Execute( + """ + using CheatEngine.Client.Lua; + using CheatEngine.SDK.Annotations.Lua; + namespace TestPlugin; + internal static partial class Globals + { + [CheatEngineLuaOperation] + [LuaGlobal("pointer")] + public static unsafe partial delegate* unmanaged[Cdecl] ReadPointer(); + } + """); + Assert.Contains(functionPointer.Diagnostics, static candidate => candidate.Id == "CECLUA1104"); + Assert.Empty(functionPointer.GeneratedSources); + } + + [Fact] + public void ApprovedSdkValuesAndClosedDtosCompileThroughTheGeneratedConsumerBoundary() + { + GeneratorRun run = GeneratorRun.Execute( + """ + using CheatEngine.Client.Lua; + using CheatEngine.SDK.Annotations.Lua; + using CheatEngine.SDK.Engine.Values; + namespace TestPlugin; + public sealed class SdkSnapshot { } + public readonly record struct SafeSnapshot(Address Address, int Version); + internal readonly struct SafeMapper : ILuaResultMapper> + { + public static System.Collections.Immutable.ImmutableArray Map(SdkSnapshot source) => default; + } + + public static partial class Globals + { + [CheatEngineLuaOperation(typeof(SafeMapper))] + [LuaGlobal("snapshot")] + public static partial SdkSnapshot ReadSnapshot(); + } + """); + + Assert.Empty(run.Diagnostics); + Compilation generatedConsumer = run.OutputCompilation.AddSyntaxTrees(CSharpSyntaxTree.ParseText( + """ + namespace TestPlugin; + public static partial class Globals + { + public static partial SdkSnapshot ReadSnapshot() => new(); + } + """, + new CSharpParseOptions(LanguageVersion.CSharp14), + cancellationToken: TestContext.Current.CancellationToken)); + AssertNoCompilerDiagnostics(generatedConsumer); + + using MemoryStream generatedImage = new(); + EmitResult generatedEmit = generatedConsumer.Emit(generatedImage, + cancellationToken: TestContext.Current.CancellationToken); + Assert.True(generatedEmit.Success, string.Join(Environment.NewLine, generatedEmit.Diagnostics)); + + Compilation downstreamConsumer = GeneratorRun.CreateConsumerCompilation( + """ + using System.Collections.Immutable; + using CheatEngine.Client.Lua; + using TestPlugin; + namespace Consumer; + public static class GeneratedApiConsumer + { + public static ILuaOperation> Create() => + Globals.CreateReadSnapshotLuaOperation(); + } + """, + generatedImage.ToArray()); + AssertNoCompilerDiagnostics(downstreamConsumer); + + using MemoryStream downstreamImage = new(); + EmitResult downstreamEmit = downstreamConsumer.Emit(downstreamImage, + cancellationToken: TestContext.Current.CancellationToken); + Assert.True(downstreamEmit.Success, string.Join(Environment.NewLine, downstreamEmit.Diagnostics)); + } + [Fact] public void UnchangedInputReusesTheIncrementalOutput() { @@ -397,4 +535,38 @@ private static string NormalizeLineEndings(string value) { return value.Replace("\r\n", "\n", StringComparison.Ordinal); } + + private static string CreateMappedOperationSource(string resultType, string? additionalDeclarations = null, + string sourceType = "SdkSnapshot") + { + return string.Join( + Environment.NewLine, + [ + "using CheatEngine.Client.Lua;", + "using CheatEngine.SDK.Annotations.Lua;", + "namespace TestPlugin;", + "internal sealed class SdkSnapshot { }", + additionalDeclarations ?? string.Empty, + "internal readonly struct UnsafeMapper : ILuaResultMapper<" + sourceType + ", " + resultType + ">", + "{", + "\tpublic static " + resultType + " Map(" + sourceType + " source) => default;", + "}", + "internal static partial class Globals", + "{", + "\t[CheatEngineLuaOperation(typeof(UnsafeMapper))]", + "\t[LuaGlobal(\"unsafe\")]", + "\tpublic static partial " + sourceType + " GetUnsafe();", + "}" + ]); + } + + private static void AssertNoCompilerDiagnostics(Compilation compilation) + { + Diagnostic[] diagnostics = + [ + .. compilation.GetDiagnostics(TestContext.Current.CancellationToken) + .Where(static diagnostic => diagnostic.Severity >= DiagnosticSeverity.Warning) + ]; + Assert.Empty(diagnostics); + } } diff --git a/tests/CheatEngine.Client.SourceGenerators.Lua.Tests/Infrastructure/GeneratorRun.cs b/tests/CheatEngine.Client.SourceGenerators.Lua.Tests/Infrastructure/GeneratorRun.cs index 10bf1d8..c3c2122 100644 --- a/tests/CheatEngine.Client.SourceGenerators.Lua.Tests/Infrastructure/GeneratorRun.cs +++ b/tests/CheatEngine.Client.SourceGenerators.Lua.Tests/Infrastructure/GeneratorRun.cs @@ -53,7 +53,7 @@ public static GeneratorRun Execute(string source) "CheatEngineClientLuaGeneratorTests", [syntaxTree], GetMetadataReferences(), - new CSharpCompilationOptions(OutputKind.DynamicallyLinkedLibrary)); + new CSharpCompilationOptions(OutputKind.DynamicallyLinkedLibrary, allowUnsafe: true)); GeneratorDriver driver = CSharpGeneratorDriver.Create( [new CheatEngineLuaGenerator().AsSourceGenerator()], parseOptions: parseOptions, @@ -79,6 +79,19 @@ public string GeneratedText(string suffix) return source.SourceText.ToString(); } + public static Compilation CreateConsumerCompilation(string source, byte[] generatedAssembly) + { + ArgumentNullException.ThrowIfNull(source); + ArgumentNullException.ThrowIfNull(generatedAssembly); + + CSharpParseOptions parseOptions = new(LanguageVersion.CSharp14); + return CSharpCompilation.Create( + "CheatEngineClientLuaGeneratedConsumerTests", + [CSharpSyntaxTree.ParseText(SourceText.From(source), parseOptions)], + [.. GetMetadataReferences(), MetadataReference.CreateFromImage(generatedAssembly)], + new CSharpCompilationOptions(OutputKind.DynamicallyLinkedLibrary, allowUnsafe: true)); + } + private static ImmutableArray GetMetadataReferences() { HashSet paths = new(StringComparer.OrdinalIgnoreCase); diff --git a/tests/CheatEngine.Client.Tests/CheatEngine.Client.Tests.csproj b/tests/CheatEngine.Client.Tests/CheatEngine.Client.Tests.csproj index 7ac8505..1247d7f 100644 --- a/tests/CheatEngine.Client.Tests/CheatEngine.Client.Tests.csproj +++ b/tests/CheatEngine.Client.Tests/CheatEngine.Client.Tests.csproj @@ -1,5 +1,9 @@ + + true + + diff --git a/tests/CheatEngine.Client.Tests/PublicClientSignatureBoundaryTests.cs b/tests/CheatEngine.Client.Tests/PublicClientSignatureBoundaryTests.cs index 1bb78b5..aebb012 100644 --- a/tests/CheatEngine.Client.Tests/PublicClientSignatureBoundaryTests.cs +++ b/tests/CheatEngine.Client.Tests/PublicClientSignatureBoundaryTests.cs @@ -3,6 +3,9 @@ using CheatEngine.Client.Extensions.DependencyInjection; using CheatEngine.Client.Hosting; using CheatEngine.Client.Memory; +using CheatEngine.SDK.Engine.Objects; +using CheatEngine.SDK.Lua.References; +using CheatEngine.SDK.Lua.State; using ReflectionAssembly = System.Reflection.Assembly; @@ -54,6 +57,41 @@ public void AllPublicClientSignaturesAreHandleFreeAndUseOnlyApprovedSdkValueType Assert.True(violations.Count == 0, string.Join(Environment.NewLine, violations)); } + [Fact] + public void RecursiveVerifierRejectsNestedSdkHandlesOwnershipDelegatesAndConstraints() + { + AssertViolation(typeof(LuaRef[]), "forbidden SDK handle"); + AssertViolation(typeof(IReadOnlyList>), "forbidden SDK handle"); + AssertViolation(typeof((int Code, CEObject Handle)), "forbidden SDK handle"); + AssertViolation(typeof(Owned<>), "forbidden SDK handle"); + AssertViolation(typeof(LeakingDto), "forbidden SDK handle"); + AssertViolation(typeof(UnsafeCallback), "forbidden SDK handle"); + AssertViolation(typeof(ConstrainedDto<>), "forbidden SDK handle"); + } + + [Fact] + public void RecursiveVerifierRejectsPointersByReferenceAndFunctionPointers() + { + AssertViolation(typeof(int).MakePointerType(), "pointer type"); + AssertViolation(typeof(LuaRef).MakeByRefType(), "forbidden SDK handle"); + + MethodInfo method = typeof(PublicClientSignatureBoundaryTests).GetMethod( + nameof(GetFunctionPointer), + BindingFlags.Static | BindingFlags.NonPublic)!; + AssertViolation(method.ReturnType, "function pointer"); + } + + [Fact] + public void RecursiveVerifierAllowsApprovedSdkValuesInsideSafeContainers() + { + List violations = []; + VerifyType(typeof(CheatEngine.SDK.Engine.Values.Address[]), "approved array", violations); + VerifyType(typeof(IReadOnlyList), "approved generic", violations); + VerifyType(typeof((CheatEngine.SDK.Engine.Values.Address Address, int Version)), "approved tuple", violations); + + Assert.Empty(violations); + } + private static IEnumerable GetAggregateClientAssemblies() { Queue pending = new( @@ -130,21 +168,36 @@ private static void VerifyDeclaredMembers(Type publicType, List violatio private static void VerifyParameters(IEnumerable parameters, MemberInfo member, List violations) + { + VerifyParameters(parameters, member, violations, [], 0); + } + + private static void VerifyParameters(IEnumerable parameters, MemberInfo member, + List violations, HashSet visited, int depth) { foreach (ParameterInfo parameter in parameters) { - VerifyType(parameter.ParameterType, $"{member} parameter '{parameter.Name}'", violations, member); + VerifyType(parameter.ParameterType, $"{member} parameter '{parameter.Name}'", violations, member, visited, + depth + 1); } } private static void VerifyGenericParameterConstraints(IEnumerable genericParameters, string? member, List violations) + { + VerifyGenericParameterConstraints(genericParameters, member, violations, [], 0); + } + + private static void VerifyGenericParameterConstraints(IEnumerable genericParameters, string? member, + List violations, HashSet visited, int depth) { foreach (Type genericParameter in genericParameters.Where(static parameter => parameter.IsGenericParameter)) { - foreach (Type constraint in genericParameter.GetGenericParameterConstraints()) + foreach (Type constraint in genericParameter.GetGenericParameterConstraints() + .OrderBy(static type => type.FullName, StringComparer.Ordinal)) { - VerifyType(constraint, $"{member} generic parameter '{genericParameter.Name}'", violations); + VerifyType(constraint, $"{member} generic parameter '{genericParameter.Name}'", violations, null, visited, + depth + 1); } } } @@ -152,58 +205,173 @@ private static void VerifyGenericParameterConstraints(IEnumerable genericP private static void VerifyType(Type? type, string? source, List violations, MemberInfo? declaringMember = null) { - if (type is null || type.IsGenericParameter) + HashSet visited = []; + VerifyType(type, source, violations, declaringMember, visited, 0); + } + + private static void VerifyType(Type? type, string? source, List violations, MemberInfo? declaringMember, + HashSet visited, int depth) + { + if (type is null || depth > 32 || !visited.Add(type)) { return; } + if (type.IsFunctionPointer) + { + violations.Add($"{source} exposes a function pointer type '{type}'."); + return; + } + if (type.IsPointer) { violations.Add($"{source} exposes a pointer type '{type}'."); return; } - if (type.HasElementType) + if (type.IsByRef) + { + VerifyType(type.GetElementType(), source, violations, declaringMember, visited, depth + 1); + return; + } + + if (type.IsArray) + { + VerifyType(type.GetElementType(), source, violations, declaringMember, visited, depth + 1); + return; + } + + if (type.IsGenericParameter) + { + VerifyGenericParameterConstraints(type.GetGenericParameterConstraints(), source, violations, visited, depth); + return; + } + + if (IsForbiddenSdkType(type, source, violations, declaringMember)) { - VerifyType(type.GetElementType(), source, violations, declaringMember); return; } if (type.IsGenericType) { Type genericDefinition = type.GetGenericTypeDefinition(); - if (genericDefinition != type) + VerifyGenericParameterConstraints(genericDefinition.GetGenericArguments(), source, violations, visited, depth); + foreach (Type argument in type.GetGenericArguments()) { - VerifyType(genericDefinition, source, violations, declaringMember); + VerifyType(argument, source, violations, declaringMember, visited, depth + 1); } + } - foreach (Type argument in type.GetGenericArguments()) + if (typeof(Delegate).IsAssignableFrom(type)) + { + MethodInfo? invoke = type.GetMethod("Invoke", BindingFlags.Public | BindingFlags.Instance); + if (invoke is not null) { - VerifyType(argument, source, violations, declaringMember); + VerifyType(invoke.ReturnType, $"{source} delegate return", violations, invoke, visited, depth + 1); + VerifyParameters(invoke.GetParameters(), invoke, violations, visited, depth); } return; } - string typeName = type.FullName ?? type.Name; - if (type.Name is "LuaState" or "LuaRef" or "CEObject" || - type.Name.StartsWith("Owned`", StringComparison.Ordinal)) + if (!ShouldInspectTypeMembers(type)) { - violations.Add($"{source} exposes forbidden SDK handle '{typeName}'."); return; } - if (type.Namespace?.Contains(".Interop", StringComparison.Ordinal) == true) + VerifyTypeHierarchy(type, source, violations, declaringMember, visited, depth); + VerifyTypeMembers(type, source, violations, visited, depth); + } + + private static bool IsForbiddenSdkType(Type type, string? source, List violations, + MemberInfo? declaringMember) + { + Type definition = type.IsGenericType ? type.GetGenericTypeDefinition() : type; + string typeName = definition.FullName ?? definition.Name; + if (definition.Name is "LuaState" or "LuaRef" or "CEObject" || + definition.Name.StartsWith("Owned`", StringComparison.Ordinal)) + { + violations.Add($"{source} exposes forbidden SDK handle '{typeName}'."); + return true; + } + + if (definition.Namespace?.Contains(".Interop", StringComparison.Ordinal) == true) { violations.Add($"{source} exposes interop namespace type '{typeName}'."); - return; + return true; } - if (type.Assembly.GetName().Name?.StartsWith("CheatEngine.SDK", StringComparison.Ordinal) == true && - (!type.IsValueType || !ApprovedSdkValueTypes.Contains(typeName)) && - !IsShippedRuntimeCapabilitiesDebt(type, declaringMember)) + if (definition.Assembly.GetName().Name?.StartsWith("CheatEngine.SDK", StringComparison.Ordinal) == true && + (!definition.IsValueType || !ApprovedSdkValueTypes.Contains(typeName)) && + !IsShippedRuntimeCapabilitiesDebt(definition, declaringMember)) { violations.Add($"{source} exposes non-approved SDK type '{typeName}'."); + return true; + } + + return false; + } + + private static bool ShouldInspectTypeMembers(Type type) + { + return type.Assembly == typeof(PublicClientSignatureBoundaryTests).Assembly || + type.Assembly.GetName().Name?.StartsWith("CheatEngine.Client", StringComparison.Ordinal) == true; + } + + private static void VerifyTypeHierarchy(Type type, string? source, List violations, + MemberInfo? declaringMember, HashSet visited, int depth) + { + if (type.BaseType is { } baseType && baseType != typeof(object) && !IsRequiredPluginBase(baseType)) + { + VerifyType(baseType, $"{source} base type", violations, declaringMember, visited, depth + 1); + } + + foreach (Type implementedInterface in type.GetInterfaces().OrderBy(static candidate => candidate.FullName, + StringComparer.Ordinal)) + { + VerifyType(implementedInterface, $"{source} interface", violations, declaringMember, visited, depth + 1); + } + } + + private static bool IsRequiredPluginBase(Type type) + { + // Hosting intentionally derives from the SDK plugin bootstrap contract; it is not an SDK owner or raw handle. + return type.FullName == "CheatEngine.SDK.Hosting.Plugin.CheatEnginePlugin"; + } + + private static void VerifyTypeMembers(Type type, string? source, List violations, HashSet visited, + int depth) + { + const BindingFlags DeclaredInstance = BindingFlags.Instance | BindingFlags.Public | BindingFlags.NonPublic | + BindingFlags.DeclaredOnly; + + foreach (FieldInfo field in type.GetFields(DeclaredInstance).OrderBy(static candidate => candidate.Name, + StringComparer.Ordinal)) + { + VerifyType(field.FieldType, $"{source} field '{field.Name}'", violations, field, visited, depth + 1); + } + + foreach (PropertyInfo property in type.GetProperties(DeclaredInstance).OrderBy(static candidate => candidate.Name, + StringComparer.Ordinal)) + { + VerifyType(property.PropertyType, $"{source} property '{property.Name}'", violations, property, visited, + depth + 1); + VerifyParameters(property.GetIndexParameters(), property, violations, visited, depth); + } + + foreach (ConstructorInfo constructor in type.GetConstructors(DeclaredInstance).OrderBy(static candidate => + candidate.ToString(), StringComparer.Ordinal)) + { + VerifyParameters(constructor.GetParameters(), constructor, violations, visited, depth); + } + + foreach (MethodInfo method in type.GetMethods(DeclaredInstance) + .Where(static candidate => !candidate.IsPrivate) + .OrderBy(static candidate => candidate.ToString(), StringComparer.Ordinal)) + { + VerifyType(method.ReturnType, $"{source} method '{method.Name}'", violations, method, visited, depth + 1); + VerifyParameters(method.GetParameters(), method, violations, visited, depth); + VerifyGenericParameterConstraints(method.GetGenericArguments(), method.ToString(), violations, visited, depth); } } @@ -215,10 +383,52 @@ private static bool IsShippedRuntimeCapabilitiesDebt(Type type, MemberInfo? decl declaringMember?.DeclaringType?.FullName == "CheatEngine.Client.Runtime.CheatEngineRuntimeSnapshot" && declaringMember switch { - ConstructorInfo => true, + FieldInfo { Name: "k__BackingField" } => true, + ConstructorInfo => true, MethodInfo { Name: "get_SdkCapabilities" } => true, PropertyInfo { Name: "SdkCapabilities" } => true, _ => false }; } + + private static void AssertViolation(Type type, string expectedFragment) + { + List violations = []; + VerifyType(type, type.FullName, violations); + + Assert.Contains(violations, violation => violation.Contains(expectedFragment, StringComparison.Ordinal)); + } + + private static unsafe delegate* unmanaged[Cdecl] GetFunctionPointer() + { + return null; + } + + private delegate void UnsafeCallback(LuaRef reference); + + private interface IUnsafeConstraint + { + LuaRef Reference + { + get; + } + } + + private sealed class ConstrainedDto + where T : IUnsafeConstraint; + + private sealed class LeakingDto : IDisposable + { + private readonly LuaRef _reference = new(); + + public void Dispose() + { + _reference.Dispose(); + } + + private LuaRef PreserveReference() + { + return _reference; + } + } }