Skip to content
Open
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
75 changes: 54 additions & 21 deletions src/ExpressiveSharp.Generator/Emitter/ExpressionTreeEmitter.cs
Original file line number Diff line number Diff line change
Expand Up @@ -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<IFieldSymbol>()
.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<IFieldSymbol>()
.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<string>();
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});");
Expand All @@ -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<T1..T7, TRest> 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)
Expand Down
Original file line number Diff line number Diff line change
@@ -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<int, int, bool> 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<Func<int, int, bool>>(Expression.Call(method, x, y), x, y);

var expanded = (Expression<Func<int, int, bool>>)
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();
}
}
Original file line number Diff line number Diff line change
@@ -0,0 +1,53 @@
// <auto-generated/>
#nullable disable

using Foo;

namespace ExpressiveSharp.Generated
{
static partial class Foo_C
{
// [Expressive]
// public bool Same => A == B;
static global::System.Linq.Expressions.Expression<global::System.Func<global::Foo.C, bool>> 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<int>).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<int>).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<global::System.Func<global::Foo.C, bool>>(expr_34, p__this);
}
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -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()
{
Expand Down
Loading