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
138 changes: 138 additions & 0 deletions src/libp2p/Libp2p.Protocols.Pubsub.Tests/TtlCacheTests.cs
Original file line number Diff line number Diff line change
@@ -0,0 +1,138 @@
// SPDX-FileCopyrightText: 2026 Demerzel Solutions Limited
// SPDX-License-Identifier: MIT

using NSubstitute;

namespace Nethermind.Libp2p.Protocols.Pubsub.Tests;

[TestFixture]
public class TtlCacheTests
{
[Test]
public void RemoveExpired_RemovesEntriesRegardlessOfKeyOrder()
{
TestTimeProvider clock = new();
using TtlCache<MessageId> cache = new(500, clock);
MessageId expiredHigh = new([0xFF]);
MessageId liveLow = new([0x01]);

cache.Add(expiredHigh);
clock.UtcNow = clock.UtcNow.AddMilliseconds(500);
cache.Add(liveLow);

cache.RemoveExpired(clock.UtcNow);

Assert.Multiple(() =>
{
Assert.That(cache.Count, Is.EqualTo(1));
Assert.That(cache.Contains(expiredHigh), Is.False);
Assert.That(cache.Contains(liveLow), Is.True);
});
}

[TestCase("Contains")]
[TestCase("TryGet")]
[TestCase("Get")]
[TestCase("ToList")]
public void ExpiredEntries_AreNotReturnedBeforeSweeping(string operation)
{
TestTimeProvider clock = new();
using TtlCache<MessageId, string> cache = new(500, clock);
MessageId id = new([0x01]);
cache.Add(id, "value");
clock.UtcNow = clock.UtcNow.AddMilliseconds(500);

// No cache read or sweep may remove the expired entry before the operation under test.
Assert.That(cache.Count, Is.EqualTo(1));
switch (operation)
{
case "Contains":
Assert.That(cache.Contains(id), Is.False);
break;
case "TryGet":
Assert.That(cache.TryGet(id, out string value), Is.False);
Assert.That(value, Is.Null);
break;
case "Get":
Assert.That(cache.Get(id), Is.Null);
break;
case "ToList":
Assert.That(cache.ToList(), Is.Empty);
break;
}
}

[Test]
public void ToList_ReturnsOnlyLiveEntries()
{
TestTimeProvider clock = new();
using TtlCache<MessageId, string> cache = new(500, clock);
cache.Add(new([0x01]), "expired");
clock.UtcNow = clock.UtcNow.AddMilliseconds(500);
cache.Add(new([0x02]), "live");

Assert.That(cache.Count, Is.EqualTo(2));
Assert.That(cache.ToList(), Is.EqualTo(new[] { "live" }));
}

[Test]
public void Add_ReplacesAnExpiredEntry()
{
TestTimeProvider clock = new();
using TtlCache<MessageId, string> cache = new(500, clock);
MessageId id = new([0x01]);
cache.Add(id, "expired");
clock.UtcNow = clock.UtcNow.AddMilliseconds(500);

Assert.That(cache.Count, Is.EqualTo(1));
cache.Add(id, "replacement");

Assert.That(cache.Count, Is.EqualTo(1));
Assert.That(cache.Get(id), Is.EqualTo("replacement"));
}

[Test]
public void Add_DoesNotReplaceOrRefreshALiveEntry()
{
TestTimeProvider clock = new();
using TtlCache<MessageId, string> cache = new(500, clock);
MessageId id = new([0x01]);
cache.Add(id, "original");
clock.UtcNow = clock.UtcNow.AddMilliseconds(250);
cache.Add(id, "replacement");

Assert.That(cache.Get(id), Is.EqualTo("original"));
clock.UtcNow = clock.UtcNow.AddMilliseconds(250);
Assert.That(cache.Contains(id), Is.False);
}

[Test]
public void RemoveExpired_RemovesLaterEntriesAfterClockMovesBackward()
{
TestTimeProvider clock = new();
using TtlCache<MessageId> cache = new(500, clock);
MessageId first = new([0x01]);
MessageId second = new([0x02]);
cache.Add(first);
clock.UtcNow = clock.UtcNow.AddMilliseconds(-250);
cache.Add(second);
clock.UtcNow = clock.UtcNow.AddMilliseconds(500);

cache.RemoveExpired(clock.UtcNow);

Assert.That(cache.Count, Is.EqualTo(1));
Assert.That(cache.Contains(first), Is.True);
Assert.That(cache.Contains(second), Is.False);
}

private sealed class TestTimeProvider : TimeProvider
{
public DateTimeOffset UtcNow { get; set; } = new(2026, 1, 1, 0, 0, 0, TimeSpan.Zero);

public override DateTimeOffset GetUtcNow() => UtcNow;

// Keep the background sweeper dormant so each test controls expiry explicitly.
public override ITimer CreateTimer(TimerCallback callback, object? state, TimeSpan dueTime, TimeSpan period)
=> Substitute.For<ITimer>();
}
}
143 changes: 112 additions & 31 deletions src/libp2p/Libp2p.Protocols.Pubsub/TtlCache.cs
Original file line number Diff line number Diff line change
Expand Up @@ -6,68 +6,149 @@ namespace Nethermind.Libp2p.Protocols.Pubsub;
internal class TtlCache<TKey, TItem> : IDisposable where TKey : notnull
{
private readonly int ttl;
private readonly TimeProvider timeProvider;
private readonly object sync = new();
private readonly Dictionary<TKey, CachedItem> items = [];
private readonly CancellationTokenSource sweeperCancellation = new();
private readonly Task sweeperTask;
private int disposed;
private readonly record struct CachedItem(TItem Item, DateTimeOffset ValidTill);

private struct CachedItem
public TtlCache(int ttl, TimeProvider? timeProvider = null)
{
public TItem Item { get; set; }
public DateTimeOffset ValidTill { get; set; }
ArgumentOutOfRangeException.ThrowIfNegativeOrZero(ttl);
this.ttl = ttl;
this.timeProvider = timeProvider ?? TimeProvider.System;
sweeperTask = Task.Run(async () =>
{
try
{
while (true)
{
await Task.Delay(TimeSpan.FromSeconds(5), this.timeProvider, sweeperCancellation.Token);
RemoveExpired(this.timeProvider.GetUtcNow());
}
}
catch (OperationCanceledException) when (sweeperCancellation.IsCancellationRequested)
{
}
});
}

private readonly SortedDictionary<TKey, CachedItem> items = [];
private bool isDisposed;
public bool Contains(TKey key) => TryGet(key, out _);

public TItem Get(TKey key) => TryGet(key, out TItem item) ? item : default!;

public TtlCache(int ttl)
internal int Count
{
this.ttl = ttl;
Task.Run(async () =>
get
{
while (!isDisposed)
lock (sync)
{
await Task.Delay(5_000);
DateTimeOffset now = DateTimeOffset.UtcNow;
lock (items)
return items.Count;
}
}
}

public bool TryGet(TKey key, out TItem item)
{
lock (sync)
{
if (items.TryGetValue(key, out CachedItem cachedItem))
{
if (cachedItem.ValidTill > timeProvider.GetUtcNow())
{
TKey[] keys = items.TakeWhile(i => i.Value.ValidTill < now).Select(i => i.Key).ToArray();
foreach (TKey keyToRemove in keys)
{
items.Remove(keyToRemove);
}
item = cachedItem.Item;
return true;
}

items.Remove(key);
}
});
}
}

public bool Contains(TKey key) => items.ContainsKey(key);
item = default!;
return false;
}

public TItem Get(TKey key) => items.GetValueOrDefault(key).Item;
internal void RemoveExpired(DateTimeOffset now)
{
lock (sync)
{
RemoveExpiredLocked(now);
}
}

public void Add(TKey key, TItem item)
{
lock (items)
lock (sync)
{
items.TryAdd(key, new CachedItem
DateTimeOffset now = timeProvider.GetUtcNow();
if (items.TryGetValue(key, out CachedItem cachedItem))
{
Item = item,
ValidTill = DateTimeOffset.UtcNow.AddMilliseconds(ttl),
});
if (cachedItem.ValidTill > now)
{
return;
}

items.Remove(key);
}

items.Add(key, new CachedItem(item, now.AddMilliseconds(ttl)));
}
}

private void RemoveExpiredLocked(DateTimeOffset now)
{
// Wall-clock adjustments can make expiration order differ from insertion order.
List<TKey>? expired = null;
foreach ((TKey key, CachedItem item) in items)
{
if (item.ValidTill <= now)
{
(expired ??= []).Add(key);
}
}

if (expired is not null)
{
foreach (TKey key in expired)
{
items.Remove(key);
}
}
}

public void Dispose()
{
isDisposed = true;
if (Interlocked.Exchange(ref disposed, 1) != 0)
{
return;
}

sweeperCancellation.Cancel();
sweeperTask.GetAwaiter().GetResult();
sweeperCancellation.Dispose();

lock (sync)
{
items.Clear();
}
}

internal IList<TItem> ToList()
{
lock (items)
lock (sync)
{
return items.Values.Select(i => i.Item).ToList();
DateTimeOffset now = timeProvider.GetUtcNow();
return items.Values
.Where(item => item.ValidTill > now)
.Select(item => item.Item)
.ToList();
}
}
}

internal class TtlCache<TKey>(int ttl) : TtlCache<TKey, bool>(ttl) where TKey : notnull
internal class TtlCache<TKey>(int ttl, TimeProvider? timeProvider = null) : TtlCache<TKey, bool>(ttl, timeProvider) where TKey : notnull
{
public void Add(TKey key) => Add(key, false);
public void Add(TKey key) => Add(key, true);
}
Loading