From dbcb24141bc0eaba7603f9b9b48ea54b78d69ed0 Mon Sep 17 00:00:00 2001 From: Koen Date: Sun, 26 Jul 2026 00:25:44 +0000 Subject: [PATCH] proper support for tuples of 8 elements or more --- .../Emitter/ExpressionTreeEmitter.cs | 75 +++++++--- .../GeneratedExpressionRuntimeTests.cs | 141 ++++++++++++++++++ ...upleBinary_Equality_8Elements.verified.txt | 53 +++++++ .../ExpressiveGenerator/TupleTests.cs | 23 +++ 4 files changed, 271 insertions(+), 21 deletions(-) create mode 100644 tests/ExpressiveSharp.Generator.Tests/ExpressiveGenerator/GeneratedExpressionRuntimeTests.cs create mode 100644 tests/ExpressiveSharp.Generator.Tests/ExpressiveGenerator/TupleTests.TupleBinary_Equality_8Elements.verified.txt diff --git a/src/ExpressiveSharp.Generator/Emitter/ExpressionTreeEmitter.cs b/src/ExpressiveSharp.Generator/Emitter/ExpressionTreeEmitter.cs index aea48924..146ec584 100644 --- a/src/ExpressiveSharp.Generator/Emitter/ExpressionTreeEmitter.cs +++ b/src/ExpressiveSharp.Generator/Emitter/ExpressionTreeEmitter.cs @@ -1536,35 +1536,25 @@ private string EmitTupleBinary(ITupleBinaryOperation tupleBinary) if (leftType is null || rightType is null) return EmitUnsupported(tupleBinary); - var leftUnderlying = leftType.TupleUnderlyingType ?? leftType; - var leftFields = leftUnderlying.GetMembers() - .OfType() - .Where(f => f.Name.StartsWith("Item")) - .OrderBy(f => f.Name) - .ToList(); + // Normalize away element names so the emitted typeof(...) matches the runtime layout. + var leftTuple = leftType.TupleUnderlyingType ?? leftType; + var rightTuple = rightType.TupleUnderlyingType ?? rightType; - var rightUnderlying = rightType.TupleUnderlyingType ?? rightType; - var rightFields = rightUnderlying.GetMembers() - .OfType() - .Where(f => f.Name.StartsWith("Item")) - .OrderBy(f => f.Name) - .ToList(); + var leftElements = leftTuple.TupleElements; + var rightElements = rightTuple.TupleElements; - if (leftFields.Count == 0 || leftFields.Count != rightFields.Count) + if (leftElements.IsDefaultOrEmpty || rightElements.IsDefaultOrEmpty + || leftElements.Length != rightElements.Length) return EmitUnsupported(tupleBinary); bool isEquality = tupleBinary.OperatorKind == BinaryOperatorKind.Equals; + var restVars = new Dictionary<(string Root, int Level), string>(); var comparisons = new List(); - for (var i = 0; i < leftFields.Count; i++) + for (var i = 0; i < leftElements.Length; i++) { - var leftFieldRef = _fieldCache.EnsureFieldInfo(leftFields[i]); - var rightFieldRef = _fieldCache.EnsureFieldInfo(rightFields[i]); - - var lAccess = NextVar(); - AppendLine($"var {lAccess} = {Expr}.Field({leftVar}, {leftFieldRef});"); - var rAccess = NextVar(); - AppendLine($"var {rAccess} = {Expr}.Field({rightVar}, {rightFieldRef});"); + var lAccess = EmitTupleElementAccess(leftVar, leftTuple, leftElements[i], i, restVars); + var rAccess = EmitTupleElementAccess(rightVar, rightTuple, rightElements[i], i, restVars); var cmpVar = NextVar(); AppendLine($"var {cmpVar} = {Expr}.Equal({lAccess}, {rAccess});"); @@ -1588,6 +1578,49 @@ private string EmitTupleBinary(ITupleBinaryOperation tupleBinary) } return resultVar; + + string EmitTupleElementAccess(string tupleVar, INamedTypeSymbol tupleType, IFieldSymbol element, int elementIndex, Dictionary<(string Root, int Level), string> restVars) + { + const int primaryTupleSize = 7; // ValueTuple has 7 direct fields; the 8th is Rest. + + if (elementIndex < primaryTupleSize) + { + var fieldRef = _fieldCache.EnsureFieldInfo(element); + var directVar = NextVar(); + AppendLine($"var {directVar} = {Expr}.Field({tupleVar}, {fieldRef});"); + return directVar; + } + + const string flags = "global::System.Reflection.BindingFlags.Public | global::System.Reflection.BindingFlags.NonPublic | global::System.Reflection.BindingFlags.Instance"; + + var currentVar = tupleVar; + var currentType = tupleType.TupleUnderlyingType ?? tupleType; + var level = 0; + while (elementIndex >= primaryTupleSize) + { + level++; + var restFqn = currentType.ToDisplayString(_fqnFormat); + currentType = (INamedTypeSymbol)currentType.TypeArguments[7]; + currentType = currentType.TupleUnderlyingType ?? currentType; + elementIndex -= primaryTupleSize; + + if (restVars.TryGetValue((tupleVar, level), out var cachedRestVar)) + { + currentVar = cachedRestVar; + continue; + } + + var restVar = NextVar(); + AppendLine($"var {restVar} = {Expr}.Field({currentVar}, typeof({restFqn}).GetField(\"Rest\", {flags}));"); + restVars[(tupleVar, level)] = restVar; + currentVar = restVar; + } + + var itemFqn = currentType.ToDisplayString(_fqnFormat); + var accessVar = NextVar(); + AppendLine($"var {accessVar} = {Expr}.Field({currentVar}, typeof({itemFqn}).GetField(\"Item{elementIndex + 1}\", {flags}));"); + return accessVar; + } } private string EmitIsPattern(IIsPatternOperation isPattern) diff --git a/tests/ExpressiveSharp.Generator.Tests/ExpressiveGenerator/GeneratedExpressionRuntimeTests.cs b/tests/ExpressiveSharp.Generator.Tests/ExpressiveGenerator/GeneratedExpressionRuntimeTests.cs new file mode 100644 index 00000000..f721c7f8 --- /dev/null +++ b/tests/ExpressiveSharp.Generator.Tests/ExpressiveGenerator/GeneratedExpressionRuntimeTests.cs @@ -0,0 +1,141 @@ +using System.Linq.Expressions; +using System.Reflection; +using Microsoft.CodeAnalysis; +using Microsoft.CodeAnalysis.CSharp; +using Microsoft.VisualStudio.TestTools.UnitTesting; +using ExpressiveSharp.Generator.Tests.Infrastructure; + +namespace ExpressiveSharp.Generator.Tests.ExpressiveGenerator; + +[TestClass] +public class GeneratedExpressionRuntimeTests : GeneratorTestBase +{ + [TestMethod] + public void PolymorphicDispatch_DerivedExpressiveWithoutGeneratedBody_DoesNotBreakBaseExpansion() + { + var baseSource = """ + namespace PolyProofBase + { + public class Animal + { + [ExpressiveSharp.Expressive] + public virtual string Label => "animal"; + } + } + """; + var baseCompilation = CSharpCompilation.Create( + "RuntimeProof.PolyBase", + new[] { CSharpSyntaxTree.ParseText(baseSource) }, + GetDefaultReferences(), + new CSharpCompilationOptions(OutputKind.DynamicallyLinkedLibrary)); + + CSharpGeneratorDriver + .Create(new global::ExpressiveSharp.Generator.ExpressiveGenerator()) + .RunGeneratorsAndUpdateCompilation(baseCompilation, out var baseWithGenerated, out _); + var baseBytes = EmitOrFail(baseWithGenerated); + + var derivedSource = """ + namespace PolyProofDerived + { + public class Dog : PolyProofBase.Animal + { + [ExpressiveSharp.Expressive] + public override string Label => "dog"; + } + } + """; + var derivedCompilation = CSharpCompilation.Create( + "RuntimeProof.PolyDerived", + new[] { CSharpSyntaxTree.ParseText(derivedSource) }, + GetDefaultReferences().Append(MetadataReference.CreateFromImage(baseBytes)), + new CSharpCompilationOptions(OutputKind.DynamicallyLinkedLibrary)); + var derivedBytes = EmitOrFail(derivedCompilation); + + var baseAssembly = Assembly.Load(baseBytes); + ResolveEventHandler resolveHandler = (_, args) => + args.Name.StartsWith("RuntimeProof.PolyBase", StringComparison.Ordinal) ? baseAssembly : null; + AppDomain.CurrentDomain.AssemblyResolve += resolveHandler; + try + { + var derivedAssembly = Assembly.Load(derivedBytes); + + var animalType = baseAssembly.GetType("PolyProofBase.Animal"); + Assert.IsNotNull(animalType); + var dogType = derivedAssembly.GetType("PolyProofDerived.Dog"); + Assert.IsNotNull(dogType); + Assert.IsTrue(animalType.IsAssignableFrom(dogType), + "Precondition: Dog must load and derive from Animal so the polymorphic scan sees it."); + + var parameter = Expression.Parameter(animalType, "a"); + var lambda = Expression.Lambda(Expression.Property(parameter, "Label"), parameter); + + var expanded = global::ExpressiveSharp.ExpressionExtensions.ExpandExpressives(lambda); + Assert.IsNotNull(expanded); + } + finally + { + AppDomain.CurrentDomain.AssemblyResolve -= resolveHandler; + } + } + + [TestMethod] + public void EightElementTupleEquality_ComparesAllElements() + { + var source = """ + namespace TupleProof + { + public static class Fx + { + [ExpressiveSharp.Expressive] + public static bool TupleEquals8(int a, int b) + => (1, 2, 3, 4, 5, 6, 7, a) == (1, 2, 3, 4, 5, 6, 7, b); + + [ExpressiveSharp.Expressive] + public static bool TupleEquals15(int a, int b) + => (1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, a) + == (1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, b); + } + } + """; + var compilation = CSharpCompilation.Create( + "RuntimeProof.Tuple8", + new[] { CSharpSyntaxTree.ParseText(source) }, + GetDefaultReferences(), + new CSharpCompilationOptions(OutputKind.DynamicallyLinkedLibrary)); + + CSharpGeneratorDriver + .Create(new global::ExpressiveSharp.Generator.ExpressiveGenerator()) + .RunGeneratorsAndUpdateCompilation(compilation, out var withGenerated, out _); + var assembly = Assembly.Load(EmitOrFail(withGenerated)); + + var equals8 = CompileExpanded(assembly, "TupleEquals8"); + Assert.IsFalse(equals8(8, 9), "Tuples differing only in the 8th element must not compare equal."); + Assert.IsTrue(equals8(8, 8)); + + var equals15 = CompileExpanded(assembly, "TupleEquals15"); + Assert.IsFalse(equals15(15, 16), "Tuples differing only in the 15th element must not compare equal."); + Assert.IsTrue(equals15(15, 15)); + } + + private static Func CompileExpanded(Assembly assembly, string methodName) + { + var method = assembly.GetType("TupleProof.Fx")!.GetMethod(methodName)!; + var x = Expression.Parameter(typeof(int), "x"); + var y = Expression.Parameter(typeof(int), "y"); + var lambda = Expression.Lambda>(Expression.Call(method, x, y), x, y); + + var expanded = (Expression>) + global::ExpressiveSharp.ExpressionExtensions.ExpandExpressives(lambda); + return expanded.Compile(); + } + + private static byte[] EmitOrFail(Compilation compilation) + { + using var stream = new MemoryStream(); + var emitResult = compilation.Emit(stream); + Assert.IsTrue(emitResult.Success, + "Fixture compilation must succeed:\n" + string.Join("\n", + emitResult.Diagnostics.Where(d => d.Severity == DiagnosticSeverity.Error))); + return stream.ToArray(); + } +} diff --git a/tests/ExpressiveSharp.Generator.Tests/ExpressiveGenerator/TupleTests.TupleBinary_Equality_8Elements.verified.txt b/tests/ExpressiveSharp.Generator.Tests/ExpressiveGenerator/TupleTests.TupleBinary_Equality_8Elements.verified.txt new file mode 100644 index 00000000..54ac54a8 --- /dev/null +++ b/tests/ExpressiveSharp.Generator.Tests/ExpressiveGenerator/TupleTests.TupleBinary_Equality_8Elements.verified.txt @@ -0,0 +1,53 @@ +// +#nullable disable + +using Foo; + +namespace ExpressiveSharp.Generated +{ + static partial class Foo_C + { + // [Expressive] + // public bool Same => A == B; + static global::System.Linq.Expressions.Expression> Same_Expression() + { + var p__this = global::System.Linq.Expressions.Expression.Parameter(typeof(global::Foo.C), "@this"); + var expr_0 = global::System.Linq.Expressions.Expression.Property(p__this, typeof(global::Foo.C).GetProperty("A", global::System.Reflection.BindingFlags.Public | global::System.Reflection.BindingFlags.NonPublic | global::System.Reflection.BindingFlags.Instance)); // A + var expr_1 = global::System.Linq.Expressions.Expression.Property(p__this, typeof(global::Foo.C).GetProperty("B", global::System.Reflection.BindingFlags.Public | global::System.Reflection.BindingFlags.NonPublic | global::System.Reflection.BindingFlags.Instance)); // B + var expr_2 = global::System.Linq.Expressions.Expression.Field(expr_0, typeof((int, int, int, int, int, int, int, int)).GetField("Item1", global::System.Reflection.BindingFlags.Public | global::System.Reflection.BindingFlags.NonPublic | global::System.Reflection.BindingFlags.Instance)); + var expr_3 = global::System.Linq.Expressions.Expression.Field(expr_1, typeof((int, int, int, int, int, int, int, int)).GetField("Item1", global::System.Reflection.BindingFlags.Public | global::System.Reflection.BindingFlags.NonPublic | global::System.Reflection.BindingFlags.Instance)); + var expr_4 = global::System.Linq.Expressions.Expression.Equal(expr_2, expr_3); + var expr_5 = global::System.Linq.Expressions.Expression.Field(expr_0, typeof((int, int, int, int, int, int, int, int)).GetField("Item2", global::System.Reflection.BindingFlags.Public | global::System.Reflection.BindingFlags.NonPublic | global::System.Reflection.BindingFlags.Instance)); + var expr_6 = global::System.Linq.Expressions.Expression.Field(expr_1, typeof((int, int, int, int, int, int, int, int)).GetField("Item2", global::System.Reflection.BindingFlags.Public | global::System.Reflection.BindingFlags.NonPublic | global::System.Reflection.BindingFlags.Instance)); + var expr_7 = global::System.Linq.Expressions.Expression.Equal(expr_5, expr_6); + var expr_8 = global::System.Linq.Expressions.Expression.Field(expr_0, typeof((int, int, int, int, int, int, int, int)).GetField("Item3", global::System.Reflection.BindingFlags.Public | global::System.Reflection.BindingFlags.NonPublic | global::System.Reflection.BindingFlags.Instance)); + var expr_9 = global::System.Linq.Expressions.Expression.Field(expr_1, typeof((int, int, int, int, int, int, int, int)).GetField("Item3", global::System.Reflection.BindingFlags.Public | global::System.Reflection.BindingFlags.NonPublic | global::System.Reflection.BindingFlags.Instance)); + var expr_10 = global::System.Linq.Expressions.Expression.Equal(expr_8, expr_9); + var expr_11 = global::System.Linq.Expressions.Expression.Field(expr_0, typeof((int, int, int, int, int, int, int, int)).GetField("Item4", global::System.Reflection.BindingFlags.Public | global::System.Reflection.BindingFlags.NonPublic | global::System.Reflection.BindingFlags.Instance)); + var expr_12 = global::System.Linq.Expressions.Expression.Field(expr_1, typeof((int, int, int, int, int, int, int, int)).GetField("Item4", global::System.Reflection.BindingFlags.Public | global::System.Reflection.BindingFlags.NonPublic | global::System.Reflection.BindingFlags.Instance)); + var expr_13 = global::System.Linq.Expressions.Expression.Equal(expr_11, expr_12); + var expr_14 = global::System.Linq.Expressions.Expression.Field(expr_0, typeof((int, int, int, int, int, int, int, int)).GetField("Item5", global::System.Reflection.BindingFlags.Public | global::System.Reflection.BindingFlags.NonPublic | global::System.Reflection.BindingFlags.Instance)); + var expr_15 = global::System.Linq.Expressions.Expression.Field(expr_1, typeof((int, int, int, int, int, int, int, int)).GetField("Item5", global::System.Reflection.BindingFlags.Public | global::System.Reflection.BindingFlags.NonPublic | global::System.Reflection.BindingFlags.Instance)); + var expr_16 = global::System.Linq.Expressions.Expression.Equal(expr_14, expr_15); + var expr_17 = global::System.Linq.Expressions.Expression.Field(expr_0, typeof((int, int, int, int, int, int, int, int)).GetField("Item6", global::System.Reflection.BindingFlags.Public | global::System.Reflection.BindingFlags.NonPublic | global::System.Reflection.BindingFlags.Instance)); + var expr_18 = global::System.Linq.Expressions.Expression.Field(expr_1, typeof((int, int, int, int, int, int, int, int)).GetField("Item6", global::System.Reflection.BindingFlags.Public | global::System.Reflection.BindingFlags.NonPublic | global::System.Reflection.BindingFlags.Instance)); + var expr_19 = global::System.Linq.Expressions.Expression.Equal(expr_17, expr_18); + var expr_20 = global::System.Linq.Expressions.Expression.Field(expr_0, typeof((int, int, int, int, int, int, int, int)).GetField("Item7", global::System.Reflection.BindingFlags.Public | global::System.Reflection.BindingFlags.NonPublic | global::System.Reflection.BindingFlags.Instance)); + var expr_21 = global::System.Linq.Expressions.Expression.Field(expr_1, typeof((int, int, int, int, int, int, int, int)).GetField("Item7", global::System.Reflection.BindingFlags.Public | global::System.Reflection.BindingFlags.NonPublic | global::System.Reflection.BindingFlags.Instance)); + var expr_22 = global::System.Linq.Expressions.Expression.Equal(expr_20, expr_21); + var expr_23 = global::System.Linq.Expressions.Expression.Field(expr_0, typeof((int, int, int, int, int, int, int, int)).GetField("Rest", global::System.Reflection.BindingFlags.Public | global::System.Reflection.BindingFlags.NonPublic | global::System.Reflection.BindingFlags.Instance)); + var expr_24 = global::System.Linq.Expressions.Expression.Field(expr_23, typeof(global::System.ValueTuple).GetField("Item1", global::System.Reflection.BindingFlags.Public | global::System.Reflection.BindingFlags.NonPublic | global::System.Reflection.BindingFlags.Instance)); + var expr_25 = global::System.Linq.Expressions.Expression.Field(expr_1, typeof((int, int, int, int, int, int, int, int)).GetField("Rest", global::System.Reflection.BindingFlags.Public | global::System.Reflection.BindingFlags.NonPublic | global::System.Reflection.BindingFlags.Instance)); + var expr_26 = global::System.Linq.Expressions.Expression.Field(expr_25, typeof(global::System.ValueTuple).GetField("Item1", global::System.Reflection.BindingFlags.Public | global::System.Reflection.BindingFlags.NonPublic | global::System.Reflection.BindingFlags.Instance)); + var expr_27 = global::System.Linq.Expressions.Expression.Equal(expr_24, expr_26); + var expr_28 = global::System.Linq.Expressions.Expression.AndAlso(expr_4, expr_7); + var expr_29 = global::System.Linq.Expressions.Expression.AndAlso(expr_28, expr_10); + var expr_30 = global::System.Linq.Expressions.Expression.AndAlso(expr_29, expr_13); + var expr_31 = global::System.Linq.Expressions.Expression.AndAlso(expr_30, expr_16); + var expr_32 = global::System.Linq.Expressions.Expression.AndAlso(expr_31, expr_19); + var expr_33 = global::System.Linq.Expressions.Expression.AndAlso(expr_32, expr_22); + var expr_34 = global::System.Linq.Expressions.Expression.AndAlso(expr_33, expr_27); + return global::System.Linq.Expressions.Expression.Lambda>(expr_34, p__this); + } + } +} diff --git a/tests/ExpressiveSharp.Generator.Tests/ExpressiveGenerator/TupleTests.cs b/tests/ExpressiveSharp.Generator.Tests/ExpressiveGenerator/TupleTests.cs index 337a2f81..9a9551c8 100644 --- a/tests/ExpressiveSharp.Generator.Tests/ExpressiveGenerator/TupleTests.cs +++ b/tests/ExpressiveSharp.Generator.Tests/ExpressiveGenerator/TupleTests.cs @@ -114,6 +114,29 @@ class C { return Verifier.Verify(result.GeneratedTrees[0].ToString()); } + [TestMethod] + public Task TupleBinary_Equality_8Elements() + { + var compilation = CreateCompilation( + """ + namespace Foo { + class C { + public (int, int, int, int, int, int, int, int) A { get; set; } + public (int, int, int, int, int, int, int, int) B { get; set; } + + [Expressive] + public bool Same => A == B; + } + } + """); + var result = RunExpressiveGenerator(compilation); + + Assert.AreEqual(0, result.Diagnostics.Length); + Assert.AreEqual(1, result.GeneratedTrees.Length); + + return Verifier.Verify(result.GeneratedTrees[0].ToString()); + } + [TestMethod] public Task TupleBinary_Equality() {