diff --git a/src/System.Linq.Dynamic.Core/Parser/ExpressionHelper.cs b/src/System.Linq.Dynamic.Core/Parser/ExpressionHelper.cs index 29223675..f4c389d7 100644 --- a/src/System.Linq.Dynamic.Core/Parser/ExpressionHelper.cs +++ b/src/System.Linq.Dynamic.Core/Parser/ExpressionHelper.cs @@ -11,6 +11,7 @@ namespace System.Linq.Dynamic.Core.Parser; internal class ExpressionHelper : IExpressionHelper { + private static readonly MethodInfo _containsMethod = typeof(string).GetMethod(nameof(string.Contains), new[] { typeof(string) })!; private static readonly Expression _nullExpression = Expression.Constant(null); private readonly IConstantExpressionWrapper _constantExpressionWrapper = new ConstantExpressionWrapper(); private readonly ParsingConfig _parsingConfig; @@ -129,6 +130,18 @@ public Expression GenerateStringConcat(Expression left, Expression right) return GenerateStaticMethodCall("Concat", left, right); } + public Expression GenerateStringContains(Expression left, Expression right) + { + Expression searchValue = right; + + if (right.Type == typeof(char) && right is ConstantExpression { Value: char character }) + { + searchValue = Expression.Constant(character.ToString(), typeof(string)); + } + + return Expression.Call(left, _containsMethod, searchValue); + } + public Expression GenerateSubtract(Expression left, Expression right) { return Expression.Subtract(left, right); diff --git a/src/System.Linq.Dynamic.Core/Parser/ExpressionParser.cs b/src/System.Linq.Dynamic.Core/Parser/ExpressionParser.cs index 2dd76e97..21987c82 100644 --- a/src/System.Linq.Dynamic.Core/Parser/ExpressionParser.cs +++ b/src/System.Linq.Dynamic.Core/Parser/ExpressionParser.cs @@ -365,14 +365,10 @@ private Expression ParseIn() _textParser.NextToken(); + var expressions = new Dictionary(); + if (_textParser.CurrentToken.Id == TokenId.OpenParen) // literals (or other inline list) { - var values = new List(); - var comparisons = new List(); - Expression? containsLeft = null; - string? containsLeftText = null; - var canUseContains = true; - while (_textParser.CurrentToken.Id != TokenId.CloseParen) { _textParser.NextToken(); @@ -380,46 +376,7 @@ private Expression ParseIn() // we need to parse unary expressions because otherwise 'in' clause will fail in use cases like 'in (-1, -1)' or 'in (!true)' Expression right = ParseUnary(); - // if the identifier is an Enum (or nullable Enum), try to convert the right-side also to an Enum. - if (TypeHelper.GetNonNullableType(left.Type).GetTypeInfo().IsEnum) - { - if (right is ConstantExpression constantExprRight) - { - right = ParseEnumToConstantExpression(token.Pos, left.Type, constantExprRight); - } - else if (_expressionHelper.TryUnwrapAsConstantExpression(right, out var unwrappedConstantExprRight)) - { - right = ParseEnumToConstantExpression(token.Pos, left.Type, unwrappedConstantExprRight); - } - } - - // else, check for direct type match - else if (left.Type != right.Type) - { - CheckAndPromoteOperands(typeof(IEqualitySignatures), TokenId.DoubleEqual, "==", ref left, ref right, token.Pos); - } - - var equalsExpression = _expressionHelper.GenerateEqual(left, right); - comparisons.Add(equalsExpression); - - if (canUseContains && equalsExpression is BinaryExpression binaryExpression && binaryExpression.NodeType == ExpressionType.Equal) - { - containsLeft ??= binaryExpression.Left; - containsLeftText ??= binaryExpression.Left.ToString(); - - if (containsLeft.Type != binaryExpression.Left.Type || !string.Equals(containsLeftText, binaryExpression.Left.ToString(), StringComparison.Ordinal) || binaryExpression.Right.Type != containsLeft.Type) - { - canUseContains = false; - } - else - { - values.Add(binaryExpression.Right); - } - } - else - { - canUseContains = false; - } + expressions.Add(right, token.Pos); if (_textParser.CurrentToken.Id == TokenId.End) { @@ -427,16 +384,7 @@ private Expression ParseIn() } } - if (canUseContains && containsLeft != null) - { - var typeArgs = new[] { containsLeft.Type }; - var args = new Expression[] { Expression.NewArrayInit(containsLeft.Type, values), containsLeft }; - accumulate = Expression.Call(typeof(Enumerable), nameof(Enumerable.Contains), typeArgs, args); - } - else - { - accumulate = _expressionHelper.GenerateBinaryOrElseTree(comparisons); - } + accumulate = ProcessInExpressions(accumulate, expressions); // Since this started with an open paren, make sure to move off the close _textParser.NextToken(); @@ -445,16 +393,35 @@ private Expression ParseIn() { Expression right = ParsePrimary(); - if (!typeof(IEnumerable).IsAssignableFrom(right.Type)) + if (!TypeHelper.TryGetAsEnumerable(right.Type, out _)) { - throw ParseError(_textParser.CurrentToken.Pos, Res.IdentifierImplementingInterfaceExpected, typeof(IEnumerable)); + throw ParseError(_textParser.CurrentToken.Pos, Res.IdentifierImplementingInterfaceExpected, typeof(IEnumerable<>)); } - var typeArgs = new[] { left.Type }; - - var args = new[] { right, left }; + // Handle `it.TestEnum in @0`, and the @0 should be a object like a List. + if (_symbols.Count > 0 && right is ConstantExpression constantExprRight && constantExprRight.Value != null) + { + foreach (var item in (IEnumerable)constantExprRight.Value) + { + expressions.Add(Expression.Constant(item), token.Pos); + } - accumulate = Expression.Call(typeof(Enumerable), nameof(Enumerable.Contains), typeArgs, args); + accumulate = ProcessInExpressions(accumulate, expressions); + } + else + { + // Handle `'y' in Name` and `"x" in Name` where Name is a string + if (right.Type == typeof(string)) + { + accumulate = _expressionHelper.GenerateStringContains(right, left); + } + else + { + var typeArgs = new[] { left.Type }; + var args = new[] { right, left }; + accumulate = Expression.Call(typeof(Enumerable), nameof(Enumerable.Contains), typeArgs, args); + } + } } else { @@ -470,6 +437,71 @@ private Expression ParseIn() return accumulate; } + private Expression ProcessInExpressions(Expression left, Dictionary expressions) + { + var values = new List(); + var comparisons = new List(); + Expression? containsLeft = null; + string? containsLeftText = null; + var canUseContains = true; + + for (int i = 0; i < expressions.Count; i++) + { + var right = expressions.ElementAt(i).Key; + var tokenPos = expressions.ElementAt(i).Value; + + // if the identifier is an Enum (or nullable Enum), try to convert the right-side also to an Enum. + if (TypeHelper.GetNonNullableType(left.Type).GetTypeInfo().IsEnum) + { + if (right is ConstantExpression constantExprRight) + { + right = ParseEnumToConstantExpression(tokenPos, left.Type, constantExprRight); + } + else if (_expressionHelper.TryUnwrapAsConstantExpression(right, out var unwrappedConstantExprRight)) + { + right = ParseEnumToConstantExpression(tokenPos, left.Type, unwrappedConstantExprRight); + } + } + + // else, check for direct type match + else if (left.Type != right.Type) + { + CheckAndPromoteOperands(typeof(IEqualitySignatures), TokenId.DoubleEqual, "==", ref left, ref right, tokenPos); + } + + var equalsExpression = _expressionHelper.GenerateEqual(left, right); + comparisons.Add(equalsExpression); + + if (canUseContains && equalsExpression is BinaryExpression binaryExpression && binaryExpression.NodeType == ExpressionType.Equal) + { + containsLeft ??= binaryExpression.Left; + containsLeftText ??= binaryExpression.Left.ToString(); + + if (containsLeft.Type != binaryExpression.Left.Type || !string.Equals(containsLeftText, binaryExpression.Left.ToString(), StringComparison.Ordinal) || binaryExpression.Right.Type != containsLeft.Type) + { + canUseContains = false; + } + else + { + values.Add(binaryExpression.Right); + } + } + else + { + canUseContains = false; + } + } + + if (canUseContains && containsLeft != null) + { + var typeArgs = new[] { containsLeft.Type }; + var args = new Expression[] { Expression.NewArrayInit(containsLeft.Type, values), containsLeft }; + return Expression.Call(typeof(Enumerable), nameof(Enumerable.Contains), typeArgs, args); + } + + return _expressionHelper.GenerateBinaryOrElseTree(comparisons); + } + // &, | bitwise operators private Expression ParseLogicalAndOrOperator() { @@ -2081,9 +2113,9 @@ private Expression ParseMemberAccess(Type? type, Expression? expression, string? throw ParseError(errorPos, Res.UnknownPropertyOrField, id, TypeHelper.GetTypeName(type)); } - private bool TryFindPropertyOrField(Type type, string id, Expression? expression, [NotNullWhen(true)] out Expression? propertyOrFieldExpression) + private bool TryFindPropertyOrField(Type type, string memberName, Expression? expression, [NotNullWhen(true)] out Expression? propertyOrFieldExpression) { - var member = FindPropertyOrField(type, id, expression == null); + var member = FindPropertyOrField(type, memberName, expression == null); switch (member) { case PropertyInfo property: diff --git a/src/System.Linq.Dynamic.Core/Parser/IExpressionHelper.cs b/src/System.Linq.Dynamic.Core/Parser/IExpressionHelper.cs index fecf495c..ead79ec4 100644 --- a/src/System.Linq.Dynamic.Core/Parser/IExpressionHelper.cs +++ b/src/System.Linq.Dynamic.Core/Parser/IExpressionHelper.cs @@ -26,7 +26,9 @@ internal interface IExpressionHelper Expression GenerateStringConcat(Expression left, Expression right); - Expression GenerateSubtract(Expression left, Expression right); + Expression GenerateStringContains(Expression left, Expression right); + + Expression GenerateSubtract(Expression left, Expression right); void OptimizeForEqualityIfPossible(ref Expression left, ref Expression right); diff --git a/src/System.Linq.Dynamic.Core/Parser/TypeHelper.cs b/src/System.Linq.Dynamic.Core/Parser/TypeHelper.cs index 1cfd533a..fd200930 100644 --- a/src/System.Linq.Dynamic.Core/Parser/TypeHelper.cs +++ b/src/System.Linq.Dynamic.Core/Parser/TypeHelper.cs @@ -13,20 +13,19 @@ internal static bool IsDynamicClass(Type type) internal static bool TryGetAsEnumerable(Type type, [NotNullWhen(true)] out Type? enumerableType) { - if (type.IsArray) - { - enumerableType = typeof(IEnumerable<>).MakeGenericType(type.GetElementType()!); - return true; - } - if (type.GetTypeInfo().IsGenericType && type.GetGenericTypeDefinition() == typeof(IEnumerable<>)) { enumerableType = type; return true; } - enumerableType = null; - return false; + enumerableType = type + .GetInterfaces() + .FirstOrDefault(i => + i.GetTypeInfo().IsGenericType && + i.GetGenericTypeDefinition() == typeof(IEnumerable<>)); + + return enumerableType is not null; } public static bool TryGetFirstGenericArgument(Type type, [NotNullWhen(true)] out Type? genericType) diff --git a/test/System.Linq.Dynamic.Core.Tests/EntitiesTests.In.cs b/test/System.Linq.Dynamic.Core.Tests/EntitiesTests.In.cs index 12e5301c..c48f955a 100644 --- a/test/System.Linq.Dynamic.Core.Tests/EntitiesTests.In.cs +++ b/test/System.Linq.Dynamic.Core.Tests/EntitiesTests.In.cs @@ -1,7 +1,6 @@ using System.Linq.Dynamic.Core.Tests.Helpers.Entities; #if EFCORE -using System.Collections.Generic; using Microsoft.EntityFrameworkCore; #else using System.Data.Entity; @@ -20,10 +19,16 @@ public partial class EntitiesTests public void Entities_Where_In_And() { // Arrange - var expected = _context.Blogs.Include(b => b.Posts).Where(b => new[] { 1000, 1001, 1002 }.Contains(b.BlogId) && new[] { "Blog1", "Blog2" }.Contains(b.Name)).ToArray(); + var expected = _context.Blogs.Include(b => b.Posts) + .Where(b => + new[] { 1000, 1001, 1002 }.Contains(b.BlogId) && new[] { "Blog1", "Blog2" }.Contains(b.Name) && b.Name.Contains("o") && b.Name.Contains("g") + ) + .ToArray(); // Act - var test = _context.Blogs.Include(b => b.Posts).Where(@"BlogId in (1000, 1001, 1002) and Name in (""Blog1"", ""Blog2"")").ToArray(); + var test = _context.Blogs.Include(b => b.Posts) + .Where(@"BlogId in (1000, 1001, 1002) and Name in (""Blog1"", ""Blog2"") && 'o' in Name && Name.Contains(""g"") && ""l"" in Name") + .ToArray(); // Assert Assert.Equal(expected, test); diff --git a/test/System.Linq.Dynamic.Core.Tests/ExpressionTests.cs b/test/System.Linq.Dynamic.Core.Tests/ExpressionTests.cs index 84e5917b..8c8c00c0 100644 --- a/test/System.Linq.Dynamic.Core.Tests/ExpressionTests.cs +++ b/test/System.Linq.Dynamic.Core.Tests/ExpressionTests.cs @@ -1310,10 +1310,22 @@ public void ExpressionTests_In_Enum() var expected = qry.Where(x => new[] { TestEnum.Var1, TestEnum.Var2 }.Contains(x.TestEnum)).ToArray(); var result1 = qry.Where("it.TestEnum in (\"Var1\", \"Var2\")").ToArray(); var result2 = qry.Where("it.TestEnum in (0, 1)").ToArray(); + var result3 = qry.Where("it.TestEnum in @0", new[] { TestEnum.Var1, TestEnum.Var2 }); + var objectList = new List { "Var1", "Var2" }; + var result4 = qry.Where("it.TestEnum in @0", objectList); + var result5 = qry.Where("it.TestEnum in @0", GetVar1AndVar2()); // Assert - Check.That(result1).ContainsExactly(expected); - Check.That(result2).ContainsExactly(expected); + Assert.Equivalent(result1, expected); + Assert.Equivalent(result2, expected); + Assert.Equivalent(result3, expected); + Assert.Equivalent(result4, expected); + Assert.Equivalent(result5, expected); + } + + private static List GetVar1AndVar2() + { + return new List { "Var1", "Var" + "2" }; } [Fact] @@ -1330,10 +1342,18 @@ public void ExpressionTests_In_EnumIsNullable() var expected = new[] { model1, model2 }; var result1 = qry.Where("it.TestEnumNullable in (\"Var1\", \"Var2\")").ToArray(); var result2 = qry.Where("it.TestEnumNullable in (0, 1)").ToArray(); + var result3 = qry.Where("it.TestEnumNullable in @0", new[] { TestEnum.Var1, TestEnum.Var2 }); + + var objectList = new List { "Var1", "Var2" }; + var result4 = qry.Where("it.TestEnumNullable in @0", objectList); + var result5 = qry.Where("it.TestEnumNullable in @0", GetVar1AndVar2()); // Assert - Check.That(result1).ContainsExactly(expected); - Check.That(result2).ContainsExactly(expected); + Assert.Equivalent(result1, expected); + Assert.Equivalent(result2, expected); + Assert.Equivalent(result3, expected); + Assert.Equivalent(result4, expected); + Assert.Equivalent(result5, expected); } [Fact] diff --git a/test/System.Linq.Dynamic.Core.Tests/Parser/ExpressionParserTests.cs b/test/System.Linq.Dynamic.Core.Tests/Parser/ExpressionParserTests.cs index 8e826518..987a533b 100644 --- a/test/System.Linq.Dynamic.Core.Tests/Parser/ExpressionParserTests.cs +++ b/test/System.Linq.Dynamic.Core.Tests/Parser/ExpressionParserTests.cs @@ -266,13 +266,14 @@ public void Parse_ParseMultipleInOperators() { // Arrange ParameterExpression[] parameters = [ParameterExpressionHelper.CreateParameterExpression(typeof(Company), "x")]; - var sut = new ExpressionParser(parameters, "MainCompanyId in (1, 2) and Name in (\"A\", \"B\") && 'y' in Name && 'z' in Name", null, null); + var values = new object?[] { new List { 42, 43 }, new List { 100, 100 + 1 } }; + var sut = new ExpressionParser(parameters, "MainCompanyId in (1, 2) and Name in (\"A\", \"B\") && 'y' in Name && \"z\" in Name and MainCompanyId in @0 and MainCompanyId in @1", values, null); // Act var parsedExpression = sut.Parse(null).ToString(); // Assert - Check.That(parsedExpression).Equals("(((new [] {1, 2}.Contains(x.MainCompanyId) AndAlso new [] {\"A\", \"B\"}.Contains(x.Name)) AndAlso x.Name.Contains(y)) AndAlso x.Name.Contains(z))"); + Assert.Equal("(((((new [] {1, 2}.Contains(x.MainCompanyId) AndAlso new [] {\"A\", \"B\"}.Contains(x.Name)) AndAlso x.Name.Contains(\"y\")) AndAlso x.Name.Contains(\"z\")) AndAlso new [] {42, 43}.Contains(x.MainCompanyId)) AndAlso new [] {100, 101}.Contains(x.MainCompanyId))", parsedExpression); } [Fact] @@ -280,16 +281,15 @@ public void Parse_ParseMultipleInAndNotInOperators() { // Arrange ParameterExpression[] parameters = [ParameterExpressionHelper.CreateParameterExpression(typeof(Company), "x")]; - var sut = new ExpressionParser(parameters, "MainCompanyId in (1, 2) and Name not in (\"A\", \"B\") && 'y' in Name && 'z' not in Name", null, null); + var sut = new ExpressionParser(parameters, "MainCompanyId in (1, 2) and Name not in (\"A\", \"B\") && 'y' in Name && \"z\" not in Name", null, null); // Act var parsedExpression = sut.Parse(null).ToString(); // Assert - Check.That(parsedExpression).Equals("(((new [] {1, 2}.Contains(x.MainCompanyId) AndAlso Not(new [] {\"A\", \"B\"}.Contains(x.Name))) AndAlso x.Name.Contains(y)) AndAlso Not(x.Name.Contains(z)))"); + Assert.Equal("(((new [] {1, 2}.Contains(x.MainCompanyId) AndAlso Not(new [] {\"A\", \"B\"}.Contains(x.Name))) AndAlso x.Name.Contains(\"y\")) AndAlso Not(x.Name.Contains(\"z\")))", parsedExpression); } - [Fact] public void Parse_In_FallsBackTo_OrElse_When_ContainsCannotBeUsed() { @@ -312,13 +312,13 @@ public void Parse_ParseMultipleInAndNotInAndNot_InOperators() { // Arrange ParameterExpression[] parameters = [ParameterExpressionHelper.CreateParameterExpression(typeof(Company), "x")]; - var sut = new ExpressionParser(parameters, "MainCompanyId in (1, 2) and MainCompanyId not in (3, 4) and Name not_in (\"A\", \"B\") && 'y' in Name && 'z' not in Name && 's' not_in Name", null, null); + var sut = new ExpressionParser(parameters, "MainCompanyId in (1, 2) and MainCompanyId not in (3, 4) and Name not_in (\"A\", \"B\") && 'y' in Name && \"z\" not in Name && 's' not_in Name", null, null); // Act var parsedExpression = sut.Parse(null).ToString(); // Assert - Check.That(parsedExpression).Equals("(((((new [] {1, 2}.Contains(x.MainCompanyId) AndAlso Not(new [] {3, 4}.Contains(x.MainCompanyId))) AndAlso Not(new [] {\"A\", \"B\"}.Contains(x.Name))) AndAlso x.Name.Contains(y)) AndAlso Not(x.Name.Contains(z))) AndAlso Not(x.Name.Contains(s)))"); + Assert.Equal("(((((new [] {1, 2}.Contains(x.MainCompanyId) AndAlso Not(new [] {3, 4}.Contains(x.MainCompanyId))) AndAlso Not(new [] {\"A\", \"B\"}.Contains(x.Name))) AndAlso x.Name.Contains(\"y\")) AndAlso Not(x.Name.Contains(\"z\"))) AndAlso Not(x.Name.Contains(\"s\")))", parsedExpression); } [Fact] @@ -326,13 +326,13 @@ public void Parse_ParseInWrappedInParenthesis() { // Arrange ParameterExpression[] parameters = [ParameterExpressionHelper.CreateParameterExpression(typeof(Company), "x")]; - var sut = new ExpressionParser(parameters, "(MainCompanyId in @0)", [new long?[] { 1, 2 }], null); + var sut = new ExpressionParser(parameters, "(MainCompanyId in @0)", [new long?[] { 1, (long) int.MaxValue + 1 }], null); // Act var parsedExpression = sut.Parse(null).ToString(); // Assert - Check.That(parsedExpression).Equals("value(System.Nullable`1[System.Int64][]).Contains(x.MainCompanyId)"); + Check.That(parsedExpression).Equals("new [] {1, 2147483648}.Contains(x.MainCompanyId)"); } [Fact]