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
44 changes: 44 additions & 0 deletions Extensions.Test/DictionaryExtensionsTests.cs
Original file line number Diff line number Diff line change
Expand Up @@ -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<string, List<int>> dictionary = new(comparer);
List<int> rival = [];
comparer.OnSecondHash = () => dictionary.TryAdd("key1", rival);

List<int> result = dictionary.GetOrCreate("key1", []);

Assert.AreSame(dictionary["key1"], result);
}

[TestMethod]
public void GetOrCreateConcurrentDictionaryShouldReturnSameInstanceToParallelCallers()
{
ConcurrentDictionary<string, ConcurrentBag<int>> dictionary = new();

Parallel.For(0, 1000, i => dictionary.GetOrCreate("key1", []).Add(i));

Assert.HasCount(1000, dictionary["key1"]);
}

private sealed class RacingComparer : IEqualityComparer<string>
{
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()
{
Expand Down
11 changes: 2 additions & 9 deletions Extensions/DictionaryExtensions.cs
Original file line number Diff line number Diff line change
Expand Up @@ -3,7 +3,6 @@
namespace ktsu.Extensions;

using System.Collections.Concurrent;
using System.Diagnostics;

/// <summary>
/// Extension methods for dictionaries.
Expand Down Expand Up @@ -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);
}

/// <summary>
Expand Down
Loading