diff --git a/Extensions.Test/DictionaryExtensionsTests.cs b/Extensions.Test/DictionaryExtensionsTests.cs index 1a317b4..d9d8aee 100644 --- a/Extensions.Test/DictionaryExtensionsTests.cs +++ b/Extensions.Test/DictionaryExtensionsTests.cs @@ -65,6 +65,50 @@ public void GetOrCreateConcurrentDictionaryShouldAddAndReturnDefaultValue() Assert.AreEqual(99, dictionary["key1"]); } + [TestMethod] + public void GetOrCreateConcurrentDictionaryShouldReturnStoredValueWhenAnotherCallerAddsFirst() + { + // 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 GetOrCreateConcurrentDictionaryShouldReturnSameInstanceToParallelCallers() + { + 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; + + public Action? OnSecondHash { get; set; } + + public bool Equals(string? x, string? y) => string.Equals(x, y, StringComparison.Ordinal); + + public int GetHashCode(string obj) + { + if (++hashCount == 2) + { + OnSecondHash?.Invoke(); + } + + return StringComparer.Ordinal.GetHashCode(obj); + } + } + [TestMethod] public void GetOrCreateShouldThrowArgumentNullExceptionWhenDictionaryIsNull() { diff --git a/Extensions/DictionaryExtensions.cs b/Extensions/DictionaryExtensions.cs index 6ac8804..a219967 100644 --- a/Extensions/DictionaryExtensions.cs +++ b/Extensions/DictionaryExtensions.cs @@ -3,7 +3,6 @@ namespace ktsu.Extensions; using System.Collections.Concurrent; -using System.Diagnostics; /// /// Extension methods for dictionaries. @@ -93,14 +92,8 @@ public static class DictionaryExtensions } #pragma warning restore KTSU0004 // Use Ensure.NotNull instead of manual null check - if (dictionary.TryGetValue(key, out TVal? val)) - { - return val; - } - - bool result = dictionary.TryAdd(key, defaultValue); - Debug.Assert(result); - return defaultValue; + // GetOrAdd is atomic, so when two callers race on a missing key both get the stored value + return dictionary.GetOrAdd(key, defaultValue); } ///