diff --git a/src/Orleans.Core.Abstractions/Runtime/GrainReference.cs b/src/Orleans.Core.Abstractions/Runtime/GrainReference.cs index 13f53615b76..479ba58daf7 100644 --- a/src/Orleans.Core.Abstractions/Runtime/GrainReference.cs +++ b/src/Orleans.Core.Abstractions/Runtime/GrainReference.cs @@ -267,6 +267,18 @@ public class GrainReference : IAddressable, IEquatable, ISpanFor [NonSerialized] private readonly IdSpan _key; + [NonSerialized] + private object? _messageTargetCache; + + internal object? MessageTargetCache + { + get => Volatile.Read(ref _messageTargetCache); + set => Volatile.Write(ref _messageTargetCache, value); + } + + internal void ClearMessageTargetCache(object expected) + => Interlocked.CompareExchange(ref _messageTargetCache, null, expected); + /// /// Gets the grain reference functionality which is shared by all grain references of a given type. /// diff --git a/src/Orleans.Core/Caching/ConcurrentLruCache.cs b/src/Orleans.Core/Caching/ConcurrentLruCache.cs index 5f1aef163cb..0cfee70840d 100644 --- a/src/Orleans.Core/Caching/ConcurrentLruCache.cs +++ b/src/Orleans.Core/Caching/ConcurrentLruCache.cs @@ -64,6 +64,12 @@ public ConcurrentLruCache(int capacity) : this(capacity, comparer: null) { } + protected virtual void UpdateItem(LruItem item, V value) => item.Value = value; + + protected virtual void OnItemRemoved(LruItem item) + { + } + /// /// Initializes a new instance of the ConcurrentLruCore class with the specified capacity and expire-after-access time to live. /// @@ -161,15 +167,25 @@ public V Get(K key) /// public bool TryGet(K key, [MaybeNullWhen(false)] out V value) { - if (_dictionary.TryGetValue(key, out var item)) + if (TryGetItem(key, out var item)) { value = item.Value; - Touch(item); - _telemetryPolicy.IncrementHit(); return true; } value = default; + return false; + } + + protected bool TryGetItem(K key, [NotNullWhen(true)] out LruItem? item) + { + if (_dictionary.TryGetValue(key, out item)) + { + TouchItem(item); + _telemetryPolicy.IncrementHit(); + return true; + } + _telemetryPolicy.IncrementMiss(); return false; } @@ -326,6 +342,7 @@ private void OnRemove(LruItem item, ItemRemovedReason reason) // from the queue. item.WasAccessed = false; item.WasRemoved = true; + OnItemRemoved(item); if (reason == ItemRemovedReason.Evicted) { @@ -351,7 +368,7 @@ public bool TryUpdate(K key, V value) { var oldValue = existing.Value; - existing.Value = value; + UpdateItem(existing, value); UpdateTimestamp(existing); _telemetryPolicy.IncrementUpdated(); @@ -853,7 +870,7 @@ private static ItemDestination RouteCold(LruItem item) /// The value. // NOTE: Internal for testing [DebuggerDisplay("[{Key}] = {Value}")] - internal sealed class LruItem(K key, V value, long timestamp = 0) + internal class LruItem(K key, V value, long timestamp = 0) { private V _data = value; @@ -1057,11 +1074,13 @@ private async Task RunExpirationLoop() } } - private LruItem CreateItem(K key, V value) => - new(key, value, _expiresAfterAccess ? _timeProvider.GetTimestamp() : 0); + protected virtual LruItem CreateItem(K key, V value) => + new(key, value, GetCurrentTimestamp()); + + protected long GetCurrentTimestamp() => _expiresAfterAccess ? _timeProvider.GetTimestamp() : 0; [MethodImpl(MethodImplOptions.AggressiveInlining)] - private void Touch(LruItem item) + protected void TouchItem(LruItem item) { if (_expiresAfterAccess) { diff --git a/src/Orleans.Core/Runtime/GrainReferenceRuntime.cs b/src/Orleans.Core/Runtime/GrainReferenceRuntime.cs index 2cde50e9d0b..56246b243aa 100644 --- a/src/Orleans.Core/Runtime/GrainReferenceRuntime.cs +++ b/src/Orleans.Core/Runtime/GrainReferenceRuntime.cs @@ -107,7 +107,13 @@ public object Cast(IAddressable grain, Type grainInterface) } var interfaceType = this.interfaceTypeResolver.GetGrainInterfaceType(grainInterface); - return this.referenceActivator.CreateReference(grainId, interfaceType); + var result = this.referenceActivator.CreateReference(grainId, interfaceType); + if (grain is GrainReference source) + { + result.MessageTargetCache = source.MessageTargetCache; + } + + return result; } /// diff --git a/src/Orleans.Runtime/Core/InsideRuntimeClient.cs b/src/Orleans.Runtime/Core/InsideRuntimeClient.cs index d605549c37e..c077a7c0f5f 100644 --- a/src/Orleans.Runtime/Core/InsideRuntimeClient.cs +++ b/src/Orleans.Runtime/Core/InsideRuntimeClient.cs @@ -216,7 +216,7 @@ public void SendRequest( } this.messagingTrace.OnSendRequest(message); - this.MessageCenter.AddressAndSendMessage(message); + this.MessageCenter.AddressAndSendMessage(message, target); } public void SendResponse(Message request, Response response) diff --git a/src/Orleans.Runtime/GrainDirectory/CachedGrainLocator.cs b/src/Orleans.Runtime/GrainDirectory/CachedGrainLocator.cs index bc2c1b6e083..a11c68e0bf2 100644 --- a/src/Orleans.Runtime/GrainDirectory/CachedGrainLocator.cs +++ b/src/Orleans.Runtime/GrainDirectory/CachedGrainLocator.cs @@ -218,6 +218,20 @@ public void UpdateCache(GrainId grainId, SiloAddress siloAddress) } public void InvalidateCache(GrainId grainId) => cache.Remove(grainId); public void InvalidateCache(GrainAddress address) => cache.Remove(address); + + internal bool TryGetCacheEntry(GrainId grainId, SiloAddress siloAddress, [NotNullWhen(true)] out GrainDirectoryCacheEntry? entry) + { + if (cache is IGrainDirectoryCacheEntrySource entrySource + && entrySource.TryGetEntry(grainId, out entry) + && entry.Address.SiloAddress?.Equals(siloAddress) == true) + { + return true; + } + + entry = null; + return false; + } + public bool TryLookupInCache(GrainId grainId, [NotNullWhen(true)] out GrainAddress? address) { var grainType = grainId.Type; diff --git a/src/Orleans.Runtime/GrainDirectory/DhtGrainLocator.cs b/src/Orleans.Runtime/GrainDirectory/DhtGrainLocator.cs index 44dd19dd708..d4926d17038 100644 --- a/src/Orleans.Runtime/GrainDirectory/DhtGrainLocator.cs +++ b/src/Orleans.Runtime/GrainDirectory/DhtGrainLocator.cs @@ -81,6 +81,17 @@ public static DhtGrainLocator FromLocalGrainDirectory(LocalGrainDirectory localG public void InvalidateCache(GrainAddress address) => _localGrainDirectory.InvalidateCacheEntry(address); public bool TryLookupInCache(GrainId grainId, [NotNullWhen(true)] out GrainAddress? address) => _localGrainDirectory.TryLocalLookup(grainId, out address); + internal bool TryGetCacheEntry(GrainId grainId, SiloAddress siloAddress, [NotNullWhen(true)] out GrainDirectoryCacheEntry? entry) + { + if (_localGrainDirectory is LocalGrainDirectory directory) + { + return directory.TryGetCacheEntry(grainId, siloAddress, out entry); + } + + entry = null; + return false; + } + private class BatchedDeregistrationWorker { private const int OperationBatchSizeLimit = 2_000; diff --git a/src/Orleans.Runtime/GrainDirectory/GrainDirectoryCacheEntry.cs b/src/Orleans.Runtime/GrainDirectory/GrainDirectoryCacheEntry.cs new file mode 100644 index 00000000000..6539f7678a8 --- /dev/null +++ b/src/Orleans.Runtime/GrainDirectory/GrainDirectoryCacheEntry.cs @@ -0,0 +1,180 @@ +using System.Threading; +using Orleans.Caching; + +namespace Orleans.Runtime.GrainDirectory; + +internal sealed class GrainDirectoryCacheEntry + : ConcurrentLruCache.LruItem, + IDisposable +{ + private static readonly object Invalidated = new(); + private static readonly object Updating = new(); + private readonly LruGrainDirectoryCache? _owner; + private readonly WeakReference _referenceHandle; + private object? _messageTarget; + + public GrainDirectoryCacheEntry(GrainAddress address, int version) + : this(owner: null, address.GrainId, (address, version), timestamp: 0) + { + } + + public GrainDirectoryCacheEntry( + LruGrainDirectoryCache? owner, + GrainId grainId, + (GrainAddress ActivationAddress, int Version) value, + long timestamp) + : base(grainId, value, timestamp) + { + _owner = owner; + _referenceHandle = new(this); + } + + public GrainAddress Address => Value.ActivationAddress; + + public int Version => Value.Version; + + public WeakReference ReferenceHandle => _referenceHandle; + + public bool TryTouch() + { + if (!IsValid) + { + return false; + } + + _owner?.Touch(this); + return IsValid; + } + + public bool IsValid + { + get + { + var target = Volatile.Read(ref _messageTarget); + return !ReferenceEquals(target, Invalidated) && !ReferenceEquals(target, Updating); + } + } + + public bool TryGetMessageTarget(out object? messageTarget) + { + var target = Volatile.Read(ref _messageTarget); + if (ReferenceEquals(target, Invalidated) || ReferenceEquals(target, Updating)) + { + messageTarget = null; + return false; + } + + messageTarget = target; + return messageTarget is not null; + } + + public bool TrySetMessageTarget(object messageTarget, GrainAddress expectedAddress) + { + ArgumentNullException.ThrowIfNull(messageTarget); + ArgumentNullException.ThrowIfNull(expectedAddress); + if (!Address.Matches(expectedAddress) || !TrySetMessageTargetCore(messageTarget)) + { + return false; + } + + if (Address.Matches(expectedAddress)) + { + return true; + } + + ClearMessageTarget(messageTarget); + return false; + } + + public bool TrySetMessageTarget(object messageTarget, SiloAddress expectedSilo) + { + ArgumentNullException.ThrowIfNull(messageTarget); + ArgumentNullException.ThrowIfNull(expectedSilo); + if (Address.SiloAddress?.Equals(expectedSilo) != true || !TrySetMessageTargetCore(messageTarget)) + { + return false; + } + + if (Address.SiloAddress?.Equals(expectedSilo) == true) + { + return true; + } + + ClearMessageTarget(messageTarget); + return false; + } + + private bool TrySetMessageTargetCore(object messageTarget) + { + var current = Volatile.Read(ref _messageTarget); + if (ReferenceEquals(current, Invalidated) || ReferenceEquals(current, Updating)) + { + return false; + } + + return ReferenceEquals(current, messageTarget) + || current is null && Interlocked.CompareExchange(ref _messageTarget, messageTarget, null) is null; + } + + public void ClearMessageTarget(object messageTarget) + { + ArgumentNullException.ThrowIfNull(messageTarget); + Interlocked.CompareExchange(ref _messageTarget, null, messageTarget); + } + + public void Invalidate() => Interlocked.Exchange(ref _messageTarget, Invalidated); + + public void Dispose() => Invalidate(); + + internal void Update((GrainAddress ActivationAddress, int Version) value) + { + var updateStarted = TryBeginUpdate(); + try + { + Value = value; + } + finally + { + if (updateStarted) + { + EndUpdate(); + } + } + } + + public void ClearMessageTarget() + { + while (true) + { + var current = Volatile.Read(ref _messageTarget); + if (current is null || ReferenceEquals(current, Invalidated)) + { + return; + } + + if (ReferenceEquals(Interlocked.CompareExchange(ref _messageTarget, null, current), current)) + { + return; + } + } + } + + internal bool TryBeginUpdate() + { + while (true) + { + var current = Volatile.Read(ref _messageTarget); + if (ReferenceEquals(current, Invalidated)) + { + return false; + } + + if (ReferenceEquals(Interlocked.CompareExchange(ref _messageTarget, Updating, current), current)) + { + return true; + } + } + } + + internal void EndUpdate() => Interlocked.CompareExchange(ref _messageTarget, null, Updating); +} diff --git a/src/Orleans.Runtime/GrainDirectory/GrainLocator.cs b/src/Orleans.Runtime/GrainDirectory/GrainLocator.cs index 85bdcf8afac..4e5355a85ae 100644 --- a/src/Orleans.Runtime/GrainDirectory/GrainLocator.cs +++ b/src/Orleans.Runtime/GrainDirectory/GrainLocator.cs @@ -54,6 +54,26 @@ public GrainLocator(GrainLocatorResolver grainLocatorResolver, DirectoryInstrume public bool TryLookupInCache(GrainId grainId, [NotNullWhen(true)] out GrainAddress? address) => GetGrainLocator(grainId.Type).TryLookupInCache(grainId, out address); + internal bool TryGetCacheEntry( + GrainId grainId, + SiloAddress siloAddress, + [NotNullWhen(true)] out GrainDirectoryCacheEntry? entry) + { + var grainLocator = GetGrainLocator(grainId.Type); + return grainLocator switch + { + CachedGrainLocator cached => cached.TryGetCacheEntry(grainId, siloAddress, out entry), + DhtGrainLocator dht => dht.TryGetCacheEntry(grainId, siloAddress, out entry), + _ => ReturnFalse(out entry), + }; + + static bool ReturnFalse(out GrainDirectoryCacheEntry? result) + { + result = null; + return false; + } + } + public void InvalidateCache(GrainId grainId) => GetGrainLocator(grainId.Type).InvalidateCache(grainId); public void InvalidateCache(GrainAddress address) => GetGrainLocator(address.GrainId.Type).InvalidateCache(address); diff --git a/src/Orleans.Runtime/GrainDirectory/IGrainDirectoryCache.cs b/src/Orleans.Runtime/GrainDirectory/IGrainDirectoryCache.cs index 58201553733..7bd43ab8c11 100644 --- a/src/Orleans.Runtime/GrainDirectory/IGrainDirectoryCache.cs +++ b/src/Orleans.Runtime/GrainDirectory/IGrainDirectoryCache.cs @@ -65,4 +65,9 @@ public static bool LookUp(this IGrainDirectoryCache cache, GrainId key, [NotNull return cache.LookUp(key, out result, out _); } } + + internal interface IGrainDirectoryCacheEntrySource + { + bool TryGetEntry(GrainId key, [NotNullWhen(true)] out GrainDirectoryCacheEntry? entry); + } } diff --git a/src/Orleans.Runtime/GrainDirectory/LocalGrainDirectory.cs b/src/Orleans.Runtime/GrainDirectory/LocalGrainDirectory.cs index 8426509d090..fe5db81b70a 100644 --- a/src/Orleans.Runtime/GrainDirectory/LocalGrainDirectory.cs +++ b/src/Orleans.Runtime/GrainDirectory/LocalGrainDirectory.cs @@ -836,6 +836,19 @@ public bool TryLocalLookup(GrainId grainId, [NotNullWhen(true)] out GrainAddress return IsDefunctActivation(cache, clusterMembershipService.CurrentSnapshot) ? null : cache; } + internal bool TryGetCacheEntry(GrainId grainId, SiloAddress siloAddress, [NotNullWhen(true)] out GrainDirectoryCacheEntry? entry) + { + if (DirectoryCache is IGrainDirectoryCacheEntrySource entrySource + && entrySource.TryGetEntry(grainId, out entry) + && entry.Address.SiloAddress?.Equals(siloAddress) == true) + { + return true; + } + + entry = null; + return false; + } + public async Task LookupAsync(GrainId grainId, int hopCount = 0) { if (hopCount > 0) diff --git a/src/Orleans.Runtime/GrainDirectory/LruGrainDirectoryCache.cs b/src/Orleans.Runtime/GrainDirectory/LruGrainDirectoryCache.cs index c99efdf18b8..ab15c2fb7b0 100644 --- a/src/Orleans.Runtime/GrainDirectory/LruGrainDirectoryCache.cs +++ b/src/Orleans.Runtime/GrainDirectory/LruGrainDirectoryCache.cs @@ -6,7 +6,7 @@ namespace Orleans.Runtime.GrainDirectory; -internal sealed class LruGrainDirectoryCache : ConcurrentLruCache, IGrainDirectoryCache, IAsyncDisposable +internal sealed class LruGrainDirectoryCache : ConcurrentLruCache, IGrainDirectoryCache, IGrainDirectoryCacheEntrySource, IAsyncDisposable { private static readonly Func<(GrainAddress Address, int Version), GrainAddress, bool> ActivationAddressesMatch = (value, state) => GrainAddress.MatchesGrainIdAndSilo(state, value.Address); private readonly IDisposable _cacheSizeRegistration; @@ -58,9 +58,35 @@ public bool LookUp(GrainId key, [NotNullWhen(true)] out GrainAddress? result, ou } } + public bool TryGetEntry(GrainId key, [NotNullWhen(true)] out GrainDirectoryCacheEntry? entry) + { + if (TryGetItem(key, out var item) + && item is GrainDirectoryCacheEntry result + && result.IsValid) + { + entry = result; + return true; + } + + entry = null; + return false; + } + + protected override LruItem CreateItem(GrainId key, (GrainAddress ActivationAddress, int Version) value) + => new GrainDirectoryCacheEntry(this, key, value, GetCurrentTimestamp()); + + internal void Touch(GrainDirectoryCacheEntry entry) => TouchItem(entry); + + protected override void UpdateItem(LruItem item, (GrainAddress ActivationAddress, int Version) value) + => ((GrainDirectoryCacheEntry)item).Update(value); + + protected override void OnItemRemoved(LruItem item) + => ((GrainDirectoryCacheEntry)item).Invalidate(); + public new async ValueTask DisposeAsync() { _cacheSizeRegistration.Dispose(); + Clear(); await base.DisposeAsync(); } diff --git a/src/Orleans.Runtime/Messaging/MessageCenter.cs b/src/Orleans.Runtime/Messaging/MessageCenter.cs index d885ddf352b..86a0ca63073 100644 --- a/src/Orleans.Runtime/Messaging/MessageCenter.cs +++ b/src/Orleans.Runtime/Messaging/MessageCenter.cs @@ -1,6 +1,8 @@ using System; using System.Collections.Generic; using System.Diagnostics; +using System.Diagnostics.CodeAnalysis; +using System.Runtime.CompilerServices; using System.Threading.Tasks; using Microsoft.Extensions.Logging; using Microsoft.Extensions.Options; @@ -141,7 +143,11 @@ public Action? SniffIncomingMessage get => this.sniffIncomingMessageHandler; } - public void SendMessage(Message msg) + public void SendMessage(Message msg) => SendMessage(msg, null); + + private void SendMessage( + Message msg, + GrainDirectoryCacheEntry? directoryCacheEntry) { Debug.Assert(!msg.IsLocalOnly); @@ -186,6 +192,16 @@ public void SendMessage(Message msg) return; } + if (directoryCacheEntry is not null && msg.TargetSilo?.Matches(_siloAddress) == true) + { + Debug.Assert(!msg.TargetGrain.IsClient()); + Debug.Assert(msg.Direction is Message.Directions.Request or Message.Directions.OneWay); + MessagingEvents.EmitSent(msg); + _messagingInstruments.LocalMessagesSentCounterAggregator.Add(1); + ReceiveMessage(msg, directoryCacheEntry); + return; + } + // First check to see if it's really destined for a proxied client, instead of a local grain. if (TryDeliverToProxy(msg)) { @@ -208,7 +224,7 @@ public void SendMessage(Message msg) _messagingInstruments.LocalMessagesSentCounterAggregator.Add(1); - this.ReceiveMessage(msg); + this.ReceiveMessage(msg, directoryCacheEntry); } else { @@ -479,17 +495,25 @@ private static bool MayForward(Message message, SiloMessagingOptions messagingOp /// - add ordering info and maintain send order /// /// - internal Task AddressAndSendMessage(Message message) + internal Task AddressAndSendMessage(Message message, GrainReference? target = null) { try { + if (TryGetDirectoryCacheEntry(target, message, out var directoryCacheEntry)) + { + var targetSilo = directoryCacheEntry.Address.SiloAddress!; + message.TargetSilo = targetSilo; + SendMessage(message, directoryCacheEntry); + return Task.CompletedTask; + } + var messageAddressingTask = placementService.AddressMessage(message); - if (messageAddressingTask.Status != TaskStatus.RanToCompletion) + if (!messageAddressingTask.IsCompletedSuccessfully) { - return SendMessageAsync(messageAddressingTask, message); + return SendMessageAsync(messageAddressingTask, message, target); } - SendMessage(message); + SendMessage(message, CaptureDirectoryCacheEntry(target, message)); } catch (Exception ex) { @@ -498,7 +522,7 @@ internal Task AddressAndSendMessage(Message message) return Task.CompletedTask; - async Task SendMessageAsync(Task addressMessageTask, Message m) + async Task SendMessageAsync(Task addressMessageTask, Message m, GrainReference? target) { try { @@ -510,7 +534,7 @@ async Task SendMessageAsync(Task addressMessageTask, Message m) return; } - SendMessage(m); + SendMessage(m, CaptureDirectoryCacheEntry(target, m)); } void OnAddressingFailure(Message m, Exception ex) @@ -520,6 +544,71 @@ void OnAddressingFailure(Message m, Exception ex) } } + [MethodImpl(MethodImplOptions.AggressiveInlining)] + internal bool TryGetDirectoryCacheEntry( + GrainReference? target, + Message message, + [NotNullWhen(true)] out GrainDirectoryCacheEntry? entry) + { + var cache = target?.MessageTargetCache; + if (cache is not WeakReference handle + || !handle.TryGetTarget(out var candidate)) + { + if (target is not null && cache is not null) + { + target.ClearMessageTargetCache(cache); + } + + entry = null; + return false; + } + + if (message.CacheInvalidationHeader is null && candidate.IsValid) + { + var candidateAddress = candidate.Address; + Debug.Assert(candidateAddress.GrainId.Equals(message.TargetGrain)); + if (candidateAddress.SiloAddress is { } targetSilo + && targetSilo.Matches(_siloAddress) + && candidate.TryTouch()) + { + entry = candidate; + return true; + } + } + + var clearHandle = !candidate.IsValid; + if (!clearHandle) + { + var candidateAddress = candidate.Address; + clearHandle = candidateAddress.SiloAddress is not { } candidateSilo + || !candidateSilo.Matches(_siloAddress); + } + + if (clearHandle) + { + target!.ClearMessageTargetCache(handle); + } + + entry = null; + return false; + } + + private GrainDirectoryCacheEntry? CaptureDirectoryCacheEntry(GrainReference? target, Message message) + { + if (target is null + || message.CacheInvalidationHeader is not null + || message.TargetSilo is not { } targetSilo + || !targetSilo.Matches(_siloAddress) + || !placementService.IsUsingGrainDirectory(message.TargetGrain) + || !_grainLocator.TryGetCacheEntry(message.TargetGrain, targetSilo, out var entry)) + { + return null; + } + + target.MessageTargetCache = entry.ReferenceHandle; + return entry; + } + internal void SendResponse(Message request, Response response) { // create the response @@ -534,13 +623,19 @@ internal void SendResponse(Message request, Response response) SendMessage(message); } - public void ReceiveMessage(Message msg) + public void ReceiveMessage(Message msg) => ReceiveMessage(msg, null); + + private void ReceiveMessage(Message msg, GrainDirectoryCacheEntry? directoryCacheEntry) { Debug.Assert(!msg.IsLocalOnly); try { this.messagingTrace.OnIncomingMessageAgentReceiveMessage(msg); - if (TryDeliverToProxy(msg)) + if (directoryCacheEntry is not null) + { + ReceiveApplicationMessage(msg, directoryCacheEntry); + } + else if (TryDeliverToProxy(msg)) { return; } @@ -550,19 +645,7 @@ public void ReceiveMessage(Message msg) } else { - var targetActivation = catalog.GetOrCreateActivation( - msg.TargetGrain, - msg.RequestContextData, - rehydrationContext: null); - - if (targetActivation is null) - { - ProcessMessageToNonExistentActivation(msg); - return; - } - - targetActivation.ReceiveMessage(msg); - _messageObserver?.Invoke(msg); + ReceiveApplicationMessage(msg, null); } } catch (Exception ex) @@ -579,6 +662,51 @@ void HandleReceiveFailure(Message msg, Exception ex) } } + private void ReceiveApplicationMessage(Message message, GrainDirectoryCacheEntry? directoryCacheEntry) + { + IGrainContext? targetActivation = GetCachedActivation(directoryCacheEntry); + if (targetActivation is null) + { + targetActivation = catalog.GetOrCreateActivation( + message.TargetGrain, + message.RequestContextData, + rehydrationContext: null); + + if (targetActivation is ActivationData { IsValid: true } activation + && directoryCacheEntry is { } entry + && entry.Address.Matches(activation.Address)) + { + entry.TrySetMessageTarget(activation, activation.Address); + } + } + + if (targetActivation is null) + { + ProcessMessageToNonExistentActivation(message); + return; + } + + targetActivation.ReceiveMessage(message); + _messageObserver?.Invoke(message); + } + + [MethodImpl(MethodImplOptions.AggressiveInlining)] + private static ActivationData? GetCachedActivation(GrainDirectoryCacheEntry? directoryCacheEntry) + { + if (directoryCacheEntry?.TryGetMessageTarget(out var target) == true + && target is { } messageTarget) + { + if (messageTarget is ActivationData { IsValid: true } activation) + { + return activation; + } + + directoryCacheEntry.ClearMessageTarget(messageTarget); + } + + return null; + } + private void ProcessMessageToNonExistentActivation(Message msg) { var target = msg.TargetGrain; diff --git a/src/Orleans.Runtime/Placement/PlacementService.cs b/src/Orleans.Runtime/Placement/PlacementService.cs index a491cc7e1b0..09b1ab8b136 100644 --- a/src/Orleans.Runtime/Placement/PlacementService.cs +++ b/src/Orleans.Runtime/Placement/PlacementService.cs @@ -103,6 +103,9 @@ Task ITestAccessor.GetOrPlaceActivationAsync(Message message) private bool IsStopping => _shutdownCts.IsCancellationRequested; + internal bool IsUsingGrainDirectory(GrainId grainId) + => _strategyResolver.GetPlacementStrategy(grainId.Type).IsUsingGrainDirectory; + void ILifecycleParticipant.Participate(ISiloLifecycle lifecycle) { lifecycle.Subscribe( diff --git a/test/Benchmarks/Caching/GrainDirectoryCacheBenchmark.cs b/test/Benchmarks/Caching/GrainDirectoryCacheBenchmark.cs new file mode 100644 index 00000000000..d8891de7796 --- /dev/null +++ b/test/Benchmarks/Caching/GrainDirectoryCacheBenchmark.cs @@ -0,0 +1,127 @@ +using BenchmarkDotNet.Attributes; +using Orleans.Caching; +using Orleans.Runtime; +using Orleans.Runtime.GrainDirectory; + +namespace Benchmarks.Caching; + +[MemoryDiagnoser] +[BenchmarkCategory("Caching")] +public class GrainDirectoryCacheBenchmark +{ + private ConcurrentLruCache _tupleCache = null!; + private LruGrainDirectoryCache _entryCache = null!; + private IGrainDirectoryCacheEntrySource _entrySource = null!; + private WeakReference _entryHandle = null!; + private GrainId _target; + + [GlobalSetup] + public void Setup() + { + _tupleCache = new( + capacity: 1_024, + timeToLive: TimeSpan.FromMinutes(10), + timeProvider: TimeProvider.System); + _entryCache = new( + maxCacheSize: 1_024, + maxCacheTTL: TimeSpan.FromMinutes(10), + timeProvider: TimeProvider.System); + _entrySource = _entryCache; + + for (var i = 0; i < 1_024; i++) + { + var grainId = GrainId.Create("benchmark", i.ToString()); + var address = new GrainAddress + { + GrainId = grainId, + ActivationId = ActivationId.NewId(), + SiloAddress = SiloAddress.FromParsableString("127.0.0.1:11111@1"), + }; + _tupleCache.AddOrUpdate(grainId, (address, i)); + _entryCache.AddOrUpdate(address, i); + if (i == 512) + { + _target = grainId; + } + } + + if (!_entrySource.TryGetEntry(_target, out var entry)) + { + throw new InvalidOperationException("Benchmark entry was not added."); + } + + _entryHandle = entry.ReferenceHandle; + } + + [Benchmark(Baseline = true)] + public GrainAddress TupleValue() + { + _tupleCache.TryGet(_target, out var entry); + return entry.Address; + } + + [Benchmark] + public GrainAddress SharedEntryRaw() + { + _entryCache.TryGet(_target, out var entry); + return entry.ActivationAddress; + } + + [Benchmark] + public GrainAddress SharedEntryValidated() + { + _entryCache.LookUp(_target, out var address, out _); + return address!; + } + + [Benchmark] + public GrainAddress SharedEntrySource() + { + _entrySource.TryGetEntry(_target, out var entry); + return entry!.Address; + } + + [Benchmark] + public GrainAddress SharedHandle() + { + _entryHandle.TryGetTarget(out var entry); + entry!.TryTouch(); + return entry.Address; + } + + [GlobalCleanup] + public async Task Cleanup() + { + await _tupleCache.DisposeAsync(); + await _entryCache.DisposeAsync(); + } +} + +[MemoryDiagnoser] +[BenchmarkCategory("Caching")] +public class GrainDirectoryCacheAllocationBenchmark +{ + private GrainId _grainId; + private GrainAddress _address = null!; + + [GlobalSetup] + public void Setup() + { + _grainId = GrainId.Create("benchmark", "allocation"); + _address = new GrainAddress + { + GrainId = _grainId, + ActivationId = ActivationId.NewId(), + SiloAddress = SiloAddress.FromParsableString("127.0.0.1:11111@1"), + }; + } + + [Benchmark(Baseline = true)] + public object TupleEntry() + => new ConcurrentLruCache.LruItem( + _grainId, + (_address, 1)); + + [Benchmark] + public object SharedEntry() => new GrainDirectoryCacheEntry(_address, version: 1); +} diff --git a/test/Benchmarks/Ping/AdaptiveConcurrencyLoadGenerator.cs b/test/Benchmarks/Ping/AdaptiveConcurrencyLoadGenerator.cs index c072a867690..29fc36aa7b5 100644 --- a/test/Benchmarks/Ping/AdaptiveConcurrencyLoadGenerator.cs +++ b/test/Benchmarks/Ping/AdaptiveConcurrencyLoadGenerator.cs @@ -156,9 +156,58 @@ public async Task RunForeverAsync(CancellationToken cancellationToken = default) break; } } + } - private async Task RunPhaseAsync(TimeSpan duration, bool isWarmup) + public async Task RunFixedConcurrencyAsync( + int repetitions, + CancellationToken cancellationToken = default) + { + ArgumentOutOfRangeException.ThrowIfLessThan(repetitions, 1); + _cts = CancellationTokenSource.CreateLinkedTokenSource(cancellationToken); + var states = Enumerable.Range(0, _currentConcurrency).Select(_getStateForWorker).ToArray(); + + try + { + await RunPhaseAsync(_warmupDuration, isWarmup: true, states); + cancellationToken.ThrowIfCancellationRequested(); + GC.Collect(); + + var results = new FixedConcurrencyMeasurement[repetitions]; + for (var i = 0; i < results.Length; i++) + { + cancellationToken.ThrowIfCancellationRequested(); + var measurement = await RunPhaseAsync(_measurementInterval, isWarmup: false, states); + cancellationToken.ThrowIfCancellationRequested(); + if (measurement.Failures > 0) + { + throw new InvalidOperationException($"Measurement completed with {measurement.Failures:N0} failed requests."); + } + + results[i] = new( + measurement.Throughput, + measurement.LatencyP50Microseconds, + measurement.LatencyP95Microseconds, + measurement.LatencyP99Microseconds, + measurement.AllocatedBytesPerRequest, + measurement.Gen0CollectionsPerMillionRequests, + measurement.CpuUtilization, + measurement.LockContentionsPerMillionRequests); + } + + return results; + } + finally + { + await _cts.CancelAsync(); + _cts.Dispose(); + } + } + + private async Task RunPhaseAsync( + TimeSpan duration, + bool isWarmup, + IReadOnlyList? workerStates = null) { _completedBlocks = Channel.CreateUnbounded( new UnboundedChannelOptions @@ -171,18 +220,15 @@ private async Task RunPhaseAsync(TimeSpan duration, bool isWarmup) var workerTasks = new List(); // Link to main cancellation token so Ctrl+C stops workers immediately using var workerCts = CancellationTokenSource.CreateLinkedTokenSource(_cts.Token); - var states = new Dictionary(); + var aggregator = Task.Run(() => AggregateBlocksAsync(duration, isWarmup, workerCts)); // Start initial workers for (int i = 0; i < _currentConcurrency; i++) { - var state = _getStateForWorker(i); - states[i] = state; + var state = workerStates is null ? _getStateForWorker(i) : workerStates[i]; workerTasks.Add(RunWorkerAsync(state, i, workerCts.Token)); } - var aggregator = Task.Run(() => AggregateBlocksAsync(duration, isWarmup, workerCts)); - var measurement = await aggregator; // Signal workers to stop @@ -220,20 +266,30 @@ private async Task AggregateBlocksAsync(TimeSpan duration, bool isW completionsBySample.Add(0); } - long totalCompleted = 0; + long totalSuccessful = 0; long totalFailures = 0; long minStartTime = long.MaxValue; long maxEndTime = long.MinValue; + var latencySamples = new List(); + var allocatedBytesBefore = GC.GetTotalAllocatedBytes(precise: false); + var gen0CollectionsBefore = GC.CollectionCount(0); + var lockContentionsBefore = Monitor.LockContentionCount; + using var process = Process.GetCurrentProcess(); + var processorTimeBefore = process.TotalProcessorTime; void RecordBlock(WorkBlock block) { - totalCompleted += block.Completed; + totalSuccessful += block.Successes; totalFailures += block.Failures; if (block.StartTimestamp < minStartTime) minStartTime = block.StartTimestamp; if (block.EndTimestamp > maxEndTime) maxEndTime = block.EndTimestamp; + if (block.FirstRequestEndTimestamp > block.StartTimestamp) + { + latencySamples.Add(block.FirstRequestEndTimestamp - block.StartTimestamp); + } var sampleIndex = GetSampleIndex(block.EndTimestamp); - completionsBySample[sampleIndex] += block.Completed; + completionsBySample[sampleIndex] += block.Successes; } int GetSampleIndex(long timestamp) @@ -286,18 +342,32 @@ int GetSampleIndex(long timestamp) RecordBlock(block); } - var totalSeconds = totalCompleted == 0 || maxEndTime <= minStartTime + var totalSeconds = totalSuccessful == 0 || maxEndTime <= minStartTime ? duration.TotalSeconds : (maxEndTime - minStartTime) / StopwatchTickPerSecond; - var throughput = totalSeconds > 0 ? totalCompleted / totalSeconds : 0; + var throughput = totalSeconds > 0 ? totalSuccessful / totalSeconds : 0; var sampleEndTime = maxEndTime > endTime ? maxEndTime : endTime; var samples = CreateThroughputSamples(completionsBySample, startTime, sampleEndTime, sampleTicks); - var measurement = new Measurement(throughput, samples); + var allocatedBytes = GC.GetTotalAllocatedBytes(precise: false) - allocatedBytesBefore; + var gen0Collections = GC.CollectionCount(0) - gen0CollectionsBefore; + var lockContentions = Monitor.LockContentionCount - lockContentionsBefore; + var processorTime = process.TotalProcessorTime - processorTimeBefore; + var measurement = new Measurement( + throughput, + samples, + totalFailures, + GetPercentileMicroseconds(latencySamples, 0.50), + GetPercentileMicroseconds(latencySamples, 0.95), + GetPercentileMicroseconds(latencySamples, 0.99), + totalSuccessful == 0 ? 0 : allocatedBytes / (double)totalSuccessful, + totalSuccessful == 0 ? 0 : gen0Collections * 1_000_000d / totalSuccessful, + totalSeconds <= 0 ? 0 : processorTime.TotalSeconds / (totalSeconds * Environment.ProcessorCount) * 100, + totalSuccessful == 0 ? 0 : lockContentions * 1_000_000d / totalSuccessful); if (isWarmup) { var failureInfo = totalFailures > 0 ? $" ({totalFailures} failures)" : ""; - Console.WriteLine($" Warmup: {throughput:N0}/s, {totalCompleted:N0} requests in {totalSeconds:F1}s{failureInfo}"); + Console.WriteLine($" Warmup: {throughput:N0}/s, {totalSuccessful:N0} requests in {totalSeconds:F1}s{failureInfo}"); } return measurement; @@ -417,6 +487,18 @@ private static double[] CreateThroughputSamples(IReadOnlyList completionsB return samples; } + private static double GetPercentileMicroseconds(List samples, double percentile) + { + if (samples.Count == 0) + { + return 0; + } + + samples.Sort(); + var index = (int)Math.Ceiling(percentile * samples.Count) - 1; + return samples[Math.Clamp(index, 0, samples.Count - 1)] * 1_000_000d / StopwatchTickPerSecond; + } + private bool IsStatisticallySignificantImprovement(Measurement candidate, Measurement baseline) { if (!baseline.HasValue) @@ -500,14 +582,42 @@ private static double GetTwoSidedTCriticalValue95(double degreesOfFreedom) }; } + public readonly record struct FixedConcurrencyMeasurement( + double Throughput, + double LatencyP50Microseconds, + double LatencyP95Microseconds, + double LatencyP99Microseconds, + double AllocatedBytesPerRequest, + double Gen0CollectionsPerMillionRequests, + double CpuUtilization, + double LockContentionsPerMillionRequests); + private readonly struct Measurement { - public Measurement(double throughput, double[] samples) + public Measurement( + double throughput, + double[] samples, + long failures, + double latencyP50Microseconds, + double latencyP95Microseconds, + double latencyP99Microseconds, + double allocatedBytesPerRequest, + double gen0CollectionsPerMillionRequests, + double cpuUtilization, + double lockContentionsPerMillionRequests) { Throughput = throughput; SampleCount = samples.Length; SampleMean = SampleCount == 0 ? throughput : samples.Average(); SampleVariance = CalculateSampleVariance(samples, SampleMean); + Failures = failures; + LatencyP50Microseconds = latencyP50Microseconds; + LatencyP95Microseconds = latencyP95Microseconds; + LatencyP99Microseconds = latencyP99Microseconds; + AllocatedBytesPerRequest = allocatedBytesPerRequest; + Gen0CollectionsPerMillionRequests = gen0CollectionsPerMillionRequests; + CpuUtilization = cpuUtilization; + LockContentionsPerMillionRequests = lockContentionsPerMillionRequests; HasValue = true; } @@ -516,6 +626,14 @@ public Measurement(double throughput, double[] samples) public int SampleCount { get; } public double SampleMean { get; } public double SampleVariance { get; } + public long Failures { get; } + public double LatencyP50Microseconds { get; } + public double LatencyP95Microseconds { get; } + public double LatencyP99Microseconds { get; } + public double AllocatedBytesPerRequest { get; } + public double Gen0CollectionsPerMillionRequests { get; } + public double CpuUtilization { get; } + public double LockContentionsPerMillionRequests { get; } private static double CalculateSampleVariance(double[] samples, double mean) { @@ -558,6 +676,11 @@ private async Task RunWorkerAsync(TState state, int workerId, CancellationToken { workBlock.Failures++; } + + if (workBlock.Completed == 1) + { + workBlock.FirstRequestEndTimestamp = Stopwatch.GetTimestamp(); + } } workBlock.EndTimestamp = Stopwatch.GetTimestamp(); @@ -583,6 +706,7 @@ private async Task RunWorkerAsync(TState state, int workerId, CancellationToken private struct WorkBlock { public long StartTimestamp; + public long FirstRequestEndTimestamp; public long EndTimestamp; public int Successes; public int Failures; diff --git a/test/Benchmarks/Ping/AdaptivePingBenchmark.cs b/test/Benchmarks/Ping/AdaptivePingBenchmark.cs index 0ad87561b74..85f1251b49e 100644 --- a/test/Benchmarks/Ping/AdaptivePingBenchmark.cs +++ b/test/Benchmarks/Ping/AdaptivePingBenchmark.cs @@ -269,4 +269,147 @@ public static async Task RunAllScenariosAsync(int maxStableRounds = DefaultMaxSt Console.WriteLine(); } + + public static async Task RunDeterministicMatrixAsync( + int repetitions = 5, + TimeSpan? warmupDuration = null, + TimeSpan? measurementInterval = null) + { + var scenarios = new (BenchmarkMode Mode, int NumSilos)[] + { + (BenchmarkMode.HostedClient, 1), + (BenchmarkMode.ExternalClient, 1), + (BenchmarkMode.ExternalClient, 2), + (BenchmarkMode.SiloToSilo, 2), + }; + var results = new List(); + + foreach (var (mode, numSilos) in scenarios) + { + results.AddRange(await MeasureDeterministicScenarioAsync(mode, numSilos, repetitions, warmupDuration, measurementInterval)); + + GC.Collect(); + await Task.Delay(1000); + } + + PrintDeterministicResults(results); + } + + public static async Task RunDeterministicScenarioAsync( + BenchmarkMode mode, + int numSilos, + int repetitions = 5, + TimeSpan? warmupDuration = null, + TimeSpan? measurementInterval = null) + { + var results = await MeasureDeterministicScenarioAsync(mode, numSilos, repetitions, warmupDuration, measurementInterval); + PrintDeterministicResults(results); + } + + private static async Task> MeasureDeterministicScenarioAsync( + BenchmarkMode mode, + int numSilos, + int repetitions, + TimeSpan? warmupDuration, + TimeSpan? measurementInterval) + { + int[] concurrencyLevels = [1, 16, 100, 250, 500]; + var results = new List(concurrencyLevels.Length); + using var benchmark = new AdaptivePingBenchmark(mode, numSilos); + try + { + foreach (var concurrency in concurrencyLevels) + { + benchmark._cts.Token.ThrowIfCancellationRequested(); + var loadGenerator = new AdaptiveConcurrencyLoadGenerator( + issueRequest: grain => grain.Run(), + getStateForWorker: workerId => benchmark.GetGrainFactory().GetGrain(workerId), + requestsPerBlock: DefaultRequestsPerBlock, + warmupDuration: warmupDuration ?? TimeSpan.FromSeconds(5), + measurementInterval: measurementInterval ?? TimeSpan.FromSeconds(3), + minConcurrency: concurrency, + maxConcurrency: concurrency, + initialConcurrency: concurrency, + maxStableRounds: 1, + initialStepSize: 1, + sampleInterval: DefaultSampleInterval, + minimumRelativeImprovement: 0); + var measurements = await loadGenerator.RunFixedConcurrencyAsync(repetitions, benchmark._cts.Token); + var throughput = measurements.Select(static measurement => measurement.Throughput).Order().ToArray(); + var result = new DeterministicResult( + benchmark.Description, + concurrency, + GetPercentile(throughput, 0.50), + GetPercentile(throughput, 0.95), + GetPercentile(throughput, 0.99), + GetMedian(measurements, static measurement => measurement.LatencyP50Microseconds), + GetMedian(measurements, static measurement => measurement.LatencyP95Microseconds), + GetMedian(measurements, static measurement => measurement.LatencyP99Microseconds), + GetMedian(measurements, static measurement => measurement.AllocatedBytesPerRequest), + GetMedian(measurements, static measurement => measurement.Gen0CollectionsPerMillionRequests), + GetMedian(measurements, static measurement => measurement.CpuUtilization), + GetMedian(measurements, static measurement => measurement.LockContentionsPerMillionRequests)); + results.Add(result); + Console.WriteLine( + $"{benchmark.Description}, concurrency {concurrency}: " + + $"P50 {result.P50Throughput:N0}/s, P95 {result.P95Throughput:N0}/s, " + + $"first-completion sample {result.LatencyP50Microseconds:N1}/{result.LatencyP95Microseconds:N1}/{result.LatencyP99Microseconds:N1} us"); + } + } + finally + { + await benchmark.ShutdownAsync(); + } + + return results; + } + + private static void PrintDeterministicResults(List results) + { + Console.WriteLine(); + Console.WriteLine("## Deterministic Ping Benchmark Results"); + Console.WriteLine(); + Console.WriteLine("| Scenario | Concurrency | P50 throughput | P95 throughput | P99 throughput | First-completion sample P50/P95/P99 (us) | B/request | Gen0/M requests | CPU | Contentions/M requests |"); + Console.WriteLine("|----------|------------:|---------------:|---------------:|---------------:|-------------------------:|----------:|----------------:|----:|-----------------------:|"); + foreach (var result in results) + { + Console.WriteLine( + $"| {result.Description} | {result.Concurrency} | {result.P50Throughput:N0}/s | " + + $"{result.P95Throughput:N0}/s | {result.P99Throughput:N0}/s | " + + $"{result.LatencyP50Microseconds:N1}/{result.LatencyP95Microseconds:N1}/{result.LatencyP99Microseconds:N1} | " + + $"{result.AllocatedBytesPerRequest:N1} | {result.Gen0CollectionsPerMillionRequests:N2} | " + + $"{result.CpuUtilization:N1}% | {result.LockContentionsPerMillionRequests:N2} |"); + } + + Console.WriteLine(); + } + + private static double GetPercentile(double[] sortedSamples, double percentile) + { + var index = (int)Math.Ceiling(percentile * sortedSamples.Length) - 1; + return sortedSamples[Math.Clamp(index, 0, sortedSamples.Length - 1)]; + } + + private static double GetMedian( + AdaptiveConcurrencyLoadGenerator.FixedConcurrencyMeasurement[] measurements, + Func.FixedConcurrencyMeasurement, double> selector) + { + var values = measurements.Select(selector).Order().ToArray(); + var midpoint = values.Length / 2; + return values.Length % 2 == 0 ? (values[midpoint - 1] + values[midpoint]) / 2 : values[midpoint]; + } + + private readonly record struct DeterministicResult( + string Description, + int Concurrency, + double P50Throughput, + double P95Throughput, + double P99Throughput, + double LatencyP50Microseconds, + double LatencyP95Microseconds, + double LatencyP99Microseconds, + double AllocatedBytesPerRequest, + double Gen0CollectionsPerMillionRequests, + double CpuUtilization, + double LockContentionsPerMillionRequests); } diff --git a/test/Benchmarks/Program.cs b/test/Benchmarks/Program.cs index 1536a92e789..523fe9f3ee2 100644 --- a/test/Benchmarks/Program.cs +++ b/test/Benchmarks/Program.cs @@ -231,6 +231,26 @@ internal class Program { AdaptivePingBenchmark.RunAllScenariosAsync().GetAwaiter().GetResult(); }, + ["DeterministicPing_All"] = _ => + { + AdaptivePingBenchmark.RunDeterministicMatrixAsync().GetAwaiter().GetResult(); + }, + ["DeterministicPing_HostedClient"] = _ => + { + AdaptivePingBenchmark.RunDeterministicScenarioAsync(AdaptivePingBenchmark.BenchmarkMode.HostedClient, numSilos: 1).GetAwaiter().GetResult(); + }, + ["DeterministicPing_ClientToOneSilo"] = _ => + { + AdaptivePingBenchmark.RunDeterministicScenarioAsync(AdaptivePingBenchmark.BenchmarkMode.ExternalClient, numSilos: 1).GetAwaiter().GetResult(); + }, + ["DeterministicPing_ClientToTwoSilos"] = _ => + { + AdaptivePingBenchmark.RunDeterministicScenarioAsync(AdaptivePingBenchmark.BenchmarkMode.ExternalClient, numSilos: 2).GetAwaiter().GetResult(); + }, + ["DeterministicPing_SiloToSilo"] = _ => + { + AdaptivePingBenchmark.RunDeterministicScenarioAsync(AdaptivePingBenchmark.BenchmarkMode.SiloToSilo, numSilos: 2).GetAwaiter().GetResult(); + }, ["ConcurrentPing_OneSilo_Forever"] = _ => { new PingBenchmark(numSilos: 1, startClient: true).PingConcurrentForever().GetAwaiter().GetResult(); diff --git a/test/Orleans.Runtime.Tests/Directories/GrainDirectoryCacheFactoryTests.cs b/test/Orleans.Runtime.Tests/Directories/GrainDirectoryCacheFactoryTests.cs index 12c23e881ce..8ce9cafa67d 100644 --- a/test/Orleans.Runtime.Tests/Directories/GrainDirectoryCacheFactoryTests.cs +++ b/test/Orleans.Runtime.Tests/Directories/GrainDirectoryCacheFactoryTests.cs @@ -1,3 +1,4 @@ +using System.Reflection; using Microsoft.Extensions.DependencyInjection; using Microsoft.Extensions.Time.Testing; using Orleans.Configuration; @@ -256,4 +257,653 @@ private sealed class DisposableGrainDirectoryCache : TestGrainDirectoryCache, ID public void Dispose() => DisposeCalled = true; } + + [Fact] + public async Task CreateGrainDirectoryCache_AddOrUpdateUpdatesSharedRouteHandle() + { + var (cache, entrySource) = CreateEntryCache(); + var disposableCache = Assert.IsAssignableFrom(cache); + var grainId = CreateGrainId(); + var originalAddress = CreateGrainAddress(grainId, port: 11111); + var replacementAddress = CreateGrainAddress(grainId, port: 22222); + var originalTarget = CreateMessageTarget(); + var replacementTarget = CreateMessageTarget(); + + try + { + cache.AddOrUpdate(originalAddress, version: 1); + var originalEntry = GetEntry(entrySource, grainId); + Assert.True(originalEntry.TrySetMessageTarget(originalTarget, originalEntry.Address)); + + cache.AddOrUpdate(replacementAddress, version: 2); + var replacementEntry = GetEntry(entrySource, grainId); + + Assert.Same(originalEntry, replacementEntry); + Assert.True(replacementEntry.IsValid); + Assert.Equal(replacementAddress, replacementEntry.Address); + Assert.Equal(2, replacementEntry.Version); + Assert.False(replacementEntry.TryGetMessageTarget(out _)); + Assert.True(replacementEntry.TrySetMessageTarget(replacementTarget, replacementEntry.Address)); + AssertMessageTarget(replacementEntry, replacementTarget); + Assert.True(cache.LookUp(grainId, out var result, out var version)); + Assert.Equal(replacementAddress, result); + Assert.Equal(2, version); + } + finally + { + await disposableCache.DisposeAsync(); + } + } + + [Fact] + public async Task CreateGrainDirectoryCache_RemoveByGrainIdInvalidatesRouteHandle() + { + var (cache, entrySource) = CreateEntryCache(); + var disposableCache = Assert.IsAssignableFrom(cache); + var address = CreateGrainAddress(CreateGrainId(), port: 11111); + var target = CreateMessageTarget(); + + try + { + cache.AddOrUpdate(address, version: 3); + var entry = GetEntry(entrySource, address.GrainId); + Assert.True(entry.TrySetMessageTarget(target, entry.Address)); + + Assert.True(cache.Remove(address.GrainId)); + + AssertInvalidEntry(entry); + Assert.False(cache.LookUp(address.GrainId, out _, out _)); + Assert.False(cache.Remove(address.GrainId)); + AssertInvalidEntry(entry); + } + finally + { + await disposableCache.DisposeAsync(); + } + } + + [Fact] + public async Task CreateGrainDirectoryCache_RemoveByAddressInvalidatesOnlyMatchingRouteHandle() + { + var (cache, entrySource) = CreateEntryCache(); + var disposableCache = Assert.IsAssignableFrom(cache); + var grainId = CreateGrainId(); + var address = CreateGrainAddress(grainId, port: 11111); + var mismatchedAddress = CreateGrainAddress(grainId, port: 22222); + var controlAddress = CreateGrainAddress(CreateGrainId(), port: 33333); + var target = CreateMessageTarget(); + var controlTarget = CreateMessageTarget(); + + try + { + cache.AddOrUpdate(address, version: 4); + cache.AddOrUpdate(controlAddress, version: 5); + var entry = GetEntry(entrySource, grainId); + var controlEntry = GetEntry(entrySource, controlAddress.GrainId); + Assert.True(entry.TrySetMessageTarget(target, entry.Address)); + Assert.True(controlEntry.TrySetMessageTarget(controlTarget, controlEntry.Address)); + + Assert.False(cache.Remove(mismatchedAddress)); + AssertMessageTarget(entry, target); + Assert.True(cache.LookUp(grainId, out var retainedAddress, out var retainedVersion)); + Assert.Equal(address, retainedAddress); + Assert.Equal(4, retainedVersion); + + Assert.True(cache.Remove(address)); + + AssertInvalidEntry(entry); + Assert.False(cache.LookUp(grainId, out _, out _)); + AssertMessageTarget(controlEntry, controlTarget); + Assert.True(cache.LookUp(controlAddress.GrainId, out var controlResult, out var controlVersion)); + Assert.Equal(controlAddress, controlResult); + Assert.Equal(5, controlVersion); + } + finally + { + await disposableCache.DisposeAsync(); + } + } + + [Fact] + public async Task CreateGrainDirectoryCache_ClearInvalidatesAllRouteHandles() + { + var (cache, entrySource) = CreateEntryCache(); + var disposableCache = Assert.IsAssignableFrom(cache); + var addresses = new[] + { + CreateGrainAddress(CreateGrainId(), port: 11111), + CreateGrainAddress(CreateGrainId(), port: 22222), + CreateGrainAddress(CreateGrainId(), port: 33333) + }; + var entries = new GrainDirectoryCacheEntry[addresses.Length]; + + try + { + for (var i = 0; i < addresses.Length; i++) + { + cache.AddOrUpdate(addresses[i], version: i + 1); + entries[i] = GetEntry(entrySource, addresses[i].GrainId); + Assert.True(entries[i].TrySetMessageTarget(CreateMessageTarget(), entries[i].Address)); + } + + cache.Clear(); + + Assert.Empty(cache.KeyValues); + for (var i = 0; i < addresses.Length; i++) + { + AssertInvalidEntry(entries[i]); + Assert.False(cache.LookUp(addresses[i].GrainId, out _, out _)); + } + + cache.Clear(); + Assert.Empty(cache.KeyValues); + foreach (var entry in entries) + { + AssertInvalidEntry(entry); + } + } + finally + { + await disposableCache.DisposeAsync(); + } + } + + [Fact] + public async Task CreateGrainDirectoryCache_ExpirationInvalidatesRouteHandle() + { + var timeProvider = new FakeTimeProvider(); + var timeToLive = TimeSpan.FromMinutes(1); + var (cache, entrySource) = CreateEntryCache(cacheSize: 10, timeToLive, timeProvider); + var disposableCache = Assert.IsAssignableFrom(cache); + using var listener = new ConcurrentLruCacheExpirationCleanupListener(cache); + var expiredAddress = CreateGrainAddress(CreateGrainId(), port: 11111); + var freshAddress = CreateGrainAddress(CreateGrainId(), port: 22222); + + try + { + cache.AddOrUpdate(expiredAddress, version: 6); + var expiredEntry = GetEntry(entrySource, expiredAddress.GrainId); + Assert.True(expiredEntry.TrySetMessageTarget(CreateMessageTarget(), expiredEntry.Address)); + + timeProvider.Advance(timeToLive); + Assert.Equal(0, await listener.WaitForCleanupAsync()); + + cache.AddOrUpdate(freshAddress, version: 7); + var freshEntry = GetEntry(entrySource, freshAddress.GrainId); + var freshTarget = CreateMessageTarget(); + Assert.True(freshEntry.TrySetMessageTarget(freshTarget, freshEntry.Address)); + + timeProvider.Advance(timeToLive); + Assert.Equal(1, await listener.WaitForCleanupAsync()); + + AssertInvalidEntry(expiredEntry); + Assert.False(cache.LookUp(expiredAddress.GrainId, out _, out _)); + AssertMessageTarget(freshEntry, freshTarget); + Assert.True(cache.LookUp(freshAddress.GrainId, out var result, out var version)); + Assert.Equal(freshAddress, result); + Assert.Equal(7, version); + } + finally + { + await disposableCache.DisposeAsync(); + } + } + + [Fact] + public async Task CreateGrainDirectoryCache_RetainedHandleTouchRefreshesExpiration() + { + var timeProvider = new FakeTimeProvider(); + var timeToLive = TimeSpan.FromMinutes(1); + var (cache, entrySource) = CreateEntryCache(cacheSize: 10, timeToLive, timeProvider); + var disposableCache = Assert.IsAssignableFrom(cache); + using var listener = new ConcurrentLruCacheExpirationCleanupListener(cache); + var address = CreateGrainAddress(CreateGrainId(), port: 11111); + + try + { + cache.AddOrUpdate(address, version: 1); + var entry = GetEntry(entrySource, address.GrainId); + + timeProvider.Advance(timeToLive); + Assert.Equal(0, await listener.WaitForCleanupAsync()); + Assert.True(entry.TryTouch()); + + timeProvider.Advance(timeToLive); + Assert.Equal(0, await listener.WaitForCleanupAsync()); + Assert.True(entry.IsValid); + Assert.Single(cache.KeyValues); + + timeProvider.Advance(timeToLive); + Assert.Equal(1, await listener.WaitForCleanupAsync()); + AssertInvalidEntry(entry); + } + finally + { + await disposableCache.DisposeAsync(); + } + } + + [Fact] + public async Task CreateGrainDirectoryCache_EvictionInvalidatesRouteHandle() + { + var (cache, entrySource) = CreateEntryCache(cacheSize: 3); + var disposableCache = Assert.IsAssignableFrom(cache); + var addresses = new[] + { + CreateGrainAddress(CreateGrainId(), port: 11111), + CreateGrainAddress(CreateGrainId(), port: 22222), + CreateGrainAddress(CreateGrainId(), port: 33333), + CreateGrainAddress(CreateGrainId(), port: 44444) + }; + var entries = new GrainDirectoryCacheEntry[addresses.Length]; + var targets = new IGrainContext[addresses.Length]; + + try + { + for (var i = 0; i < 3; i++) + { + cache.AddOrUpdate(addresses[i], version: i + 1); + entries[i] = GetEntry(entrySource, addresses[i].GrainId); + targets[i] = CreateMessageTarget(); + Assert.True(entries[i].TrySetMessageTarget(targets[i], entries[i].Address)); + } + + cache.AddOrUpdate(addresses[3], version: 4); + entries[3] = GetEntry(entrySource, addresses[3].GrainId); + targets[3] = CreateMessageTarget(); + Assert.True(entries[3].TrySetMessageTarget(targets[3], entries[3].Address)); + + AssertInvalidEntry(entries[0]); + Assert.False(entrySource.TryGetEntry(addresses[0].GrainId, out _)); + for (var i = 1; i < entries.Length; i++) + { + Assert.Same(entries[i], GetEntry(entrySource, addresses[i].GrainId)); + AssertMessageTarget(entries[i], targets[i]); + Assert.True(cache.LookUp(addresses[i].GrainId, out var result, out var version)); + Assert.Equal(addresses[i], result); + Assert.Equal(i + 1, version); + } + + Assert.Equal(3, cache.KeyValues.Count()); + } + finally + { + await disposableCache.DisposeAsync(); + } + } + + [Fact] + public void GrainDirectoryCacheEntry_DisposedEntryCannotBindMessageTarget() + { + var address = CreateGrainAddress(CreateGrainId(), port: 11111); + var entry = new GrainDirectoryCacheEntry(address, version: 8); + + entry.Dispose(); + + AssertInvalidEntry(entry); + Assert.Equal(address, entry.Address); + Assert.Equal(8, entry.Version); + } + + [Fact] + public void GrainDirectoryCacheEntry_SecondTargetCannotReplaceBoundTarget() + { + var entry = new GrainDirectoryCacheEntry(CreateGrainAddress(CreateGrainId(), port: 11111), version: 9); + var originalTarget = CreateMessageTarget(); + + Assert.True(entry.TrySetMessageTarget(originalTarget, entry.Address)); + Assert.False(entry.TrySetMessageTarget(CreateMessageTarget(), entry.Address)); + AssertMessageTarget(entry, originalTarget); + } + + [Fact] + public void GrainDirectoryCacheEntry_ClearRequiresBoundTargetIdentity() + { + var entry = new GrainDirectoryCacheEntry(CreateGrainAddress(CreateGrainId(), port: 11111), version: 10); + var originalTarget = CreateMessageTarget(); + var replacementTarget = CreateMessageTarget(); + Assert.True(entry.TrySetMessageTarget(originalTarget, entry.Address)); + + entry.ClearMessageTarget(CreateMessageTarget()); + AssertMessageTarget(entry, originalTarget); + + entry.ClearMessageTarget(originalTarget); + Assert.False(entry.TryGetMessageTarget(out _)); + Assert.True(entry.TrySetMessageTarget(replacementTarget, entry.Address)); + AssertMessageTarget(entry, replacementTarget); + } + + [Fact] + public async Task GrainDirectoryCacheEntry_StaleAddressCannotRebindAfterUpdate() + { + var (cache, entrySource) = CreateEntryCache(); + var disposableCache = Assert.IsAssignableFrom(cache); + var grainId = CreateGrainId(); + var originalAddress = CreateGrainAddress(grainId, port: 11111); + var replacementAddress = CreateGrainAddress(grainId, port: 22222); + + try + { + cache.AddOrUpdate(originalAddress, version: 1); + var entry = GetEntry(entrySource, grainId); + cache.AddOrUpdate(replacementAddress, version: 2); + + Assert.False(entry.TrySetMessageTarget(CreateMessageTarget(), originalAddress)); + Assert.False(entry.TryGetMessageTarget(out _)); + Assert.Equal(replacementAddress, entry.Address); + } + finally + { + await disposableCache.DisposeAsync(); + } + } + + [Fact] + public void GrainDirectoryCacheEntry_UpdateBlocksBindingUntilNewAddressIsPublished() + { + var grainId = CreateGrainId(); + var originalAddress = CreateGrainAddress(grainId, port: 11111); + var replacementAddress = CreateGrainAddress(grainId, port: 22222); + var entry = new GrainDirectoryCacheEntry(originalAddress, version: 1); + var originalTarget = CreateMessageTarget(); + + Assert.True(entry.TrySetMessageTarget(originalTarget, originalAddress)); + + Assert.True(entry.TryBeginUpdate()); + Assert.False(entry.IsValid); + Assert.False(entry.TryGetMessageTarget(out _)); + Assert.False(entry.TrySetMessageTarget(CreateMessageTarget(), originalAddress)); + + entry.Value = (replacementAddress, 2); + + Assert.False(entry.IsValid); + Assert.False(entry.TrySetMessageTarget(CreateMessageTarget(), replacementAddress)); + + entry.EndUpdate(); + + Assert.True(entry.IsValid); + Assert.Equal(replacementAddress, entry.Address); + Assert.Equal(2, entry.Version); + Assert.False(entry.TryGetMessageTarget(out _)); + Assert.False(entry.TrySetMessageTarget(CreateMessageTarget(), originalAddress)); + + var replacementTarget = CreateMessageTarget(); + Assert.True(entry.TrySetMessageTarget(replacementTarget, replacementAddress)); + AssertMessageTarget(entry, replacementTarget); + } + + [Fact] + public void GrainDirectoryCacheEntry_InvalidationDuringUpdateCannotReopenBinding() + { + var grainId = CreateGrainId(); + var originalAddress = CreateGrainAddress(grainId, port: 11111); + var replacementAddress = CreateGrainAddress(grainId, port: 22222); + var entry = new GrainDirectoryCacheEntry(originalAddress, version: 1); + + Assert.True(entry.TryBeginUpdate()); + + entry.Invalidate(); + entry.Value = (replacementAddress, 2); + entry.EndUpdate(); + + AssertInvalidEntry(entry); + Assert.Equal(replacementAddress, entry.Address); + Assert.Equal(2, entry.Version); + } + + [Fact] + public async Task CreateGrainDirectoryCache_RemoveByGrainIdReleasesMessageTargetReference() + { + var (cache, entrySource) = CreateEntryCache(); + var disposableCache = Assert.IsAssignableFrom(cache); + var address = CreateGrainAddress(CreateGrainId(), port: 11111); + + try + { + cache.AddOrUpdate(address, version: 9); + var entry = GetEntry(entrySource, address.GrainId); + var targetReference = BindMessageTarget(entry); + + Assert.True(cache.Remove(address.GrainId)); + AssertInvalidEntry(entry); + AssertEventuallyCollected(targetReference); + } + finally + { + await disposableCache.DisposeAsync(); + } + } + + [Fact] + public void CreateGrainDirectoryCache_LongLivedHandleDoesNotRetainRemovedEntry() + { + var handle = CreateRemovedEntryHandle(); + + AssertEventuallyCollected(handle); + GC.KeepAlive(handle); + } + + [Fact] + public async Task CreateGrainDirectoryCache_DisposeInvalidatesHandleAndReleasesTarget() + { + var (cache, entrySource) = CreateEntryCache(); + var disposableCache = Assert.IsAssignableFrom(cache); + var address = CreateGrainAddress(CreateGrainId(), port: 11111); + cache.AddOrUpdate(address, version: 1); + var entry = GetEntry(entrySource, address.GrainId); + var targetReference = BindMessageTarget(entry); + + await disposableCache.DisposeAsync(); + + AssertInvalidEntry(entry); + AssertEventuallyCollected(targetReference); + } + + [Fact] + public void CreateGrainDirectoryCache_HighCardinalityHandlesRemainBoundedAndReleaseObjectGraphs() + { + const int cacheSize = 64; + const int entryCount = 5_000; + var (handles, targets) = CreateHighCardinalityHandles(cacheSize, entryCount); + + Assert.Equal(entryCount, handles.Length); + Assert.Equal(entryCount, targets.Length); + AssertEventuallyCollected(handles); + AssertEventuallyCollected(targets); + GC.KeepAlive(handles); + GC.KeepAlive(targets); + } + + private static (IGrainDirectoryCache Cache, IGrainDirectoryCacheEntrySource EntrySource) CreateEntryCache( + int cacheSize = 10, + TimeSpan? timeToLive = null, + FakeTimeProvider? timeProvider = null) + { + var services = new ServiceCollection(); + if (timeProvider is not null) + { + services.AddKeyedSingleton(TimeProviderNames.GrainDirectory, timeProvider); + } + + var cache = GrainDirectoryCacheFactory.CreateGrainDirectoryCache( + services.BuildServiceProvider(), + new GrainDirectoryOptions + { + CacheSize = cacheSize, + MaximumCacheTTL = timeToLive ?? TimeSpan.FromHours(1) + }); + + return (cache, Assert.IsAssignableFrom(cache)); + } + + private static GrainDirectoryCacheEntry GetEntry(IGrainDirectoryCacheEntrySource entrySource, GrainId grainId) + { + Assert.True(entrySource.TryGetEntry(grainId, out var entry)); + return Assert.IsType(entry); + } + + private static void AssertMessageTarget(GrainDirectoryCacheEntry entry, IGrainContext expected) + { + Assert.True(entry.IsValid); + Assert.True(entry.TryGetMessageTarget(out var actual)); + Assert.Same(expected, actual); + } + + private static void AssertInvalidEntry(GrainDirectoryCacheEntry entry) + { + Assert.False(entry.IsValid); + Assert.False(entry.TryGetMessageTarget(out var actual)); + Assert.Null(actual); + Assert.False(entry.TrySetMessageTarget(CreateMessageTarget(), entry.Address)); + } + + [System.Runtime.CompilerServices.MethodImpl(System.Runtime.CompilerServices.MethodImplOptions.NoInlining)] + private static WeakReference BindMessageTarget(GrainDirectoryCacheEntry entry) + { + var target = CreateMessageTarget(); + Assert.True(entry.TrySetMessageTarget(target, entry.Address)); + return new WeakReference(target); + } + + [System.Runtime.CompilerServices.MethodImpl(System.Runtime.CompilerServices.MethodImplOptions.NoInlining)] + private static WeakReference AddRemoveAndGetEntryHandle( + IGrainDirectoryCache cache, + IGrainDirectoryCacheEntrySource entrySource, + GrainAddress address) + { + cache.AddOrUpdate(address, version: 1); + var entry = GetEntry(entrySource, address.GrainId); + var handle = entry.ReferenceHandle; + Assert.True(cache.Remove(address.GrainId)); + Assert.False(handle.TryGetTarget(out var retained) && retained.IsValid); + return handle; + } + + [System.Runtime.CompilerServices.MethodImpl(System.Runtime.CompilerServices.MethodImplOptions.NoInlining)] + private static WeakReference CreateRemovedEntryHandle() + { + var (cache, entrySource) = CreateEntryCache(); + var disposableCache = Assert.IsAssignableFrom(cache); + try + { + return AddRemoveAndGetEntryHandle( + cache, + entrySource, + CreateGrainAddress(CreateGrainId(), port: 11111)); + } + finally + { + disposableCache.DisposeAsync().AsTask().GetAwaiter().GetResult(); + } + } + + [System.Runtime.CompilerServices.MethodImpl(System.Runtime.CompilerServices.MethodImplOptions.NoInlining)] + private static ( + WeakReference[] Handles, + WeakReference[] Targets) CreateHighCardinalityHandles(int cacheSize, int entryCount) + { + var (cache, entrySource) = CreateEntryCache(cacheSize); + var disposableCache = Assert.IsAssignableFrom(cache); + var handles = new WeakReference[entryCount]; + var targets = new WeakReference[entryCount]; + try + { + for (var i = 0; i < entryCount; i++) + { + var address = CreateGrainAddress(CreateGrainId(), port: 11111 + (i % 100)); + cache.AddOrUpdate(address, version: i); + var entry = GetEntry(entrySource, address.GrainId); + var target = new object(); + Assert.True(entry.TrySetMessageTarget(target, address)); + handles[i] = entry.ReferenceHandle; + targets[i] = new(target); + } + + Assert.InRange(cache.KeyValues.Count(), 0, cacheSize); + return (handles, targets); + } + finally + { + disposableCache.DisposeAsync().AsTask().GetAwaiter().GetResult(); + } + } + + [System.Runtime.CompilerServices.MethodImpl(System.Runtime.CompilerServices.MethodImplOptions.NoInlining)] + private static void AssertEventuallyCollected(WeakReference targetReference) + { + for (var attempt = 0; attempt < 5; attempt++) + { + ForceFullCollection(); + + if (!targetReference.TryGetTarget(out _)) + { + return; + } + } + + Assert.False(targetReference.TryGetTarget(out _)); + } + + [System.Runtime.CompilerServices.MethodImpl(System.Runtime.CompilerServices.MethodImplOptions.NoInlining)] + private static void AssertEventuallyCollected(WeakReference entryReference) + { + for (var attempt = 0; attempt < 5; attempt++) + { + ForceFullCollection(); + + if (!entryReference.TryGetTarget(out _)) + { + return; + } + } + + Assert.False(entryReference.TryGetTarget(out _)); + } + + [System.Runtime.CompilerServices.MethodImpl(System.Runtime.CompilerServices.MethodImplOptions.NoInlining)] + private static void AssertEventuallyCollected(WeakReference[] entryReferences) + { + CollectGarbage(); + Assert.All(entryReferences, entryReference => Assert.False(entryReference.TryGetTarget(out _))); + } + + [System.Runtime.CompilerServices.MethodImpl(System.Runtime.CompilerServices.MethodImplOptions.NoInlining)] + private static void AssertEventuallyCollected(WeakReference[] targetReferences) + { + CollectGarbage(); + Assert.All(targetReferences, targetReference => Assert.False(targetReference.TryGetTarget(out _))); + } + + private static void CollectGarbage() + { + for (var attempt = 0; attempt < 5; attempt++) + { + ForceFullCollection(); + } + } + + private static void ForceFullCollection() + { + GC.Collect(GC.MaxGeneration, GCCollectionMode.Forced, blocking: true, compacting: true); + GC.WaitForPendingFinalizers(); + GC.Collect(GC.MaxGeneration, GCCollectionMode.Forced, blocking: true, compacting: true); + } + + private static GrainId CreateGrainId() => GrainId.Parse($"user/{Guid.NewGuid():N}"); + + private static IGrainContext CreateMessageTarget() => DispatchProxy.Create(); + + private class GrainContextProxy : DispatchProxy + { + protected override object? Invoke(MethodInfo? targetMethod, object?[]? args) + => targetMethod?.ReturnType.IsValueType == true ? Activator.CreateInstance(targetMethod.ReturnType) : null; + } + + private static GrainAddress CreateGrainAddress(GrainId grainId, int port) => new() + { + ActivationId = ActivationId.NewId(), + GrainId = grainId, + SiloAddress = SiloAddress.FromParsableString($"127.0.0.1:{port}@1"), + MembershipVersion = new MembershipVersion(1) + }; } diff --git a/test/Orleans.Runtime.Tests/SharedEntryMessageTargetFastPathTests.cs b/test/Orleans.Runtime.Tests/SharedEntryMessageTargetFastPathTests.cs new file mode 100644 index 00000000000..34c34fee187 --- /dev/null +++ b/test/Orleans.Runtime.Tests/SharedEntryMessageTargetFastPathTests.cs @@ -0,0 +1,325 @@ +using Microsoft.Extensions.DependencyInjection; +using Orleans.Hosting; +using Orleans.Runtime; +using Orleans.Runtime.GrainDirectory; +using Orleans.Runtime.Messaging; +using Orleans.Runtime.Placement; +using Orleans.TestingHost; +using TestExtensions; +using UnitTests.GrainInterfaces; +using Xunit; + +namespace UnitTests.Messaging; + +[TestSuite("BVT")] +[TestProvider("None")] +[TestArea("Runtime")] +[TestCategory("BVT"), TestCategory("Messaging")] +public sealed class SharedEntryMessageTargetFastPathTests : IClassFixture +{ + private static long _nextGrainKey = 10_000_000; + private readonly Fixture _fixture; + + public SharedEntryMessageTargetFastPathTests(Fixture fixture) + { + _fixture = fixture; + } + + [Fact] + public async Task LocalDirectoryGrain_AfterTwoCalls_BindsEntryToExactActivation() + { + var primary = (InProcessSiloHandle)_fixture.HostedCluster.Primary!; + var grainFactory = GetPrimarySiloGrainFactory(primary); + var grain = grainFactory.GetGrain(Interlocked.Increment(ref _nextGrainKey)); + var grainReference = Assert.IsAssignableFrom(grain); + const string label = "local-fast-path"; + + RequestContext.Set(IPlacementDirector.PlacementHintKey, primary.SiloAddress); + try + { + await grain.SetLabel(label); + Assert.Equal(label, await grain.GetLabel()); + } + finally + { + RequestContext.Remove(IPlacementDirector.PlacementHintKey); + } + + var entry = GetEntry(grainReference); + Assert.True(entry.IsValid); + Assert.Equal(grainReference.GrainId, entry.Address.GrainId); + Assert.Equal(primary.SiloAddress, entry.Address.SiloAddress); + Assert.True(_fixture.HostedCluster.TryGetGrainContext(grainReference.GrainId, out var grainContext)); + var activation = Assert.IsType(grainContext); + Assert.Equal(activation.Address, entry.Address); + Assert.True(entry.TryGetMessageTarget(out var messageTarget)); + Assert.Same(activation, messageTarget); + } + + [Fact] + public async Task InvalidateCache_DisposesRetainedEntry_AndNextCallCapturesDifferentEntry() + { + var primary = (InProcessSiloHandle)_fixture.HostedCluster.Primary!; + var grainFactory = GetPrimarySiloGrainFactory(primary); + var grain = grainFactory.GetGrain(Interlocked.Increment(ref _nextGrainKey)); + var grainReference = Assert.IsAssignableFrom(grain); + const string label = "entry-before-invalidation"; + + RequestContext.Set(IPlacementDirector.PlacementHintKey, primary.SiloAddress); + try + { + await grain.SetLabel(label); + Assert.Equal(label, await grain.GetLabel()); + } + finally + { + RequestContext.Remove(IPlacementDirector.PlacementHintKey); + } + + var retainedEntry = GetEntry(grainReference); + Assert.True(retainedEntry.IsValid); + primary.ServiceProvider.GetRequiredService().InvalidateCache(grainReference.GrainId); + Assert.False(retainedEntry.IsValid); + Assert.False(retainedEntry.TryGetMessageTarget(out var disposedTarget)); + Assert.Null(disposedTarget); + Assert.Same(retainedEntry.ReferenceHandle, grainReference.MessageTargetCache); + + Assert.Equal(label, await grain.GetLabel()); + + var replacementEntry = GetEntry(grainReference); + Assert.NotSame(retainedEntry, replacementEntry); + Assert.True(replacementEntry.IsValid); + Assert.Equal(grainReference.GrainId, replacementEntry.Address.GrainId); + Assert.Equal(primary.SiloAddress, replacementEntry.Address.SiloAddress); + Assert.True(_fixture.HostedCluster.TryGetGrainContext(grainReference.GrainId, out var grainContext)); + var activation = Assert.IsType(grainContext); + Assert.True(replacementEntry.TryGetMessageTarget(out var replacementTarget)); + Assert.Same(activation, replacementTarget); + } + + [Fact] + public async Task RemoteDirectoryGrain_DoesNotBindConnectionTarget() + { + var primary = (InProcessSiloHandle)_fixture.HostedCluster.Primary!; + var secondary = _fixture.HostedCluster.Silos.Single( + silo => !silo.SiloAddress.Equals(primary.SiloAddress)); + var grainFactory = GetPrimarySiloGrainFactory(primary); + var grain = grainFactory.GetGrain(Interlocked.Increment(ref _nextGrainKey)); + var grainReference = Assert.IsAssignableFrom(grain); + const string label = "remote-fast-path"; + + RequestContext.Set(IPlacementDirector.PlacementHintKey, secondary.SiloAddress); + try + { + await grain.SetLabel(label); + Assert.Equal(label, await grain.GetLabel()); + } + finally + { + RequestContext.Remove(IPlacementDirector.PlacementHintKey); + } + + Assert.Null(grainReference.MessageTargetCache); + } + + [Fact] + public async Task LocalDirectoryGrain_DeactivationDoesNotReuseInvalidActivation() + { + var primary = (InProcessSiloHandle)_fixture.HostedCluster.Primary!; + var grainFactory = GetPrimarySiloGrainFactory(primary); + var grain = grainFactory.GetGrain(Guid.NewGuid()); + var grainReference = Assert.IsAssignableFrom(grain); + + RequestContext.Set(IPlacementDirector.PlacementHintKey, primary.SiloAddress); + try + { + _ = await grain.GetActivationId(); + _ = await grain.GetActivationId(); + } + finally + { + RequestContext.Remove(IPlacementDirector.PlacementHintKey); + } + + var originalEntry = GetEntry(grainReference); + Assert.True(originalEntry.TryGetMessageTarget(out var originalTarget)); + var originalActivation = Assert.IsType(originalTarget); + var originalActivationId = await grain.GetActivationId(); + + await grain.Deactivate(); + RequestContext.Set(IPlacementDirector.PlacementHintKey, primary.SiloAddress); + try + { + var reboundActivationId = await grain.GetActivationId(); + Assert.Equal(reboundActivationId, await grain.GetActivationId()); + + Assert.NotEqual(originalActivationId, reboundActivationId); + Assert.False(originalActivation.IsValid); + Assert.True(_fixture.HostedCluster.TryGetGrainContext(grainReference.GrainId, out var reboundContext)); + var reboundActivation = Assert.IsType(reboundContext); + Assert.NotSame(originalActivation, reboundActivation); + Assert.True(reboundActivation.IsValid); + Assert.Equal(primary.SiloAddress, reboundActivation.Address.SiloAddress); + } + finally + { + RequestContext.Remove(IPlacementDirector.PlacementHintKey); + } + } + + [Fact] + public async Task CompatibleInterfaceCast_TransfersSameEntryHandle() + { + var primary = (InProcessSiloHandle)_fixture.HostedCluster.Primary!; + var grainFactory = GetPrimarySiloGrainFactory(primary); + var writer = grainFactory.GetGrain(Interlocked.Increment(ref _nextGrainKey)); + var writerReference = Assert.IsAssignableFrom(writer); + const int value = 1729; + + RequestContext.Set(IPlacementDirector.PlacementHintKey, primary.SiloAddress); + try + { + await writer.SetValue(-1); + await writer.SetValue(value); + } + finally + { + RequestContext.Remove(IPlacementDirector.PlacementHintKey); + } + + var writerEntry = GetEntry(writerReference); + Assert.True(writerEntry.IsValid); + + var reader = writer.AsReference(); + var readerReference = Assert.IsAssignableFrom(reader); + + Assert.NotSame(writerReference, readerReference); + Assert.Same(writerEntry.ReferenceHandle, readerReference.MessageTargetCache); + Assert.Equal(writerReference.GrainId, readerReference.GrainId); + Assert.Equal(value, await reader.GetValue()); + Assert.Same(writerEntry.ReferenceHandle, readerReference.MessageTargetCache); + } + + [Fact] + public async Task ExternalClientCalls_DoNotAttachMessageTargetCache() + { + var grain = _fixture.Client.GetGrain(Interlocked.Increment(ref _nextGrainKey)); + + await grain.SetLabel("external-client"); + + Assert.Null(Assert.IsAssignableFrom(grain).MessageTargetCache); + } + + [Fact] + public async Task StatelessWorkerCalls_DoNotAttachDirectoryEntries() + { + var primary = (InProcessSiloHandle)_fixture.HostedCluster.Primary!; + var grainFactory = GetPrimarySiloGrainFactory(primary); + var grain = grainFactory.GetGrain(Guid.NewGuid()); + + await grain.Nop(); + + Assert.Null(Assert.IsAssignableFrom(grain).MessageTargetCache); + } + + [Fact] + public async Task CacheInvalidationHeader_BypassesFastPathWithoutDiscardingLiveHandle() + { + var primary = (InProcessSiloHandle)_fixture.HostedCluster.Primary!; + var grainFactory = GetPrimarySiloGrainFactory(primary); + var grain = grainFactory.GetGrain(Interlocked.Increment(ref _nextGrainKey)); + var grainReference = Assert.IsAssignableFrom(grain); + RequestContext.Set(IPlacementDirector.PlacementHintKey, primary.SiloAddress); + try + { + await grain.SetLabel("cache-header"); + Assert.Equal("cache-header", await grain.GetLabel()); + } + finally + { + RequestContext.Remove(IPlacementDirector.PlacementHintKey); + } + + var retainedEntry = GetEntry(grainReference); + var unrelatedAddress = new GrainAddress + { + GrainId = GrainId.Create("unrelated", Guid.NewGuid().ToString()), + SiloAddress = primary.SiloAddress, + }; + var message = new Message + { + Direction = Message.Directions.Request, + TargetGrain = grainReference.GrainId, + CacheInvalidationHeader = [new GrainAddressCacheUpdate(unrelatedAddress, validAddress: null)], + }; + + var messageCenter = primary.ServiceProvider.GetRequiredService(); + Assert.False(messageCenter.TryGetDirectoryCacheEntry(grainReference, message, out var selectedEntry)); + + Assert.Null(selectedEntry); + Assert.True(retainedEntry.IsValid); + Assert.Same(retainedEntry.ReferenceHandle, grainReference.MessageTargetCache); + } + + [Fact] + public void RemoteSiloEntry_IsRejectedAndClearedFromGrainReference() + { + var primary = (InProcessSiloHandle)_fixture.HostedCluster.Primary!; + var grainFactory = GetPrimarySiloGrainFactory(primary); + var grain = grainFactory.GetGrain(Interlocked.Increment(ref _nextGrainKey)); + var grainReference = Assert.IsAssignableFrom(grain); + var address = new GrainAddress + { + GrainId = grainReference.GrainId, + ActivationId = ActivationId.NewId(), + SiloAddress = SiloAddress.FromParsableString("127.0.0.1:54321@1"), + MembershipVersion = new MembershipVersion(0), + }; + var entry = new GrainDirectoryCacheEntry(address, version: 0); + grainReference.MessageTargetCache = entry.ReferenceHandle; + var message = new Message + { + Direction = Message.Directions.Request, + TargetGrain = grainReference.GrainId, + }; + + var messageCenter = primary.ServiceProvider.GetRequiredService(); + Assert.False(messageCenter.TryGetDirectoryCacheEntry(grainReference, message, out var selectedEntry)); + + Assert.Null(selectedEntry); + Assert.Null(grainReference.MessageTargetCache); + Assert.True(entry.IsValid); + } + + private static IGrainFactory GetPrimarySiloGrainFactory(InProcessSiloHandle primary) + { + var runtimeClient = primary.ServiceProvider.GetRequiredService(); + var grainFactory = primary.ServiceProvider.GetRequiredService(); + Assert.Same(runtimeClient.ConcreteGrainFactory, grainFactory); + return grainFactory; + } + + private static GrainDirectoryCacheEntry GetEntry(GrainReference grainReference) + { + var handle = Assert.IsType>(grainReference.MessageTargetCache); + Assert.True(handle.TryGetTarget(out var entry)); + return entry; + } + + public sealed class Fixture : BaseTestClusterFixture + { + protected override void ConfigureTestCluster(TestClusterBuilder builder) + { + builder.Options.InitialSilosCount = 2; + builder.AddSiloBuilderConfigurator(); + } + } + + private sealed class SiloConfigurator : ISiloConfigurator + { + public void Configure(ISiloBuilder siloBuilder) + { + siloBuilder.AddMemoryGrainStorage("MemoryStore"); + } + } +}