diff --git a/Extensions.Test/DictionaryExtensionsTests.cs b/Extensions.Test/DictionaryExtensionsTests.cs index d9d8aee..216c544 100644 --- a/Extensions.Test/DictionaryExtensionsTests.cs +++ b/Extensions.Test/DictionaryExtensionsTests.cs @@ -90,6 +90,31 @@ public void GetOrCreateConcurrentDictionaryShouldReturnSameInstanceToParallelCal Assert.HasCount(1000, dictionary["key1"]); } + [TestMethod] + public void GetOrCreateConcurrentDictionaryWithoutDefaultShouldReturnStoredValueWhenAnotherCallerAddsFirst() + { + // The comparer adds a rival value for the key the second time it hashes it, which is the + // moment between a lookup that missed and the add that follows it + RacingComparer comparer = new(); + ConcurrentDictionary> dictionary = new(comparer); + List rival = []; + comparer.OnSecondHash = () => dictionary.TryAdd("key1", rival); + + List result = dictionary.GetOrCreate("key1"); + + Assert.AreSame(dictionary["key1"], result); + } + + [TestMethod] + public void GetOrCreateConcurrentDictionaryWithoutDefaultShouldReturnSameInstanceToParallelCallers() + { + ConcurrentDictionary> dictionary = new(); + + Parallel.For(0, 1000, i => dictionary.GetOrCreate("key1").Add(i)); + + Assert.HasCount(1000, dictionary["key1"]); + } + private sealed class RacingComparer : IEqualityComparer { private int hashCount; @@ -125,6 +150,22 @@ public void GetOrCreateShouldThrowArgumentNullExceptionWhenKeyIsNull() Assert.ThrowsExactly(() => dictionary.GetOrCreate(null!)); } + [TestMethod] + public void GetOrCreateConcurrentDictionaryWithoutDefaultShouldThrowArgumentNullExceptionWhenDictionaryIsNull() + { + ConcurrentDictionary? dictionary = null!; + + Assert.ThrowsExactly(() => dictionary.GetOrCreate("key1")); + } + + [TestMethod] + public void GetOrCreateConcurrentDictionaryWithoutDefaultShouldThrowArgumentNullExceptionWhenKeyIsNull() + { + ConcurrentDictionary dictionary = new(); + + Assert.ThrowsExactly(() => dictionary.GetOrCreate(null!)); + } + [TestMethod] public void GetOrCreateShouldThrowArgumentNullExceptionWhenDefaultValueIsNull() { diff --git a/Extensions/DictionaryExtensions.cs b/Extensions/DictionaryExtensions.cs index a219967..21ea1b4 100644 --- a/Extensions/DictionaryExtensions.cs +++ b/Extensions/DictionaryExtensions.cs @@ -60,6 +60,35 @@ public static class DictionaryExtensions return defaultValue; } + /// + /// Method that gets a value from a dictionary if it exists, otherwise creates a new value and adds it to the dictionary. + /// + /// The type of the keys in the dictionary. + /// The type of the values in the dictionary. + /// The dictionary to get the value from. + /// The key to get the value for. + /// The value for the key if it exists, otherwise a new value. + public static TVal GetOrCreate(this ConcurrentDictionary dictionary, TKey key) where TKey : notnull where TVal : new() + { +#pragma warning disable KTSU0004 // Use Ensure.NotNull instead of manual null check + if (dictionary is null) + { + throw new ArgumentNullException(nameof(dictionary), "The dictionary cannot be null."); + } +#pragma warning restore KTSU0004 // Use Ensure.NotNull instead of manual null check + +#pragma warning disable KTSU0004 // Use Ensure.NotNull instead of manual null check + if (key is null) + { + throw new ArgumentNullException(nameof(key), "The key cannot be null."); + } +#pragma warning restore KTSU0004 // Use Ensure.NotNull instead of manual null check + + // Without this overload the call binds to the IDictionary one, whose lookup-then-Add throws + // when another caller adds the key in between. GetOrAdd is atomic. + return dictionary.GetOrAdd(key, _ => new TVal()); + } + /// /// Method that gets a value from a dictionary if it exists, otherwise creates a new value and adds it to the dictionary. ///