diff --git a/src/libraries/System.Collections.Immutable/src/System/Linq/ImmutableArrayExtensions.cs b/src/libraries/System.Collections.Immutable/src/System/Linq/ImmutableArrayExtensions.cs index 6fdd83dd7d4fd1..8c148f159704f9 100644 --- a/src/libraries/System.Collections.Immutable/src/System/Linq/ImmutableArrayExtensions.cs +++ b/src/libraries/System.Collections.Immutable/src/System/Linq/ImmutableArrayExtensions.cs @@ -187,28 +187,39 @@ public static bool SequenceEqual(this ImmutableArray imm /// The type of element contained by the collection. public static bool SequenceEqual(this ImmutableArray immutableArray, IEnumerable items, IEqualityComparer? comparer = null) where TDerived : TBase { - Requires.NotNull(items, nameof(items)); + if (items is ICollection itemsCol) + { + immutableArray.ThrowNullRefIfNotInitialized(); + return Enumerable.SequenceEqual(immutableArray.array, itemsCol, comparer); + } - comparer ??= EqualityComparer.Default; + return Enumerate(immutableArray, items, comparer); - int i = 0; - int n = immutableArray.Length; - foreach (TDerived item in items) + static bool Enumerate(ImmutableArray immutableArray, IEnumerable items, IEqualityComparer? comparer) { - if (i == n) - { - return false; - } + Requires.NotNull(items, nameof(items)); - if (!comparer.Equals(immutableArray[i], item)) + comparer ??= EqualityComparer.Default; + + int i = 0; + int n = immutableArray.Length; + foreach (TDerived item in items) { - return false; + if (i == n) + { + return false; + } + + if (!comparer.Equals(immutableArray[i], item)) + { + return false; + } + + i++; } - i++; + return i == n; } - - return i == n; } /// diff --git a/src/libraries/System.Collections.Immutable/tests/ImmutableArrayExtensionsTest.cs b/src/libraries/System.Collections.Immutable/tests/ImmutableArrayExtensionsTest.cs index 3415fe13ddf49b..4ce49d5e58c12b 100644 --- a/src/libraries/System.Collections.Immutable/tests/ImmutableArrayExtensionsTest.cs +++ b/src/libraries/System.Collections.Immutable/tests/ImmutableArrayExtensionsTest.cs @@ -20,6 +20,14 @@ public class ImmutableArrayExtensionsTest private static readonly ImmutableArray.Builder s_oneElementBuilder = ImmutableArray.Create(1).ToBuilder(); private static readonly ImmutableArray.Builder s_manyElementsBuilder = ImmutableArray.Create(1, 2, 3).ToBuilder(); + private IEnumerable ToEnumerable(ImmutableArray array) + { + foreach (T item in array) + { + yield return item; + } + } + [Fact] public void Select() { @@ -146,6 +154,11 @@ public void SequenceEqual() Assert.False(ImmutableArrayExtensions.SequenceEqual(s_manyElements, (IEnumerable)s_manyElements.Add(1).ToArray(), comparer)); Assert.False(ImmutableArrayExtensions.SequenceEqual(s_manyElements.Add(1), s_manyElements.Add(2).ToArray(), comparer)); Assert.False(ImmutableArrayExtensions.SequenceEqual(s_manyElements.Add(1), (IEnumerable)s_manyElements.Add(2).ToArray(), comparer)); + + Assert.True(ImmutableArrayExtensions.SequenceEqual(s_manyElements, ToEnumerable(s_manyElements), comparer)); + Assert.False(ImmutableArrayExtensions.SequenceEqual(s_manyElements, ToEnumerable(s_oneElement), comparer)); + Assert.False(ImmutableArrayExtensions.SequenceEqual(s_manyElements, ToEnumerable(s_manyElements.Add(1)), comparer)); + Assert.False(ImmutableArrayExtensions.SequenceEqual(s_manyElements.Add(1), ToEnumerable(s_manyElements.Add(2)), comparer)); } Assert.True(ImmutableArrayExtensions.SequenceEqual(s_manyElements, s_manyElements, (a, b) => true)); @@ -164,6 +177,9 @@ public void SequenceEqualEmptyDefault() TestExtensionsMethods.ValidateDefaultThisBehavior(() => ImmutableArrayExtensions.SequenceEqual(s_oneElement, s_emptyDefault)); TestExtensionsMethods.ValidateDefaultThisBehavior(() => ImmutableArrayExtensions.SequenceEqual(s_emptyDefault, s_empty)); TestExtensionsMethods.ValidateDefaultThisBehavior(() => ImmutableArrayExtensions.SequenceEqual(s_emptyDefault, s_emptyDefault)); + TestExtensionsMethods.ValidateDefaultThisBehavior(() => ImmutableArrayExtensions.SequenceEqual(s_oneElement, ToEnumerable(s_emptyDefault))); + TestExtensionsMethods.ValidateDefaultThisBehavior(() => ImmutableArrayExtensions.SequenceEqual(s_emptyDefault, ToEnumerable(s_empty))); + TestExtensionsMethods.ValidateDefaultThisBehavior(() => ImmutableArrayExtensions.SequenceEqual(s_emptyDefault, ToEnumerable(s_emptyDefault))); AssertExtensions.Throws("predicate", () => ImmutableArrayExtensions.SequenceEqual(s_emptyDefault, s_emptyDefault, (Func)null)); } @@ -173,10 +189,22 @@ public void SequenceEqualEmpty() AssertExtensions.Throws("items", () => ImmutableArrayExtensions.SequenceEqual(s_empty, (IEnumerable)null)); Assert.True(ImmutableArrayExtensions.SequenceEqual(s_empty, s_empty)); Assert.True(ImmutableArrayExtensions.SequenceEqual(s_empty, s_empty.ToArray())); + Assert.True(ImmutableArrayExtensions.SequenceEqual(s_empty, ToEnumerable(s_empty))); Assert.True(ImmutableArrayExtensions.SequenceEqual(s_empty, s_empty, (a, b) => true)); Assert.True(ImmutableArrayExtensions.SequenceEqual(s_empty, s_empty, (a, b) => false)); } + [Fact] + public void SequenceEqualSingleElement() + { + Assert.True(ImmutableArrayExtensions.SequenceEqual(s_oneElement, s_oneElement)); + Assert.False(ImmutableArrayExtensions.SequenceEqual(s_oneElement, s_empty)); + Assert.False(ImmutableArrayExtensions.SequenceEqual(s_oneElement, s_manyElements)); + Assert.True(ImmutableArrayExtensions.SequenceEqual(s_oneElement, ToEnumerable(s_oneElement))); + Assert.False(ImmutableArrayExtensions.SequenceEqual(s_oneElement, ToEnumerable(s_empty))); + Assert.False(ImmutableArrayExtensions.SequenceEqual(s_oneElement, ToEnumerable(s_manyElements))); + } + [Fact] public void Aggregate() {