Skip to content
Merged
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
86 changes: 86 additions & 0 deletions Containers.Test/ContiguousCollectionTests.cs
Original file line number Diff line number Diff line change
Expand Up @@ -528,4 +528,90 @@ public void WorksWithReferenceTypes()
Assert.AreEqual("alpha", collection[1]);
Assert.AreEqual("bravo", collection[2]);
}

[TestMethod]
public void Enumerate_RemovingDuringForeach_Throws()
{
ContiguousCollection<int> collection = [1, 2, 3, 4];

Assert.ThrowsExactly<InvalidOperationException>(() =>
{
foreach (int item in collection)
{
collection.Remove(item);
}
});
}

[TestMethod]
public void Enumerate_AddingDuringForeach_Throws()
{
ContiguousCollection<int> collection = [1, 2];

Assert.ThrowsExactly<InvalidOperationException>(() =>
{
foreach (int item in collection)
{
collection.Add(item);
}
});
}

[TestMethod]
public void Enumerate_ChangesMadeThroughEveryMutator_Throw()
{
Action<ContiguousCollection<int>>[] mutations =
[
c => c.Add(9),
c => c.Insert(0, 9),
c => c.Remove(1),
c => c.RemoveAt(0),
c => c.Clear(),
c => c[0] = 9,
];

foreach (Action<ContiguousCollection<int>> mutate in mutations)
{
ContiguousCollection<int> collection = [1, 2, 3];
using IEnumerator<int> enumerator = collection.GetEnumerator();
Assert.IsTrue(enumerator.MoveNext());

mutate(collection);

Assert.ThrowsExactly<InvalidOperationException>(() => enumerator.MoveNext());
}
}

[TestMethod]
public void Enumerate_ChangedBeforeTheFirstMoveNext_Throws()
{
ContiguousCollection<int> collection = [1, 2, 3];
using IEnumerator<int> enumerator = collection.GetEnumerator();

collection.Add(4);

Assert.ThrowsExactly<InvalidOperationException>(() => enumerator.MoveNext());
}

[TestMethod]
public void Enumerate_ChangedAfterTheLastElement_Throws()
{
ContiguousCollection<int> collection = [1, 2];
using IEnumerator<int> enumerator = collection.GetEnumerator();
Assert.IsTrue(enumerator.MoveNext());
Assert.IsTrue(enumerator.MoveNext());

collection.Add(3);

Assert.ThrowsExactly<InvalidOperationException>(() => enumerator.MoveNext());
}

[TestMethod]
public void Enumerate_Unchanged_VisitsEveryElement()
{
ContiguousCollection<int> collection = [1, 2, 3];
collection.Add(4);

Assert.AreSequenceEqual([1, 2, 3, 4], collection);
}
}
12 changes: 12 additions & 0 deletions Containers.Test/ContiguousMapTests.cs
Original file line number Diff line number Diff line change
Expand Up @@ -593,4 +593,16 @@ public void Entry_KeyAndValue_ExposeConstructorArguments()
Assert.AreEqual(7, entry.Key);
Assert.AreEqual("seven", entry.Value);
}

[TestMethod]
public void Enumerate_RemovingDuringForeach_Throws() =>
MapEnumerationAssertions.RemovingDuringForeachThrows(() => new ContiguousMap<int, string>());

[TestMethod]
public void Enumerate_ChangesMadeThroughEveryMutator_Throw() =>
MapEnumerationAssertions.EveryMutatorInvalidatesEnumerators(() => new ContiguousMap<int, string>());

[TestMethod]
public void Enumerate_Unchanged_VisitsEveryEntry() =>
MapEnumerationAssertions.UnchangedEnumerationVisitsEveryEntry(() => new ContiguousMap<int, string>());
}
63 changes: 63 additions & 0 deletions Containers.Test/ContiguousSetTests.cs
Original file line number Diff line number Diff line change
Expand Up @@ -495,4 +495,67 @@ public void Remove_WithCustomComparer_ThenAdd_KeepsSingleElement()
string[] expectedItems = ["apple"];
Assert.AreSequenceEqual(expectedItems, set);
}

[TestMethod]
public void Enumerate_RemovingDuringForeach_Throws()
{
ContiguousSet<int> set = [1, 2, 3, 4];

Assert.ThrowsExactly<InvalidOperationException>(() =>
{
foreach (int item in set)
{
set.Remove(item);
}
});
}

[TestMethod]
public void Enumerate_ChangesMadeThroughEveryMutator_Throw()
{
Action<ContiguousSet<int>>[] mutations =
[
s => s.Add(9),
s => s.Remove(1),
s => s.Clear(),
];

foreach (Action<ContiguousSet<int>> mutate in mutations)
{
ContiguousSet<int> set = [1, 2, 3];
using IEnumerator<int> enumerator = set.GetEnumerator();
Assert.IsTrue(enumerator.MoveNext());

mutate(set);

Assert.ThrowsExactly<InvalidOperationException>(() => enumerator.MoveNext());
}
}

[TestMethod]
public void Enumerate_AddingADuplicateOrRemovingAMissingItem_DoesNotThrow()
{
// Neither changes the set, so an enumeration in progress is still valid.
ContiguousSet<int> set = [1, 2, 3];
List<int> seen = [];

foreach (int item in set)
{
set.Add(item);
set.Remove(99);
seen.Add(item);
}

Assert.AreSequenceEqual([1, 2, 3], seen);
}

[TestMethod]
public void UnionWith_Itself_DoesNotThrow()
{
ContiguousSet<int> set = [1, 2, 3];

set.UnionWith(set);

Assert.AreSequenceEqual([1, 2, 3], set);
}
}
12 changes: 12 additions & 0 deletions Containers.Test/InsertionOrderMapTests.cs
Original file line number Diff line number Diff line change
Expand Up @@ -424,4 +424,16 @@ public void WorksWithCustomTypes()
string[] expectedKeyOrder = ["charlie", "alpha", "bravo"];
Assert.AreSequenceEqual(expectedKeyOrder, map.Keys);
}

[TestMethod]
public void Enumerate_RemovingDuringForeach_Throws() =>
MapEnumerationAssertions.RemovingDuringForeachThrows(() => new InsertionOrderMap<int, string>());

[TestMethod]
public void Enumerate_ChangesMadeThroughEveryMutator_Throw() =>
MapEnumerationAssertions.EveryMutatorInvalidatesEnumerators(() => new InsertionOrderMap<int, string>());

[TestMethod]
public void Enumerate_Unchanged_VisitsEveryEntry() =>
MapEnumerationAssertions.UnchangedEnumerationVisitsEveryEntry(() => new InsertionOrderMap<int, string>());
}
74 changes: 74 additions & 0 deletions Containers.Test/MapEnumerationAssertions.cs
Original file line number Diff line number Diff line change
@@ -0,0 +1,74 @@
// Copyright (c) 2023-2026 ktsu-dev contributors

namespace ktsu.Containers.Tests;

using Microsoft.VisualStudio.TestTools.UnitTesting;

/// <summary>
/// Enumeration checks shared by every map, which must all behave like <see cref="Dictionary{TKey, TValue}"/>
/// when changed during enumeration.
/// </summary>
internal static class MapEnumerationAssertions
{
private static IDictionary<int, string> Seeded(Func<IDictionary<int, string>> create, int count)
{
IDictionary<int, string> map = create();
for (int key = 1; key <= count; key++)
{
map.Add(key, $"value {key}");
}

return map;
}

internal static void RemovingDuringForeachThrows(Func<IDictionary<int, string>> create)
{
IDictionary<int, string> map = Seeded(create, 4);

Assert.ThrowsExactly<InvalidOperationException>(() =>
{
foreach (KeyValuePair<int, string> pair in map)
{
map.Remove(pair.Key);
}
});
}

internal static void EveryMutatorInvalidatesEnumerators(Func<IDictionary<int, string>> create)
{
Action<IDictionary<int, string>>[] mutations =
[
m => m.Add(9, "nine"),
m => m[9] = "nine",
m => m[1] = "uno",
m => m.Remove(1),
m => m.Remove(new KeyValuePair<int, string>(1, "value 1")),
m => m.Clear(),
];

foreach (Action<IDictionary<int, string>> mutate in mutations)
{
IDictionary<int, string> map = Seeded(create, 2);
using IEnumerator<KeyValuePair<int, string>> pairs = map.GetEnumerator();
using IEnumerator<int> keys = map.Keys.GetEnumerator();
using IEnumerator<string> values = map.Values.GetEnumerator();
Assert.IsTrue(pairs.MoveNext());
Assert.IsTrue(keys.MoveNext());
Assert.IsTrue(values.MoveNext());

mutate(map);

Assert.ThrowsExactly<InvalidOperationException>(() => pairs.MoveNext());
Assert.ThrowsExactly<InvalidOperationException>(() => keys.MoveNext());
Assert.ThrowsExactly<InvalidOperationException>(() => values.MoveNext());
}
}

internal static void UnchangedEnumerationVisitsEveryEntry(Func<IDictionary<int, string>> create)
{
IDictionary<int, string> map = Seeded(create, 2);

Assert.AreSequenceEqual([1, 2], map.Keys);
Assert.AreSequenceEqual(["value 1", "value 2"], map.Values);
}
}
12 changes: 12 additions & 0 deletions Containers.Test/OrderedMapTests.cs
Original file line number Diff line number Diff line change
Expand Up @@ -517,4 +517,16 @@ public void LargeCollection_MaintainsOrder()
Assert.AreEqual(i + 1, keys[i]);
}
}

[TestMethod]
public void Enumerate_RemovingDuringForeach_Throws() =>
Tests.MapEnumerationAssertions.RemovingDuringForeachThrows(() => new OrderedMap<int, string>());

[TestMethod]
public void Enumerate_ChangesMadeThroughEveryMutator_Throw() =>
Tests.MapEnumerationAssertions.EveryMutatorInvalidatesEnumerators(() => new OrderedMap<int, string>());

[TestMethod]
public void Enumerate_Unchanged_VisitsEveryEntry() =>
Tests.MapEnumerationAssertions.UnchangedEnumerationVisitsEveryEntry(() => new OrderedMap<int, string>());
}
17 changes: 16 additions & 1 deletion Containers/ContiguousCollection.cs
Original file line number Diff line number Diff line change
Expand Up @@ -49,6 +49,11 @@ public class ContiguousCollection<T> : ICollection<T>, IReadOnlyList<T>
/// </summary>
private T[] items;

/// <summary>
/// Incremented by every change, so an enumerator can tell the collection changed under it.
/// </summary>
private int version;

/// <summary>
/// The default initial capacity for the collection.
/// </summary>
Expand Down Expand Up @@ -88,6 +93,7 @@ public T this[int index]
ArgumentOutOfRangeException.ThrowIfNegative(index);
ArgumentOutOfRangeException.ThrowIfGreaterThanOrEqual(index, Count);
items[index] = value;
version++;
}
}

Expand Down Expand Up @@ -158,6 +164,7 @@ public void Add(T item)

items[Count] = item;
Count++;
version++;
}

/// <summary>
Expand All @@ -175,6 +182,7 @@ public void Clear()
Array.Clear(items, 0, Count);
}
Count = 0;
version++;
}

/// <summary>
Expand Down Expand Up @@ -241,6 +249,7 @@ public void RemoveAt(int index)
ArgumentOutOfRangeException.ThrowIfGreaterThanOrEqual(index, Count);

Count--;
version++;
if (index < Count)
{
Array.Copy(items, index + 1, items, index, Count - index);
Expand Down Expand Up @@ -294,6 +303,7 @@ public void Insert(int index, T item)

items[index] = item;
Count++;
version++;
}

/// <summary>
Expand Down Expand Up @@ -339,12 +349,17 @@ public void TrimExcess()
/// <remarks>
/// Enumeration benefits from the contiguous memory layout with optimal cache performance.
/// </remarks>
public IEnumerator<T> GetEnumerator()
public IEnumerator<T> GetEnumerator() => Enumerate(version);

private IEnumerator<T> Enumerate(int expected)
{
for (int i = 0; i < Count; i++)
{
Enumeration.ThrowIfModified(expected, version);
yield return items[i];
}

Enumeration.ThrowIfModified(expected, version);
}

/// <summary>
Expand Down
Loading
Loading