diff --git a/src/EFCore.PG/Storage/Internal/Mapping/NpgsqlArrayTypeMapping.cs b/src/EFCore.PG/Storage/Internal/Mapping/NpgsqlArrayTypeMapping.cs index 749e3d3a0..07a4328b9 100644 --- a/src/EFCore.PG/Storage/Internal/Mapping/NpgsqlArrayTypeMapping.cs +++ b/src/EFCore.PG/Storage/Internal/Mapping/NpgsqlArrayTypeMapping.cs @@ -181,6 +181,50 @@ private static RelationalTypeMappingParameters CreateParameters(string storeType storeType); } + private IEnumerable AsElementEnumerable(object value) + => value switch + { + IEnumerable elements => elements, + IEnumerable elements => elements.Cast(), + _ => throw new InvalidOperationException( + $"Cannot create a parameter for {GetType().Name} from value of type '{value.GetType().Name}'") + }; + + private static TConcreteCollection CreateInstance(int? count) + => (count, typeof(TConcreteCollection)) switch + { + ({ } c, var type) when type.GetConstructor([typeof(int)]) is { } ctorWithSize + => (TConcreteCollection)ctorWithSize.Invoke([c]), + var (_, type) when type.GetConstructor([]) is { } ctor + => (TConcreteCollection)ctor.Invoke(null), + var (_, type) => throw new InvalidOperationException( + $"Type {type.Name} cannot be instantiated as it does not have a public parameterless constructor") + }; + + private static object Materialize(IEnumerable elements) + { + if (typeof(TConcreteCollection).IsArray) + { + return elements.ToArray(); + } + + var count = elements.TryGetNonEnumeratedCount(out var c) ? c : (int?)null; + var collection = CreateInstance(count); + + if (collection is not ICollection destination) + { + throw new InvalidOperationException( + $"Type {typeof(TConcreteCollection).Name} cannot be populated (no ICollection<{typeof(TElement).Name}>)."); + } + + foreach (var element in elements) + { + destination.Add(element); + } + + return collection; + } + /// /// This is an internal API that supports the Entity Framework Core infrastructure and not subject to /// the same compatibility standards as public APIs. It may be changed or removed without notice in @@ -233,21 +277,18 @@ public override DbParameter CreateParameter( // In queries which compose non-server-correlated LINQ operators over an array parameter (e.g. Where(b => ids.Skip(1)...) we // get an enumerable parameter value that isn't an array/list - but those aren't supported at the Npgsql ADO level. // Detect this here and evaluate the enumerable to get a fully materialized List. - // Note that when we have a value converter (e.g. for HashSet), we don't want to convert it to a List, since the value converter - // expects the original type. + // Note that when we have a value converter (e.g. for HashSet), we don't want to convert values that already match + // the converter's model type, since the value converter expects that original type. + // However, if the value's collection shape differs from the converter model type (e.g. List vs T[] after + // type mapping inference for Intersect().Any() → &&), normalize to TConcreteCollection so Sanitize succeeds. // TODO: Make Npgsql support IList<> instead of only arrays and List<> if (value is not null && Converter is null && !value.GetType().IsArrayOrGenericList()) { - switch (value) - { - case IEnumerable elements: - value = elements.ToList(); - break; - - case IEnumerable elements: - value = elements.Cast().ToList(); - break; - } + value = AsElementEnumerable(value).ToList(); + } + else if (value is not null && Converter is not null && !Converter.ModelClrType.IsInstanceOfType(value)) + { + value = Materialize(AsElementEnumerable(value)); } var param = base.CreateParameter(command, name, value, nullable, direction); diff --git a/test/EFCore.PG.FunctionalTests/Query/ArrayArrayQueryTest.cs b/test/EFCore.PG.FunctionalTests/Query/ArrayArrayQueryTest.cs index 083d4fa22..99c8d0be6 100644 --- a/test/EFCore.PG.FunctionalTests/Query/ArrayArrayQueryTest.cs +++ b/test/EFCore.PG.FunctionalTests/Query/ArrayArrayQueryTest.cs @@ -1,3 +1,4 @@ +using System.Collections.Immutable; using Microsoft.EntityFrameworkCore.TestModels.Array; using Npgsql.EntityFrameworkCore.PostgreSQL.Internal; @@ -867,6 +868,110 @@ public virtual async Task All_Contains() #endregion Any/All + #region Intersect + + [ConditionalFact] + public virtual async Task Intersect_parameter_list_over_value_converted_array() + { + List toFindList = [SomeEnum.One, SomeEnum.Three, SomeEnum.Eight]; + + await AssertQuery(ss => ss.Set().Where(e => e.ValueConvertedArrayOfEnum.Intersect(toFindList).Any())); + + AssertSql( + """ +@toFindList={ 'One' +'Three' +'Eight' } (DbType = Object) + +SELECT s."Id", s."ArrayContainerEntityId", s."ArrayOfStringConvertedToDelimitedString", s."Byte", s."ByteArray", s."Bytea", s."EnumConvertedToInt", s."EnumConvertedToString", s."IList", s."IntArray", s."IntList", s."ListOfStringConvertedToDelimitedString", s."NonNullableText", s."NullableEnumConvertedToString", s."NullableEnumConvertedToStringWithNonNullableLambda", s."NullableIntArray", s."NullableIntList", s."NullableStringArray", s."NullableStringList", s."NullableText", s."StringArray", s."StringList", s."ValueConvertedArrayOfEnum", s."ValueConvertedListOfEnum", s."Varchar10", s."Varchar15" +FROM "SomeEntities" AS s +WHERE s."ValueConvertedArrayOfEnum" && @toFindList +"""); + } + + [ConditionalFact] + public virtual async Task Intersect_parameter_immutable_list_over_value_converted_array() + { + ImmutableList toFindImmutableList = [SomeEnum.One, SomeEnum.Three, SomeEnum.Eight]; + + await AssertQuery(ss => ss.Set().Where(e => e.ValueConvertedArrayOfEnum.Intersect(toFindImmutableList).Any())); + + AssertSql( + """ +@toFindImmutableList={ 'One' +'Three' +'Eight' } (DbType = Object) + +SELECT s."Id", s."ArrayContainerEntityId", s."ArrayOfStringConvertedToDelimitedString", s."Byte", s."ByteArray", s."Bytea", s."EnumConvertedToInt", s."EnumConvertedToString", s."IList", s."IntArray", s."IntList", s."ListOfStringConvertedToDelimitedString", s."NonNullableText", s."NullableEnumConvertedToString", s."NullableEnumConvertedToStringWithNonNullableLambda", s."NullableIntArray", s."NullableIntList", s."NullableStringArray", s."NullableStringList", s."NullableText", s."StringArray", s."StringList", s."ValueConvertedArrayOfEnum", s."ValueConvertedListOfEnum", s."Varchar10", s."Varchar15" +FROM "SomeEntities" AS s +WHERE s."ValueConvertedArrayOfEnum" && @toFindImmutableList +"""); + } + + [ConditionalFact] + public virtual async Task Intersect_parameter_array_over_value_converted_array() + { + SomeEnum[] toFindArray = [SomeEnum.One, SomeEnum.Three, SomeEnum.Eight]; + + await AssertQuery(ss => ss.Set().Where(e => e.ValueConvertedArrayOfEnum.Intersect(toFindArray).Any())); + + AssertSql( + """ +@toFindArray={ 'One' +'Three' +'Eight' } (DbType = Object) + +SELECT s."Id", s."ArrayContainerEntityId", s."ArrayOfStringConvertedToDelimitedString", s."Byte", s."ByteArray", s."Bytea", s."EnumConvertedToInt", s."EnumConvertedToString", s."IList", s."IntArray", s."IntList", s."ListOfStringConvertedToDelimitedString", s."NonNullableText", s."NullableEnumConvertedToString", s."NullableEnumConvertedToStringWithNonNullableLambda", s."NullableIntArray", s."NullableIntList", s."NullableStringArray", s."NullableStringList", s."NullableText", s."StringArray", s."StringList", s."ValueConvertedArrayOfEnum", s."ValueConvertedListOfEnum", s."Varchar10", s."Varchar15" +FROM "SomeEntities" AS s +WHERE s."ValueConvertedArrayOfEnum" && @toFindArray +"""); + } + + [ConditionalFact] + public virtual async Task Intersect_parameter_hash_set_over_value_converted_array() + { + HashSet toFindHashSet = [SomeEnum.One, SomeEnum.Three, SomeEnum.Eight]; + + await AssertQuery(ss => ss.Set().Where(e => e.ValueConvertedArrayOfEnum.Intersect(toFindHashSet).Any())); + + AssertSql( + """ +@toFindHashSet={ 'One' +'Three' +'Eight' } (DbType = Object) + +SELECT s."Id", s."ArrayContainerEntityId", s."ArrayOfStringConvertedToDelimitedString", s."Byte", s."ByteArray", s."Bytea", s."EnumConvertedToInt", s."EnumConvertedToString", s."IList", s."IntArray", s."IntList", s."ListOfStringConvertedToDelimitedString", s."NonNullableText", s."NullableEnumConvertedToString", s."NullableEnumConvertedToStringWithNonNullableLambda", s."NullableIntArray", s."NullableIntList", s."NullableStringArray", s."NullableStringList", s."NullableText", s."StringArray", s."StringList", s."ValueConvertedArrayOfEnum", s."ValueConvertedListOfEnum", s."Varchar10", s."Varchar15" +FROM "SomeEntities" AS s +WHERE s."ValueConvertedArrayOfEnum" && @toFindHashSet +"""); + } + + [ConditionalFact] + public virtual async Task Intersect_parameter_non_collection_enumerable_over_value_converted_array() + { + var toFindEnumerable = new List + { + SomeEnum.One, + SomeEnum.Three, + SomeEnum.Eight + }.Where(_ => true); + + await AssertQuery(ss => ss.Set().Where(e => e.ValueConvertedArrayOfEnum.Intersect(toFindEnumerable).Any())); + + AssertSql( + """ +@toFindEnumerable={ 'One' +'Three' +'Eight' } (DbType = Object) + +SELECT s."Id", s."ArrayContainerEntityId", s."ArrayOfStringConvertedToDelimitedString", s."Byte", s."ByteArray", s."Bytea", s."EnumConvertedToInt", s."EnumConvertedToString", s."IList", s."IntArray", s."IntList", s."ListOfStringConvertedToDelimitedString", s."NonNullableText", s."NullableEnumConvertedToString", s."NullableEnumConvertedToStringWithNonNullableLambda", s."NullableIntArray", s."NullableIntList", s."NullableStringArray", s."NullableStringList", s."NullableText", s."StringArray", s."StringList", s."ValueConvertedArrayOfEnum", s."ValueConvertedListOfEnum", s."Varchar10", s."Varchar15" +FROM "SomeEntities" AS s +WHERE s."ValueConvertedArrayOfEnum" && @toFindEnumerable +"""); + } + + #endregion + #region Other translations [ConditionalFact] diff --git a/test/EFCore.PG.FunctionalTests/Query/ArrayListQueryTest.cs b/test/EFCore.PG.FunctionalTests/Query/ArrayListQueryTest.cs index 12dadfd91..c62566d46 100644 --- a/test/EFCore.PG.FunctionalTests/Query/ArrayListQueryTest.cs +++ b/test/EFCore.PG.FunctionalTests/Query/ArrayListQueryTest.cs @@ -1,3 +1,4 @@ +using System.Collections.Immutable; using Microsoft.EntityFrameworkCore.TestModels.Array; namespace Microsoft.EntityFrameworkCore.Query; @@ -873,6 +874,110 @@ public virtual async Task All_Contains() #endregion Any/All + #region Intersect + + [ConditionalFact] + public virtual async Task Intersect_parameter_array_over_value_converted_list() + { + SomeEnum[] toFindArray = [SomeEnum.One, SomeEnum.Three, SomeEnum.Eight]; + + await AssertQuery(ss => ss.Set().Where(e => e.ValueConvertedListOfEnum.Intersect(toFindArray).Any())); + + AssertSql( + """ +@toFindArray={ 'One' +'Three' +'Eight' } (DbType = Object) + +SELECT s."Id", s."ArrayContainerEntityId", s."ArrayOfStringConvertedToDelimitedString", s."Byte", s."ByteArray", s."Bytea", s."EnumConvertedToInt", s."EnumConvertedToString", s."IList", s."IntArray", s."IntList", s."ListOfStringConvertedToDelimitedString", s."NonNullableText", s."NullableEnumConvertedToString", s."NullableEnumConvertedToStringWithNonNullableLambda", s."NullableIntArray", s."NullableIntList", s."NullableStringArray", s."NullableStringList", s."NullableText", s."StringArray", s."StringList", s."ValueConvertedArrayOfEnum", s."ValueConvertedListOfEnum", s."Varchar10", s."Varchar15" +FROM "SomeEntities" AS s +WHERE s."ValueConvertedListOfEnum" && @toFindArray +"""); + } + + [ConditionalFact] + public virtual async Task Intersect_parameter_immutable_list_over_value_converted_list() + { + ImmutableList toFindImmutableList = [SomeEnum.One, SomeEnum.Three, SomeEnum.Eight]; + + await AssertQuery(ss => ss.Set().Where(e => e.ValueConvertedListOfEnum.Intersect(toFindImmutableList).Any())); + + AssertSql( + """ +@toFindImmutableList={ 'One' +'Three' +'Eight' } (DbType = Object) + +SELECT s."Id", s."ArrayContainerEntityId", s."ArrayOfStringConvertedToDelimitedString", s."Byte", s."ByteArray", s."Bytea", s."EnumConvertedToInt", s."EnumConvertedToString", s."IList", s."IntArray", s."IntList", s."ListOfStringConvertedToDelimitedString", s."NonNullableText", s."NullableEnumConvertedToString", s."NullableEnumConvertedToStringWithNonNullableLambda", s."NullableIntArray", s."NullableIntList", s."NullableStringArray", s."NullableStringList", s."NullableText", s."StringArray", s."StringList", s."ValueConvertedArrayOfEnum", s."ValueConvertedListOfEnum", s."Varchar10", s."Varchar15" +FROM "SomeEntities" AS s +WHERE s."ValueConvertedListOfEnum" && @toFindImmutableList +"""); + } + + [ConditionalFact] + public virtual async Task Intersect_parameter_list_over_value_converted_list() + { + List toFindList = [SomeEnum.One, SomeEnum.Three, SomeEnum.Eight]; + + await AssertQuery(ss => ss.Set().Where(e => e.ValueConvertedListOfEnum.Intersect(toFindList).Any())); + + AssertSql( + """ +@toFindList={ 'One' +'Three' +'Eight' } (DbType = Object) + +SELECT s."Id", s."ArrayContainerEntityId", s."ArrayOfStringConvertedToDelimitedString", s."Byte", s."ByteArray", s."Bytea", s."EnumConvertedToInt", s."EnumConvertedToString", s."IList", s."IntArray", s."IntList", s."ListOfStringConvertedToDelimitedString", s."NonNullableText", s."NullableEnumConvertedToString", s."NullableEnumConvertedToStringWithNonNullableLambda", s."NullableIntArray", s."NullableIntList", s."NullableStringArray", s."NullableStringList", s."NullableText", s."StringArray", s."StringList", s."ValueConvertedArrayOfEnum", s."ValueConvertedListOfEnum", s."Varchar10", s."Varchar15" +FROM "SomeEntities" AS s +WHERE s."ValueConvertedListOfEnum" && @toFindList +"""); + } + + [ConditionalFact] + public virtual async Task Intersect_parameter_hash_set_over_value_converted_list() + { + HashSet toFindHashSet = [SomeEnum.One, SomeEnum.Three, SomeEnum.Eight]; + + await AssertQuery(ss => ss.Set().Where(e => e.ValueConvertedListOfEnum.Intersect(toFindHashSet).Any())); + + AssertSql( + """ +@toFindHashSet={ 'One' +'Three' +'Eight' } (DbType = Object) + +SELECT s."Id", s."ArrayContainerEntityId", s."ArrayOfStringConvertedToDelimitedString", s."Byte", s."ByteArray", s."Bytea", s."EnumConvertedToInt", s."EnumConvertedToString", s."IList", s."IntArray", s."IntList", s."ListOfStringConvertedToDelimitedString", s."NonNullableText", s."NullableEnumConvertedToString", s."NullableEnumConvertedToStringWithNonNullableLambda", s."NullableIntArray", s."NullableIntList", s."NullableStringArray", s."NullableStringList", s."NullableText", s."StringArray", s."StringList", s."ValueConvertedArrayOfEnum", s."ValueConvertedListOfEnum", s."Varchar10", s."Varchar15" +FROM "SomeEntities" AS s +WHERE s."ValueConvertedListOfEnum" && @toFindHashSet +"""); + } + + [ConditionalFact] + public virtual async Task Intersect_parameter_non_collection_enumerable_over_value_converted_list() + { + var toFindEnumerable = new List + { + SomeEnum.One, + SomeEnum.Three, + SomeEnum.Eight + }.Where(_ => true); + + await AssertQuery(ss => ss.Set().Where(e => e.ValueConvertedListOfEnum.Intersect(toFindEnumerable).Any())); + + AssertSql( + """ +@toFindEnumerable={ 'One' +'Three' +'Eight' } (DbType = Object) + +SELECT s."Id", s."ArrayContainerEntityId", s."ArrayOfStringConvertedToDelimitedString", s."Byte", s."ByteArray", s."Bytea", s."EnumConvertedToInt", s."EnumConvertedToString", s."IList", s."IntArray", s."IntList", s."ListOfStringConvertedToDelimitedString", s."NonNullableText", s."NullableEnumConvertedToString", s."NullableEnumConvertedToStringWithNonNullableLambda", s."NullableIntArray", s."NullableIntList", s."NullableStringArray", s."NullableStringList", s."NullableText", s."StringArray", s."StringList", s."ValueConvertedArrayOfEnum", s."ValueConvertedListOfEnum", s."Varchar10", s."Varchar15" +FROM "SomeEntities" AS s +WHERE s."ValueConvertedListOfEnum" && @toFindEnumerable +"""); + } + + #endregion + #region Other translations // TODO: https://github.com/dotnet/efcore/issues/30669 diff --git a/test/EFCore.PG.Tests/Storage/NpgsqlTypeMappingSourceTest.cs b/test/EFCore.PG.Tests/Storage/NpgsqlTypeMappingSourceTest.cs index b1ec3c98b..32bb2d247 100644 --- a/test/EFCore.PG.Tests/Storage/NpgsqlTypeMappingSourceTest.cs +++ b/test/EFCore.PG.Tests/Storage/NpgsqlTypeMappingSourceTest.cs @@ -1,3 +1,4 @@ +using System.Collections.Immutable; using System.Net; using System.Net.NetworkInformation; using System.Text.Json; @@ -263,6 +264,82 @@ public void Array_over_type_mapping_with_value_converter_by_clr_type_list() public void Array_over_type_mapping_with_value_converter_by_store_type() => Array_over_type_mapping_with_value_converter(CreateTypeMappingSource().FindMapping("ltree[]"), typeof(List)); + [Theory] + [InlineData(typeof(LTree[]), typeof(List))] + [InlineData(typeof(List), typeof(List))] + [InlineData(typeof(HashSet), typeof(List))] + [InlineData(typeof(LTree[]), typeof(HashSet))] + [InlineData(typeof(List), typeof(HashSet))] + [InlineData(typeof(HashSet), typeof(HashSet))] + public void CreateParameter_with_value_converter_accepts_mutable_collections(Type mappingType, Type valueType) + => CreateParameter_with_value_converter( + mappingType, Activator.CreateInstance(valueType, new LTree[] { new("foo"), new("bar") })); + + [Theory] + [InlineData(typeof(LTree[]))] + [InlineData(typeof(List))] + [InlineData(typeof(HashSet))] + public void CreateParameter_with_value_converter_accepts_immutable_list(Type mappingType) + => CreateParameter_with_value_converter(mappingType, ImmutableList.Create(new LTree("foo"), new LTree("bar"))); + + [Theory] + [InlineData(typeof(LTree[]))] + [InlineData(typeof(List))] + [InlineData(typeof(HashSet))] + public void CreateParameter_with_value_converter_accepts_array(Type mappingType) + => CreateParameter_with_value_converter(mappingType, new LTree[] { new("foo"), new("bar") }); + + [Theory] + [InlineData(typeof(LTree[]))] + [InlineData(typeof(List))] + [InlineData(typeof(HashSet))] + public void CreateParameter_with_value_converter_accepts_non_collection_enumerable(Type mappingType) + => CreateParameter_with_value_converter( + mappingType, new List { new("foo"), new("bar") }.Where(_ => true)); + + [Theory] + [InlineData(typeof(int[]), typeof(List))] + [InlineData(typeof(List), typeof(List))] + [InlineData(typeof(HashSet), typeof(List))] + [InlineData(typeof(int[]), typeof(HashSet))] + [InlineData(typeof(List), typeof(HashSet))] + [InlineData(typeof(HashSet), typeof(HashSet))] + public void CreateParameter_without_converter_accepts_mutable_collections(Type mappingType, Type valueType) + => CreateParameter_without_converter( + mappingType, Activator.CreateInstance(valueType, new[] { 1, 2 })); + + [Theory] + [InlineData(typeof(int[]))] + [InlineData(typeof(List))] + [InlineData(typeof(HashSet))] + public void CreateParameter_without_converter_accepts_immutable_list(Type mappingType) + => CreateParameter_without_converter(mappingType, ImmutableList.Create(new[] { 1, 2 })); + + [Theory] + [InlineData(typeof(int[]))] + [InlineData(typeof(List))] + [InlineData(typeof(HashSet))] + public void CreateParameter_without_converter_accepts_array(Type mappingType) + { + var mapping = CreateTypeMappingSource().FindMapping(mappingType)!; + Assert.Null(mapping.Converter); + var parameter = mapping.CreateParameter(new NpgsqlCommand(), "p", new[] { 1, 2 }); + Assert.Equal(new[] { 1, 2 }, Assert.IsType(parameter.Value)); + } + + [Fact] + public void CreateParameter_with_shape_only_converter_accepts_non_ilist_collection() + { + var mapping = CreateTypeMappingSource().FindMapping(typeof(IList))!; + Assert.NotNull(mapping.Converter); + Assert.Same(typeof(IList), mapping.Converter.ModelClrType); + Assert.Same(typeof(int[]), mapping.Converter.ProviderClrType); + Assert.Null(mapping.ElementTypeMapping!.Converter); + + var parameter = mapping.CreateParameter(new NpgsqlCommand(), "p", new HashSet { 1, 2 }); + Assert.Equivalent(new[] { 1, 2 }, Assert.IsType(parameter.Value), strict: true); + } + private void Array_over_type_mapping_with_value_converter(CoreTypeMapping mapping, Type expectedType) { var arrayMapping = (NpgsqlArrayTypeMapping)mapping; @@ -288,6 +365,22 @@ private void Array_over_type_mapping_with_value_converter(CoreTypeMapping mappin s => Assert.Equal("bar", s)); } + private void CreateParameter_with_value_converter(Type mappingType, object value) + { + var mapping = CreateTypeMappingSource().FindMapping(mappingType)!; + Assert.NotNull(mapping.Converter); + var parameter = mapping.CreateParameter(new NpgsqlCommand(), "p", value); + Assert.Equivalent(new[] { "foo", "bar" }, Assert.IsType(parameter.Value), strict: true); + } + + private void CreateParameter_without_converter(Type mappingType, object value) + { + var mapping = CreateTypeMappingSource().FindMapping(mappingType)!; + Assert.Null(mapping.Converter); + var parameter = mapping.CreateParameter(new NpgsqlCommand(), "p", value); + Assert.Equivalent(new[] { 1, 2 }, Assert.IsType>(parameter.Value), strict: true); + } + #endregion Array #region JSON